MLIR 24.0.0git
SPIRVToLLVM.cpp
Go to the documentation of this file.
1//===- SPIRVToLLVM.cpp - SPIR-V to LLVM Patterns --------------------------===//
2//
3// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.
4// See https://llvm.org/LICENSE.txt for license information.
5// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
6//
7//===----------------------------------------------------------------------===//
8//
9// This file implements patterns to convert SPIR-V dialect to LLVM dialect.
10//
11//===----------------------------------------------------------------------===//
12
20#include "mlir/IR/BuiltinOps.h"
23#include "llvm/ADT/TypeSwitch.h"
24#include "llvm/Support/FormatVariadic.h"
25
26#define DEBUG_TYPE "spirv-to-llvm-pattern"
27
28using namespace mlir;
29
30//===----------------------------------------------------------------------===//
31// Utility functions
32//===----------------------------------------------------------------------===//
33
34/// Returns true if the given type is a signed integer or vector type.
35static bool isSignedIntegerOrVector(Type type) {
36 if (type.isSignedInteger())
37 return true;
38 if (auto vecType = dyn_cast<VectorType>(type))
39 return vecType.getElementType().isSignedInteger();
40 return false;
41}
42
43/// Returns true if the given type is an unsigned integer or vector type
45 if (type.isUnsignedInteger())
46 return true;
47 if (auto vecType = dyn_cast<VectorType>(type))
48 return vecType.getElementType().isUnsignedInteger();
49 return false;
50}
51
52/// Returns the width of an integer or of the element type of an integer vector,
53/// if applicable.
54static std::optional<uint64_t> getIntegerOrVectorElementWidth(Type type) {
55 if (auto intType = dyn_cast<IntegerType>(type))
56 return intType.getWidth();
57 if (auto vecType = dyn_cast<VectorType>(type))
58 if (auto intType = dyn_cast<IntegerType>(vecType.getElementType()))
59 return intType.getWidth();
60 return std::nullopt;
61}
62
63/// Returns the bit width of integer, float or vector of float or integer values
64static unsigned getBitWidth(Type type) {
65 assert((type.isIntOrFloat() || isa<VectorType>(type)) &&
66 "bitwidth is not supported for this type");
67 if (type.isIntOrFloat())
68 return type.getIntOrFloatBitWidth();
69 auto vecType = dyn_cast<VectorType>(type);
70 auto elementType = vecType.getElementType();
71 assert(elementType.isIntOrFloat() &&
72 "only integers and floats have a bitwidth");
73 return elementType.getIntOrFloatBitWidth();
74}
75
76/// Returns the bit width of LLVMType integer or vector.
77static unsigned getLLVMTypeBitWidth(Type type) {
78 if (auto vecTy = dyn_cast<VectorType>(type))
79 type = vecTy.getElementType();
80 return cast<IntegerType>(type).getWidth();
81}
82
83/// Creates `llvm.mlir.constant` with a scalar or vector integer value,
84/// broadcasting `scalarAttr` across the vector if `srcType` is a vector.
85static Value createIntegerConstant(Location loc, Type srcType, Type dstType,
86 PatternRewriter &rewriter,
87 IntegerAttr scalarAttr) {
88 if (auto vecType = dyn_cast<VectorType>(srcType))
89 return LLVM::ConstantOp::create(
90 rewriter, loc, dstType, SplatElementsAttr::get(vecType, scalarAttr));
91 return LLVM::ConstantOp::create(rewriter, loc, dstType, scalarAttr);
92}
93
94/// Creates `llvm.mlir.constant` with all bits set for the given type.
95static Value createConstantAllBitsSet(Location loc, Type srcType, Type dstType,
96 PatternRewriter &rewriter) {
97 auto integerType = cast<IntegerType>(
98 isa<VectorType>(srcType) ? cast<VectorType>(srcType).getElementType()
99 : srcType);
100 return createIntegerConstant(loc, srcType, dstType, rewriter,
101 rewriter.getIntegerAttr(integerType, -1));
102}
103
104/// Creates `llvm.mlir.constant` with a floating-point scalar or vector value.
105static Value createFPConstant(Location loc, Type srcType, Type dstType,
106 PatternRewriter &rewriter, double value) {
107 if (auto vecType = dyn_cast<VectorType>(srcType)) {
108 auto floatType = cast<FloatType>(vecType.getElementType());
109 return LLVM::ConstantOp::create(
110 rewriter, loc, dstType,
112 rewriter.getFloatAttr(floatType, value)));
113 }
114 auto floatType = cast<FloatType>(srcType);
115 return LLVM::ConstantOp::create(rewriter, loc, dstType,
116 rewriter.getFloatAttr(floatType, value));
117}
118
119/// Utility function for bitfield ops:
120/// - `BitFieldInsert`
121/// - `BitFieldSExtract`
122/// - `BitFieldUExtract`
123/// Truncates or extends the value. If the bitwidth of the value is the same as
124/// `llvmType` bitwidth, the value remains unchanged.
126 Type llvmType,
127 PatternRewriter &rewriter) {
128 auto srcType = value.getType();
129 unsigned targetBitWidth = getLLVMTypeBitWidth(llvmType);
130 unsigned valueBitWidth = LLVM::isCompatibleType(srcType)
131 ? getLLVMTypeBitWidth(srcType)
132 : getBitWidth(srcType);
133
134 if (valueBitWidth < targetBitWidth)
135 return LLVM::ZExtOp::create(rewriter, loc, llvmType, value);
136 // If the bit widths of `Count` and `Offset` are greater than the bit width
137 // of the target type, they are truncated. Truncation is safe since `Count`
138 // and `Offset` must be no more than 64 for op behaviour to be defined. Hence,
139 // both values can be expressed in 8 bits.
140 if (valueBitWidth > targetBitWidth)
141 return LLVM::TruncOp::create(rewriter, loc, llvmType, value);
142 return value;
143}
144
145/// Broadcasts the value to vector with `numElements` number of elements.
146static Value broadcast(Location loc, Value toBroadcast, unsigned numElements,
147 const TypeConverter &typeConverter,
148 ConversionPatternRewriter &rewriter) {
149 auto vectorType = VectorType::get(numElements, toBroadcast.getType());
150 auto llvmVectorType = typeConverter.convertType(vectorType);
151 auto llvmI32Type = typeConverter.convertType(rewriter.getIntegerType(32));
152 Value broadcasted = LLVM::PoisonOp::create(rewriter, loc, llvmVectorType);
153 for (unsigned i = 0; i < numElements; ++i) {
154 auto index = LLVM::ConstantOp::create(rewriter, loc, llvmI32Type,
155 rewriter.getI32IntegerAttr(i));
156 broadcasted = LLVM::InsertElementOp::create(
157 rewriter, loc, llvmVectorType, broadcasted, toBroadcast, index);
158 }
159 return broadcasted;
160}
161
162/// Broadcasts the value. If `srcType` is a scalar, the value remains unchanged.
163static Value optionallyBroadcast(Location loc, Value value, Type srcType,
164 const TypeConverter &typeConverter,
165 ConversionPatternRewriter &rewriter) {
166 if (auto vectorType = dyn_cast<VectorType>(srcType)) {
167 unsigned numElements = vectorType.getNumElements();
168 return broadcast(loc, value, numElements, typeConverter, rewriter);
169 }
170 return value;
171}
172
173/// Utility function for bitfield ops: `BitFieldInsert`, `BitFieldSExtract` and
174/// `BitFieldUExtract`.
175/// Broadcast `Offset` and `Count` to match the type of `Base`. If `Base` is of
176/// a vector type, construct a vector that has:
177/// - same number of elements as `Base`
178/// - each element has the type that is the same as the type of `Offset` or
179/// `Count`
180/// - each element has the same value as `Offset` or `Count`
181/// Then cast `Offset` and `Count` if their bit width is different
182/// from `Base` bit width.
183static Value processCountOrOffset(Location loc, Value value, Type srcType,
184 Type dstType, const TypeConverter &converter,
185 ConversionPatternRewriter &rewriter) {
186 Value broadcasted =
187 optionallyBroadcast(loc, value, srcType, converter, rewriter);
188 return optionallyTruncateOrExtend(loc, broadcasted, dstType, rewriter);
189}
190
191/// Converts SPIR-V struct with a regular (according to `VulkanLayoutUtils`)
192/// offset to LLVM struct. Otherwise, the conversion is not supported.
194 const TypeConverter &converter) {
195 if (type != VulkanLayoutUtils::decorateType(type))
196 return nullptr;
197
198 SmallVector<Type> elementsVector;
199 if (failed(converter.convertTypes(type.getElementTypes(), elementsVector)))
200 return nullptr;
201 return LLVM::LLVMStructType::getLiteral(type.getContext(), elementsVector,
202 /*isPacked=*/false);
203}
204
205/// Converts SPIR-V struct with no offset to packed LLVM struct.
207 const TypeConverter &converter) {
208 SmallVector<Type> elementsVector;
209 if (failed(converter.convertTypes(type.getElementTypes(), elementsVector)))
210 return nullptr;
211 return LLVM::LLVMStructType::getLiteral(type.getContext(), elementsVector,
212 /*isPacked=*/true);
213}
214
215/// Creates LLVM dialect constant with the given value.
217 unsigned value) {
218 return LLVM::ConstantOp::create(
219 rewriter, loc, IntegerType::get(rewriter.getContext(), 32),
220 rewriter.getIntegerAttr(rewriter.getI32Type(), value));
221}
222
223/// Utility for `spirv.Load` and `spirv.Store` conversion.
224static LogicalResult replaceWithLoadOrStore(Operation *op, ValueRange operands,
225 ConversionPatternRewriter &rewriter,
226 const TypeConverter &typeConverter,
227 unsigned alignment, bool isVolatile,
228 bool isNonTemporal) {
229 if (auto loadOp = dyn_cast<spirv::LoadOp>(op)) {
230 auto dstType = typeConverter.convertType(loadOp.getType());
231 if (!dstType)
232 return rewriter.notifyMatchFailure(op, "type conversion failed");
233 rewriter.replaceOpWithNewOp<LLVM::LoadOp>(
234 loadOp, dstType, spirv::LoadOpAdaptor(operands).getPtr(), alignment,
235 isVolatile, isNonTemporal);
236 return success();
237 }
238 auto storeOp = cast<spirv::StoreOp>(op);
239 spirv::StoreOpAdaptor adaptor(operands);
240 rewriter.replaceOpWithNewOp<LLVM::StoreOp>(storeOp, adaptor.getValue(),
241 adaptor.getPtr(), alignment,
242 isVolatile, isNonTemporal);
243 return success();
244}
245
246//===----------------------------------------------------------------------===//
247// Type conversion
248//===----------------------------------------------------------------------===//
249
250/// Converts SPIR-V array type to LLVM array. Natural stride (according to
251/// `VulkanLayoutUtils`) is also mapped to LLVM array. This has to be respected
252/// when converting ops that manipulate array types.
253static std::optional<Type> convertArrayType(spirv::ArrayType type,
254 TypeConverter &converter) {
255 unsigned stride = type.getArrayStride();
256 Type elementType = type.getElementType();
257 auto sizeInBytes = cast<spirv::SPIRVType>(elementType).getSizeInBytes();
258 if (stride != 0 && (!sizeInBytes || *sizeInBytes != stride))
259 return std::nullopt;
260
261 auto llvmElementType = converter.convertType(elementType);
262 unsigned numElements = type.getNumElements();
263 return LLVM::LLVMArrayType::get(llvmElementType, numElements);
264}
265
266/// Converts SPIR-V pointer type to LLVM pointer. Pointer's storage class is not
267/// modelled at the moment.
269 const TypeConverter &converter,
270 spirv::ClientAPI clientAPI) {
271 unsigned addressSpace =
273 return LLVM::LLVMPointerType::get(type.getContext(), addressSpace);
274}
275
276/// Converts SPIR-V runtime array to LLVM array. Since LLVM allows indexing over
277/// the bounds, the runtime array is converted to a 0-sized LLVM array. There is
278/// no modelling of array stride at the moment.
279static std::optional<Type> convertRuntimeArrayType(spirv::RuntimeArrayType type,
280 TypeConverter &converter) {
281 if (type.getArrayStride() != 0)
282 return std::nullopt;
283 auto elementType = converter.convertType(type.getElementType());
284 return LLVM::LLVMArrayType::get(elementType, 0);
285}
286
287/// Converts SPIR-V struct to LLVM struct. There is no support of structs with
288/// member decorations. Also, only natural offset is supported.
290 const TypeConverter &converter) {
292 type.getMemberDecorations(memberDecorations);
293 if (!memberDecorations.empty())
294 return nullptr;
295 if (type.hasOffset())
296 return convertStructTypeWithOffset(type, converter);
297 return convertStructTypePacked(type, converter);
298}
299
300//===----------------------------------------------------------------------===//
301// Operation conversion
302//===----------------------------------------------------------------------===//
303
304namespace {
305
306template <typename TargetOp, typename SourceOp>
307static LogicalResult
308collectAttrsForConversion(SourceOp op,
309 typename TargetOp::Properties &properties,
310 SmallVectorImpl<NamedAttribute> &discardableAttrs) {
311 NamedAttrList attrs(op->getDiscardableAttrDictionary());
312 if (auto sourceProperties =
313 dyn_cast_or_null<DictionaryAttr>(op->getPropertiesAsAttribute()))
314 attrs.append(sourceProperties.getValue());
315
316 if constexpr (std::is_same_v<typename TargetOp::Properties,
318 llvm::append_range(discardableAttrs, attrs);
319 return success();
320 }
321
322 TargetOp::populateDefaultProperties(
323 OperationName(TargetOp::getOperationName(), op->getContext()),
324 properties);
325 if (failed(TargetOp::setPropertiesFromAttr(
326 properties, attrs.getDictionary(op->getContext()),
327 [&]() { return op.emitError("failed to convert properties"); })))
328 return failure();
329
330 auto convertedProperties = dyn_cast_or_null<DictionaryAttr>(
331 TargetOp::getPropertiesAsAttr(op->getContext(), properties));
332 for (NamedAttribute attr : attrs) {
333 StringRef name = attr.getName().getValue();
334 bool isProperty = convertedProperties && convertedProperties.contains(name);
335 if (name == "operand_segment_sizes")
336 isProperty |= convertedProperties &&
337 convertedProperties.contains("operandSegmentSizes");
338 if (name == "result_segment_sizes")
339 isProperty |= convertedProperties &&
340 convertedProperties.contains("resultSegmentSizes");
341 if (!isProperty)
342 discardableAttrs.push_back(attr);
343 }
344 return success();
345}
346
347class AccessChainPattern : public SPIRVToLLVMConversion<spirv::AccessChainOp> {
348public:
349 using SPIRVToLLVMConversion<spirv::AccessChainOp>::SPIRVToLLVMConversion;
350
351 LogicalResult
352 matchAndRewrite(spirv::AccessChainOp op, OpAdaptor adaptor,
353 ConversionPatternRewriter &rewriter) const override {
354 auto dstType =
355 getTypeConverter()->convertType(op.getComponentPtr().getType());
356 if (!dstType)
357 return rewriter.notifyMatchFailure(op, "type conversion failed");
358 // To use GEP we need to add a first 0 index to go through the pointer.
359 auto indices = llvm::to_vector<4>(adaptor.getIndices());
360 Type indexType = op.getIndices().front().getType();
361 auto llvmIndexType = getTypeConverter()->convertType(indexType);
362 if (!llvmIndexType)
363 return rewriter.notifyMatchFailure(op, "type conversion failed");
364 Value zero =
365 LLVM::ConstantOp::create(rewriter, op.getLoc(), llvmIndexType,
366 rewriter.getIntegerAttr(llvmIndexType, 0));
367 indices.insert(indices.begin(), zero);
368
369 auto elementType = getTypeConverter()->convertType(
370 cast<spirv::PointerType>(op.getBasePtr().getType()).getPointeeType());
371 if (!elementType)
372 return rewriter.notifyMatchFailure(op, "type conversion failed");
373 rewriter.replaceOpWithNewOp<LLVM::GEPOp>(op, dstType, elementType,
374 adaptor.getBasePtr(), indices);
375 return success();
376 }
377};
378
379class AddressOfPattern : public SPIRVToLLVMConversion<spirv::AddressOfOp> {
380public:
381 using SPIRVToLLVMConversion<spirv::AddressOfOp>::SPIRVToLLVMConversion;
382
383 LogicalResult
384 matchAndRewrite(spirv::AddressOfOp op, OpAdaptor adaptor,
385 ConversionPatternRewriter &rewriter) const override {
386 auto dstType = getTypeConverter()->convertType(op.getPointer().getType());
387 if (!dstType)
388 return rewriter.notifyMatchFailure(op, "type conversion failed");
389 rewriter.replaceOpWithNewOp<LLVM::AddressOfOp>(op, dstType,
390 op.getVariable());
391 return success();
392 }
393};
394
395class BitFieldInsertPattern
396 : public SPIRVToLLVMConversion<spirv::BitFieldInsertOp> {
397public:
398 using SPIRVToLLVMConversion<spirv::BitFieldInsertOp>::SPIRVToLLVMConversion;
399
400 LogicalResult
401 matchAndRewrite(spirv::BitFieldInsertOp op, OpAdaptor adaptor,
402 ConversionPatternRewriter &rewriter) const override {
403 auto srcType = op.getType();
404 auto dstType = getTypeConverter()->convertType(srcType);
405 if (!dstType)
406 return rewriter.notifyMatchFailure(op, "type conversion failed");
407 Location loc = op.getLoc();
408
409 // Process `Offset` and `Count`: broadcast and extend/truncate if needed.
410 Value offset = processCountOrOffset(loc, op.getOffset(), srcType, dstType,
411 *getTypeConverter(), rewriter);
412 Value count = processCountOrOffset(loc, op.getCount(), srcType, dstType,
413 *getTypeConverter(), rewriter);
414
415 // Create a mask with bits set outside [Offset, Offset + Count - 1].
416 Value minusOne = createConstantAllBitsSet(loc, srcType, dstType, rewriter);
417 Value maskShiftedByCount =
418 LLVM::ShlOp::create(rewriter, loc, dstType, minusOne, count);
419 Value negated = LLVM::XOrOp::create(rewriter, loc, dstType,
420 maskShiftedByCount, minusOne);
421 Value maskShiftedByCountAndOffset =
422 LLVM::ShlOp::create(rewriter, loc, dstType, negated, offset);
423 Value mask = LLVM::XOrOp::create(rewriter, loc, dstType,
424 maskShiftedByCountAndOffset, minusOne);
425
426 // Extract unchanged bits from the `Base` that are outside of
427 // [Offset, Offset + Count - 1]. Then `or` with shifted `Insert`.
428 Value baseAndMask =
429 LLVM::AndOp::create(rewriter, loc, dstType, op.getBase(), mask);
430 Value insertShiftedByOffset =
431 LLVM::ShlOp::create(rewriter, loc, dstType, op.getInsert(), offset);
432 rewriter.replaceOpWithNewOp<LLVM::OrOp>(op, dstType, baseAndMask,
433 insertShiftedByOffset);
434 return success();
435 }
436};
437
438/// Converts SPIR-V ConstantOp with scalar or vector type.
439class ConstantScalarAndVectorPattern
440 : public SPIRVToLLVMConversion<spirv::ConstantOp> {
441public:
442 using SPIRVToLLVMConversion<spirv::ConstantOp>::SPIRVToLLVMConversion;
443
444 LogicalResult
445 matchAndRewrite(spirv::ConstantOp constOp, OpAdaptor adaptor,
446 ConversionPatternRewriter &rewriter) const override {
447 auto srcType = constOp.getType();
448 if (!isa<VectorType>(srcType) && !srcType.isIntOrFloat())
449 return failure();
450
451 auto dstType = getTypeConverter()->convertType(srcType);
452 if (!dstType)
453 return rewriter.notifyMatchFailure(constOp, "type conversion failed");
454
455 // SPIR-V constant can be a signed/unsigned integer, which has to be
456 // casted to signless integer when converting to LLVM dialect. Removing the
457 // sign bit may have unexpected behaviour. However, it is better to handle
458 // it case-by-case, given that the purpose of the conversion is not to
459 // cover all possible corner cases.
460 if (isSignedIntegerOrVector(srcType) ||
461 isUnsignedIntegerOrVector(srcType)) {
462 auto signlessType = rewriter.getIntegerType(getBitWidth(srcType));
463
464 if (isa<VectorType>(srcType)) {
465 auto dstElementsAttr = cast<DenseIntElementsAttr>(constOp.getValue());
466 rewriter.replaceOpWithNewOp<LLVM::ConstantOp>(
467 constOp, dstType,
468 dstElementsAttr.mapValues(
469 signlessType, [&](const APInt &value) { return value; }));
470 return success();
471 }
472 auto srcAttr = cast<IntegerAttr>(constOp.getValue());
473 auto dstAttr = rewriter.getIntegerAttr(signlessType, srcAttr.getValue());
474 rewriter.replaceOpWithNewOp<LLVM::ConstantOp>(constOp, dstType, dstAttr);
475 return success();
476 }
477 LLVM::ConstantOp::Properties properties{};
478 SmallVector<NamedAttribute> discardableAttrs;
479 if (failed(collectAttrsForConversion<LLVM::ConstantOp>(constOp, properties,
480 discardableAttrs)))
481 return failure();
482 rewriter.replaceOpWithNewOp<LLVM::ConstantOp>(
483 constOp, dstType, adaptor.getOperands(), properties, discardableAttrs);
484 return success();
485 }
486};
487
488class BitFieldSExtractPattern
489 : public SPIRVToLLVMConversion<spirv::BitFieldSExtractOp> {
490public:
491 using SPIRVToLLVMConversion<spirv::BitFieldSExtractOp>::SPIRVToLLVMConversion;
492
493 LogicalResult
494 matchAndRewrite(spirv::BitFieldSExtractOp op, OpAdaptor adaptor,
495 ConversionPatternRewriter &rewriter) const override {
496 auto srcType = op.getType();
497 auto dstType = getTypeConverter()->convertType(srcType);
498 if (!dstType)
499 return rewriter.notifyMatchFailure(op, "type conversion failed");
500 Location loc = op.getLoc();
501
502 // Process `Offset` and `Count`: broadcast and extend/truncate if needed.
503 Value offset = processCountOrOffset(loc, op.getOffset(), srcType, dstType,
504 *getTypeConverter(), rewriter);
505 Value count = processCountOrOffset(loc, op.getCount(), srcType, dstType,
506 *getTypeConverter(), rewriter);
507
508 // Create a constant that holds the size of the `Base`.
509 IntegerType integerType;
510 if (auto vecType = dyn_cast<VectorType>(srcType))
511 integerType = cast<IntegerType>(vecType.getElementType());
512 else
513 integerType = cast<IntegerType>(srcType);
514
515 auto baseSize = rewriter.getIntegerAttr(integerType, getBitWidth(srcType));
516 Value size =
517 isa<VectorType>(srcType)
518 ? LLVM::ConstantOp::create(
519 rewriter, loc, dstType,
520 SplatElementsAttr::get(cast<ShapedType>(srcType), baseSize))
521 : LLVM::ConstantOp::create(rewriter, loc, dstType, baseSize);
522
523 // Shift `Base` left by [sizeof(Base) - (Count + Offset)], so that the bit
524 // at Offset + Count - 1 is the most significant bit now.
525 Value countPlusOffset =
526 LLVM::AddOp::create(rewriter, loc, dstType, count, offset);
527 Value amountToShiftLeft =
528 LLVM::SubOp::create(rewriter, loc, dstType, size, countPlusOffset);
529 Value baseShiftedLeft = LLVM::ShlOp::create(
530 rewriter, loc, dstType, op.getBase(), amountToShiftLeft);
531
532 // Shift the result right, filling the bits with the sign bit.
533 Value amountToShiftRight =
534 LLVM::AddOp::create(rewriter, loc, dstType, offset, amountToShiftLeft);
535 rewriter.replaceOpWithNewOp<LLVM::AShrOp>(op, dstType, baseShiftedLeft,
536 amountToShiftRight);
537 return success();
538 }
539};
540
541class BitFieldUExtractPattern
542 : public SPIRVToLLVMConversion<spirv::BitFieldUExtractOp> {
543public:
544 using SPIRVToLLVMConversion<spirv::BitFieldUExtractOp>::SPIRVToLLVMConversion;
545
546 LogicalResult
547 matchAndRewrite(spirv::BitFieldUExtractOp op, OpAdaptor adaptor,
548 ConversionPatternRewriter &rewriter) const override {
549 auto srcType = op.getType();
550 auto dstType = getTypeConverter()->convertType(srcType);
551 if (!dstType)
552 return rewriter.notifyMatchFailure(op, "type conversion failed");
553 Location loc = op.getLoc();
554
555 // Process `Offset` and `Count`: broadcast and extend/truncate if needed.
556 Value offset = processCountOrOffset(loc, op.getOffset(), srcType, dstType,
557 *getTypeConverter(), rewriter);
558 Value count = processCountOrOffset(loc, op.getCount(), srcType, dstType,
559 *getTypeConverter(), rewriter);
560
561 // Create a mask with bits set at [0, Count - 1].
562 Value minusOne = createConstantAllBitsSet(loc, srcType, dstType, rewriter);
563 Value maskShiftedByCount =
564 LLVM::ShlOp::create(rewriter, loc, dstType, minusOne, count);
565 Value mask = LLVM::XOrOp::create(rewriter, loc, dstType, maskShiftedByCount,
566 minusOne);
567
568 // Shift `Base` by `Offset` and apply the mask on it.
569 Value shiftedBase =
570 LLVM::LShrOp::create(rewriter, loc, dstType, op.getBase(), offset);
571 rewriter.replaceOpWithNewOp<LLVM::AndOp>(op, dstType, shiftedBase, mask);
572 return success();
573 }
574};
575
576class BranchConversionPattern : public SPIRVToLLVMConversion<spirv::BranchOp> {
577public:
578 using SPIRVToLLVMConversion<spirv::BranchOp>::SPIRVToLLVMConversion;
579
580 LogicalResult
581 matchAndRewrite(spirv::BranchOp branchOp, OpAdaptor adaptor,
582 ConversionPatternRewriter &rewriter) const override {
583 rewriter.replaceOpWithNewOp<LLVM::BrOp>(branchOp, adaptor.getOperands(),
584 branchOp.getTarget());
585 return success();
586 }
587};
588
589class BranchConditionalConversionPattern
590 : public SPIRVToLLVMConversion<spirv::BranchConditionalOp> {
591public:
592 using SPIRVToLLVMConversion<
593 spirv::BranchConditionalOp>::SPIRVToLLVMConversion;
594
595 LogicalResult
596 matchAndRewrite(spirv::BranchConditionalOp op, OpAdaptor adaptor,
597 ConversionPatternRewriter &rewriter) const override {
598 // If branch weights exist, map them to 32-bit integer vector.
599 DenseI32ArrayAttr branchWeights = nullptr;
600 if (auto weights = op.getBranchWeights()) {
601 SmallVector<int32_t> weightValues;
602 for (auto weight : weights->getAsRange<IntegerAttr>())
603 weightValues.push_back(weight.getInt());
604 branchWeights = DenseI32ArrayAttr::get(getContext(), weightValues);
605 }
606
607 rewriter.replaceOpWithNewOp<LLVM::CondBrOp>(
608 op, op.getCondition(), op.getTrueBlockArguments(),
609 op.getFalseBlockArguments(), branchWeights, op.getTrueBlock(),
610 op.getFalseBlock());
611 return success();
612 }
613};
614
615/// Converts `spirv.getCompositeExtract` to `llvm.extractvalue` if the container
616/// type is an aggregate type (struct or array). Otherwise, converts to
617/// `llvm.extractelement` that operates on vectors.
618class CompositeExtractPattern
619 : public SPIRVToLLVMConversion<spirv::CompositeExtractOp> {
620public:
621 using SPIRVToLLVMConversion<spirv::CompositeExtractOp>::SPIRVToLLVMConversion;
622
623 LogicalResult
624 matchAndRewrite(spirv::CompositeExtractOp op, OpAdaptor adaptor,
625 ConversionPatternRewriter &rewriter) const override {
626 auto dstType = this->getTypeConverter()->convertType(op.getType());
627 if (!dstType)
628 return rewriter.notifyMatchFailure(op, "type conversion failed");
629
630 Type containerType = op.getComposite().getType();
631 if (isa<VectorType>(containerType)) {
632 Location loc = op.getLoc();
633 IntegerAttr value = cast<IntegerAttr>(op.getIndices()[0]);
634 Value index = createI32ConstantOf(loc, rewriter, value.getInt());
635 rewriter.replaceOpWithNewOp<LLVM::ExtractElementOp>(
636 op, dstType, adaptor.getComposite(), index);
637 return success();
638 }
639
640 rewriter.replaceOpWithNewOp<LLVM::ExtractValueOp>(
641 op, adaptor.getComposite(),
642 LLVM::convertArrayToIndices(op.getIndices()));
643 return success();
644 }
645};
646
647/// Converts `spirv.getCompositeInsert` to `llvm.insertvalue` if the container
648/// type is an aggregate type (struct or array). Otherwise, converts to
649/// `llvm.insertelement` that operates on vectors.
650class CompositeInsertPattern
651 : public SPIRVToLLVMConversion<spirv::CompositeInsertOp> {
652public:
653 using SPIRVToLLVMConversion<spirv::CompositeInsertOp>::SPIRVToLLVMConversion;
654
655 LogicalResult
656 matchAndRewrite(spirv::CompositeInsertOp op, OpAdaptor adaptor,
657 ConversionPatternRewriter &rewriter) const override {
658 auto dstType = this->getTypeConverter()->convertType(op.getType());
659 if (!dstType)
660 return rewriter.notifyMatchFailure(op, "type conversion failed");
661
662 Type containerType = op.getComposite().getType();
663 if (isa<VectorType>(containerType)) {
664 Location loc = op.getLoc();
665 IntegerAttr value = cast<IntegerAttr>(op.getIndices()[0]);
666 Value index = createI32ConstantOf(loc, rewriter, value.getInt());
667 rewriter.replaceOpWithNewOp<LLVM::InsertElementOp>(
668 op, dstType, adaptor.getComposite(), adaptor.getObject(), index);
669 return success();
670 }
671
672 rewriter.replaceOpWithNewOp<LLVM::InsertValueOp>(
673 op, adaptor.getComposite(), adaptor.getObject(),
674 LLVM::convertArrayToIndices(op.getIndices()));
675 return success();
676 }
677};
678
679/// Converts SPIR-V operations that have straightforward LLVM equivalent
680/// into LLVM dialect operations.
681template <typename SPIRVOp, typename LLVMOp>
682class DirectConversionPattern : public SPIRVToLLVMConversion<SPIRVOp> {
683public:
684 using SPIRVToLLVMConversion<SPIRVOp>::SPIRVToLLVMConversion;
685
686 LogicalResult
687 matchAndRewrite(SPIRVOp op, typename SPIRVOp::Adaptor adaptor,
688 ConversionPatternRewriter &rewriter) const override {
689 auto dstType = this->getTypeConverter()->convertType(op.getType());
690 if (!dstType)
691 return rewriter.notifyMatchFailure(op, "type conversion failed");
692 typename LLVMOp::Properties properties{};
693 SmallVector<NamedAttribute> discardableAttrs;
694 if (failed(collectAttrsForConversion<LLVMOp>(op, properties,
695 discardableAttrs)))
696 return failure();
697 rewriter.template replaceOpWithNewOp<LLVMOp>(
698 op, dstType, adaptor.getOperands(), properties, discardableAttrs);
699 return success();
700 }
701};
702
703/// Converts SPIR-V extended arithmetic ops (`spirv.IAddCarry`,
704/// `spirv.ISubBorrow`) that produce a two-member struct of {low-order bits,
705/// carry/borrow} into the matching LLVM `*.with.overflow` intrinsic. The
706/// intrinsic yields an `i1` carry/borrow, which is zero-extended to the full
707/// component width and re-packed into the SPIR-V result struct.
708template <typename SPIRVOp, typename LLVMOp>
709class ArithmeticWithOverflowPattern : public SPIRVToLLVMConversion<SPIRVOp> {
710public:
711 using SPIRVToLLVMConversion<SPIRVOp>::SPIRVToLLVMConversion;
712
713 LogicalResult
714 matchAndRewrite(SPIRVOp op, typename SPIRVOp::Adaptor adaptor,
715 ConversionPatternRewriter &rewriter) const override {
716 Type dstType = this->getTypeConverter()->convertType(op.getType());
717 if (!dstType)
718 return rewriter.notifyMatchFailure(op, "type conversion failed");
719
720 Location loc = op.getLoc();
721 Type operandType = adaptor.getOperand1().getType();
722 Type overflowType = rewriter.getI1Type();
723 if (auto vecType = dyn_cast<VectorType>(operandType))
724 overflowType = VectorType::get(vecType.getShape(), overflowType);
725
726 Type intrType = LLVM::LLVMStructType::getLiteral(
727 rewriter.getContext(), {operandType, overflowType});
728 Value intrResult = LLVMOp::create(
729 rewriter, loc, intrType, adaptor.getOperand1(), adaptor.getOperand2());
730 Value lowBits = LLVM::ExtractValueOp::create(rewriter, loc, intrResult, 0);
731 Value overflow = LLVM::ExtractValueOp::create(rewriter, loc, intrResult, 1);
732 overflow = LLVM::ZExtOp::create(rewriter, loc, operandType, overflow);
733
734 Value result = LLVM::PoisonOp::create(rewriter, loc, dstType);
735 result = LLVM::InsertValueOp::create(rewriter, loc, result, lowBits,
736 ArrayRef<int64_t>{0});
737 result = LLVM::InsertValueOp::create(rewriter, loc, result, overflow,
738 ArrayRef<int64_t>{1});
739 rewriter.replaceOp(op, result);
740 return success();
741 }
742};
743
744/// Converts `spirv.ExecutionMode` into a global struct constant that holds
745/// execution mode information.
746class ExecutionModePattern
747 : public SPIRVToLLVMConversion<spirv::ExecutionModeOp> {
748public:
749 using SPIRVToLLVMConversion<spirv::ExecutionModeOp>::SPIRVToLLVMConversion;
750
751 LogicalResult
752 matchAndRewrite(spirv::ExecutionModeOp op, OpAdaptor adaptor,
753 ConversionPatternRewriter &rewriter) const override {
754 // First, create the global struct's name that would be associated with
755 // this entry point's execution mode. We set it to be:
756 // __spv__{SPIR-V module name}_{function name}_execution_mode_info_{mode}
757 ModuleOp module = op->getParentOfType<ModuleOp>();
758 spirv::ExecutionModeAttr executionModeAttr = op.getExecutionModeAttr();
759 std::string moduleName;
760 if (module.getName().has_value())
761 moduleName = "_" + module.getName()->str();
762 else
763 moduleName = "";
764 std::string executionModeInfoName = llvm::formatv(
765 "__spv_{0}_{1}_execution_mode_info_{2}", moduleName, op.getFn().str(),
766 static_cast<uint32_t>(executionModeAttr.getValue()));
767
768 MLIRContext *context = rewriter.getContext();
769 OpBuilder::InsertionGuard guard(rewriter);
770 rewriter.setInsertionPointToStart(module.getBody());
771
772 // Create a struct type, corresponding to the C struct below.
773 // struct {
774 // int32_t executionMode;
775 // int32_t values[]; // optional values
776 // };
777 auto llvmI32Type = IntegerType::get(context, 32);
778 SmallVector<Type, 2> fields;
779 fields.push_back(llvmI32Type);
780 ArrayAttr values = op.getValues();
781 if (!values.empty()) {
782 auto arrayType = LLVM::LLVMArrayType::get(llvmI32Type, values.size());
783 fields.push_back(arrayType);
784 }
785 auto structType = LLVM::LLVMStructType::getLiteral(context, fields);
786
787 // Create `llvm.mlir.global` with initializer region containing one block.
788 auto global = LLVM::GlobalOp::create(
789 rewriter, UnknownLoc::get(context), structType, /*isConstant=*/true,
790 LLVM::Linkage::External, executionModeInfoName, Attribute(),
791 /*alignment=*/0);
792 Location loc = global.getLoc();
793 Region &region = global.getInitializerRegion();
794 Block *block = rewriter.createBlock(&region);
795
796 // Initialize the struct and set the execution mode value.
797 rewriter.setInsertionPointToStart(block);
798 Value structValue = LLVM::PoisonOp::create(rewriter, loc, structType);
799 Value executionMode = LLVM::ConstantOp::create(
800 rewriter, loc, llvmI32Type,
801 rewriter.getI32IntegerAttr(
802 static_cast<uint32_t>(executionModeAttr.getValue())));
803 SmallVector<int64_t> position{0};
804 structValue = LLVM::InsertValueOp::create(rewriter, loc, structValue,
805 executionMode, position);
806
807 // Insert extra operands if they exist into execution mode info struct.
808 for (unsigned i = 0, e = values.size(); i < e; ++i) {
809 auto attr = values.getValue()[i];
810 Value entry = LLVM::ConstantOp::create(rewriter, loc, llvmI32Type, attr);
811 structValue = LLVM::InsertValueOp::create(
812 rewriter, loc, structValue, entry, ArrayRef<int64_t>({1, i}));
813 }
814 LLVM::ReturnOp::create(rewriter, loc, ArrayRef<Value>({structValue}));
815 rewriter.eraseOp(op);
816 return success();
817 }
818};
819
820/// Converts `spirv.GlobalVariable` to `llvm.mlir.global`. Note that SPIR-V
821/// global returns a pointer, whereas in LLVM dialect the global holds an actual
822/// value. This difference is handled by `spirv.mlir.addressof` and
823/// `llvm.mlir.addressof`ops that both return a pointer.
824class GlobalVariablePattern
825 : public SPIRVToLLVMConversion<spirv::GlobalVariableOp> {
826public:
827 template <typename... Args>
828 GlobalVariablePattern(spirv::ClientAPI clientAPI, Args &&...args)
829 : SPIRVToLLVMConversion<spirv::GlobalVariableOp>(
830 std::forward<Args>(args)...),
831 clientAPI(clientAPI) {}
832
833 LogicalResult
834 matchAndRewrite(spirv::GlobalVariableOp op, OpAdaptor adaptor,
835 ConversionPatternRewriter &rewriter) const override {
836 // Currently, there is no support of initialization with a constant value in
837 // SPIR-V dialect. Specialization constants are not considered as well.
838 if (op.getInitializer())
839 return failure();
840
841 auto srcType = cast<spirv::PointerType>(op.getType());
842 auto dstType = getTypeConverter()->convertType(srcType.getPointeeType());
843 if (!dstType)
844 return rewriter.notifyMatchFailure(op, "type conversion failed");
845
846 // Limit conversion to the current invocation only or `StorageBuffer`
847 // required by SPIR-V runner.
848 // This is okay because multiple invocations are not supported yet.
849 auto storageClass = srcType.getStorageClass();
850 switch (storageClass) {
851 case spirv::StorageClass::Input:
852 case spirv::StorageClass::Private:
853 case spirv::StorageClass::Output:
854 case spirv::StorageClass::StorageBuffer:
855 case spirv::StorageClass::UniformConstant:
856 break;
857 default:
858 return failure();
859 }
860
861 // LLVM dialect spec: "If the global value is a constant, storing into it is
862 // not allowed.". This corresponds to SPIR-V 'Input' and 'UniformConstant'
863 // storage class that is read-only.
864 bool isConstant = (storageClass == spirv::StorageClass::Input) ||
865 (storageClass == spirv::StorageClass::UniformConstant);
866 // SPIR-V spec: "By default, functions and global variables are private to a
867 // module and cannot be accessed by other modules. However, a module may be
868 // written to export or import functions and global (module scope)
869 // variables.". Therefore, map 'Private' storage class to private linkage,
870 // 'Input' and 'Output' to external linkage.
871 auto linkage = storageClass == spirv::StorageClass::Private
872 ? LLVM::Linkage::Private
873 : LLVM::Linkage::External;
874 StringAttr locationAttrName = op.getLocationAttrName();
875 IntegerAttr locationAttr = op.getLocationAttr();
876 auto newGlobalOp = rewriter.replaceOpWithNewOp<LLVM::GlobalOp>(
877 op, dstType, isConstant, linkage, op.getSymName(), Attribute(),
878 /*alignment=*/0, storageClassToAddressSpace(clientAPI, storageClass));
879
880 // Attach location attribute if applicable
881 if (locationAttr)
882 newGlobalOp->setDiscardableAttr(locationAttrName, locationAttr);
883
884 return success();
885 }
886
887private:
888 spirv::ClientAPI clientAPI;
889};
890
891/// Converts SPIR-V cast ops that do not have straightforward LLVM
892/// equivalent in LLVM dialect.
893template <typename SPIRVOp, typename LLVMExtOp, typename LLVMTruncOp>
894class IndirectCastPattern : public SPIRVToLLVMConversion<SPIRVOp> {
895public:
896 using SPIRVToLLVMConversion<SPIRVOp>::SPIRVToLLVMConversion;
897
898 LogicalResult
899 matchAndRewrite(SPIRVOp op, typename SPIRVOp::Adaptor adaptor,
900 ConversionPatternRewriter &rewriter) const override {
901
902 Type fromType = op.getOperand().getType();
903 Type toType = op.getType();
904
905 auto dstType = this->getTypeConverter()->convertType(toType);
906 if (!dstType)
907 return rewriter.notifyMatchFailure(op, "type conversion failed");
908
909 if (getBitWidth(fromType) < getBitWidth(toType)) {
910 rewriter.template replaceOpWithNewOp<LLVMExtOp>(op, dstType,
911 adaptor.getOperands());
912 return success();
913 }
914 if (getBitWidth(fromType) > getBitWidth(toType)) {
915 rewriter.template replaceOpWithNewOp<LLVMTruncOp>(op, dstType,
916 adaptor.getOperands());
917 return success();
918 }
919 return failure();
920 }
921};
922
923class FunctionCallPattern
924 : public SPIRVToLLVMConversion<spirv::FunctionCallOp> {
925public:
926 using SPIRVToLLVMConversion<spirv::FunctionCallOp>::SPIRVToLLVMConversion;
927
928 LogicalResult
929 matchAndRewrite(spirv::FunctionCallOp callOp, OpAdaptor adaptor,
930 ConversionPatternRewriter &rewriter) const override {
931 LLVM::CallOp::Properties properties{};
932 SmallVector<NamedAttribute> discardableAttrs;
933 if (failed(collectAttrsForConversion<LLVM::CallOp>(callOp, properties,
934 discardableAttrs)))
935 return failure();
936 properties.operandSegmentSizes = {
937 static_cast<int32_t>(adaptor.getOperands().size()), 0};
938 properties.op_bundle_sizes = rewriter.getDenseI32ArrayAttr({});
939
940 if (callOp.getNumResults() == 0) {
941 rewriter.replaceOpWithNewOp<LLVM::CallOp>(callOp, TypeRange(),
942 adaptor.getOperands(),
943 properties, discardableAttrs);
944 return success();
945 }
946
947 // Function returns a single result.
948 auto dstType = getTypeConverter()->convertType(callOp.getType(0));
949 if (!dstType)
950 return rewriter.notifyMatchFailure(callOp, "type conversion failed");
951 rewriter.replaceOpWithNewOp<LLVM::CallOp>(
952 callOp, dstType, adaptor.getOperands(), properties, discardableAttrs);
953 return success();
954 }
955};
956
957/// Converts SPIR-V floating-point comparisons to llvm.fcmp "predicate"
958template <typename SPIRVOp, LLVM::FCmpPredicate predicate>
959class FComparePattern : public SPIRVToLLVMConversion<SPIRVOp> {
960public:
961 using SPIRVToLLVMConversion<SPIRVOp>::SPIRVToLLVMConversion;
962
963 LogicalResult
964 matchAndRewrite(SPIRVOp op, typename SPIRVOp::Adaptor adaptor,
965 ConversionPatternRewriter &rewriter) const override {
966
967 auto dstType = this->getTypeConverter()->convertType(op.getType());
968 if (!dstType)
969 return rewriter.notifyMatchFailure(op, "type conversion failed");
970
971 rewriter.template replaceOpWithNewOp<LLVM::FCmpOp>(
972 op, dstType, predicate, op.getOperand1(), op.getOperand2());
973 return success();
974 }
975};
976
977/// Converts SPIR-V integer comparisons to llvm.icmp "predicate"
978template <typename SPIRVOp, LLVM::ICmpPredicate predicate>
979class IComparePattern : public SPIRVToLLVMConversion<SPIRVOp> {
980public:
981 using SPIRVToLLVMConversion<SPIRVOp>::SPIRVToLLVMConversion;
982
983 LogicalResult
984 matchAndRewrite(SPIRVOp op, typename SPIRVOp::Adaptor adaptor,
985 ConversionPatternRewriter &rewriter) const override {
986
987 auto dstType = this->getTypeConverter()->convertType(op.getType());
988 if (!dstType)
989 return rewriter.notifyMatchFailure(op, "type conversion failed");
990
991 rewriter.template replaceOpWithNewOp<LLVM::ICmpOp>(
992 op, dstType, predicate, op.getOperand1(), op.getOperand2());
993 return success();
994 }
995};
996
997class InverseSqrtPattern
998 : public SPIRVToLLVMConversion<spirv::GLInverseSqrtOp> {
999public:
1000 using SPIRVToLLVMConversion<spirv::GLInverseSqrtOp>::SPIRVToLLVMConversion;
1001
1002 LogicalResult
1003 matchAndRewrite(spirv::GLInverseSqrtOp op, OpAdaptor adaptor,
1004 ConversionPatternRewriter &rewriter) const override {
1005 auto srcType = op.getType();
1006 auto dstType = getTypeConverter()->convertType(srcType);
1007 if (!dstType)
1008 return rewriter.notifyMatchFailure(op, "type conversion failed");
1009
1010 Location loc = op.getLoc();
1011 Value one = createFPConstant(loc, srcType, dstType, rewriter, 1.0);
1012 Value sqrt = LLVM::SqrtOp::create(rewriter, loc, dstType, op.getOperand());
1013 rewriter.replaceOpWithNewOp<LLVM::FDivOp>(op, dstType, one, sqrt);
1014 return success();
1015 }
1016};
1017
1018/// Converts `spirv.VectorTimesScalar` to a broadcast of the scalar followed by
1019/// an `llvm.fmul`.
1020class VectorTimesScalarPattern
1021 : public SPIRVToLLVMConversion<spirv::VectorTimesScalarOp> {
1022public:
1023 using SPIRVToLLVMConversion<
1024 spirv::VectorTimesScalarOp>::SPIRVToLLVMConversion;
1025
1026 LogicalResult
1027 matchAndRewrite(spirv::VectorTimesScalarOp op, OpAdaptor adaptor,
1028 ConversionPatternRewriter &rewriter) const override {
1029 Type srcType = op.getType();
1030 Type dstType = getTypeConverter()->convertType(srcType);
1031 if (!dstType)
1032 return rewriter.notifyMatchFailure(op, "type conversion failed");
1033
1034 unsigned numElements = op.getVector().getType().getNumElements();
1035 Value broadcasted = broadcast(op.getLoc(), adaptor.getScalar(), numElements,
1036 *getTypeConverter(), rewriter);
1037 rewriter.replaceOpWithNewOp<LLVM::FMulOp>(op, dstType, adaptor.getVector(),
1038 broadcasted);
1039 return success();
1040 }
1041};
1042
1043/// Converts `spirv.SNegate` to `0 - x`.
1044class SNegatePattern : public SPIRVToLLVMConversion<spirv::SNegateOp> {
1045public:
1046 using SPIRVToLLVMConversion<spirv::SNegateOp>::SPIRVToLLVMConversion;
1047
1048 LogicalResult
1049 matchAndRewrite(spirv::SNegateOp op, OpAdaptor adaptor,
1050 ConversionPatternRewriter &rewriter) const override {
1051 Type srcType = op.getType();
1052 Type dstType = getTypeConverter()->convertType(srcType);
1053 if (!dstType)
1054 return rewriter.notifyMatchFailure(op, "type conversion failed");
1055
1056 Location loc = op.getLoc();
1057 IntegerAttr zeroAttr = rewriter.getIntegerAttr(
1058 cast<IntegerType>(getElementTypeOrSelf(srcType)), 0);
1059 Value zero =
1060 createIntegerConstant(loc, srcType, dstType, rewriter, zeroAttr);
1061 rewriter.replaceOpWithNewOp<LLVM::SubOp>(op, dstType, zero,
1062 adaptor.getOperand());
1063 return success();
1064 }
1065};
1066
1067/// Converts the GLSL clamp ops (FClamp, SClamp, UClamp) into a nested
1068/// min/max sequence, following the op semantics `min(max(x, minVal), maxVal)`.
1069template <typename SPIRVOp, typename LLVMMinOp, typename LLVMMaxOp>
1070class ClampPattern : public SPIRVToLLVMConversion<SPIRVOp> {
1071public:
1072 using SPIRVToLLVMConversion<SPIRVOp>::SPIRVToLLVMConversion;
1073
1074 LogicalResult
1075 matchAndRewrite(SPIRVOp op, typename SPIRVOp::Adaptor adaptor,
1076 ConversionPatternRewriter &rewriter) const override {
1077 Type dstType = this->getTypeConverter()->convertType(op.getType());
1078 if (!dstType)
1079 return rewriter.notifyMatchFailure(op, "type conversion failed");
1080
1081 Location loc = op.getLoc();
1082 Value max = LLVMMaxOp::create(rewriter, loc, dstType, adaptor.getX(),
1083 adaptor.getY());
1084 rewriter.template replaceOpWithNewOp<LLVMMinOp>(op, dstType, max,
1085 adaptor.getZ());
1086 return success();
1087 }
1088};
1089
1090/// Converts `spirv.FMod` to `x - y * floor(x / y)`. The SPIR-V op requires the
1091/// result to take the sign of the divisor, whereas `llvm.frem` keeps the sign
1092/// of the dividend, so `frem` cannot be used directly.
1093class FModPattern : public SPIRVToLLVMConversion<spirv::FModOp> {
1094public:
1095 using SPIRVToLLVMConversion<spirv::FModOp>::SPIRVToLLVMConversion;
1096
1097 LogicalResult
1098 matchAndRewrite(spirv::FModOp op, OpAdaptor adaptor,
1099 ConversionPatternRewriter &rewriter) const override {
1100 Type dstType = getTypeConverter()->convertType(op.getType());
1101 if (!dstType)
1102 return rewriter.notifyMatchFailure(op, "type conversion failed");
1103
1104 Location loc = op.getLoc();
1105 Value lhs = adaptor.getOperand1();
1106 Value rhs = adaptor.getOperand2();
1107 Value div = LLVM::FDivOp::create(rewriter, loc, dstType, lhs, rhs);
1108 Value floored = LLVM::FFloorOp::create(rewriter, loc, dstType, div);
1109 Value scaled = LLVM::FMulOp::create(rewriter, loc, dstType, rhs, floored);
1110 rewriter.replaceOpWithNewOp<LLVM::FSubOp>(op, dstType, lhs, scaled);
1111 return success();
1112 }
1113};
1114
1115/// Converts `spirv.SMod` to a signed remainder corrected to take the sign of
1116/// the divisor. `llvm.srem` keeps the sign of the dividend, so the result is
1117/// adjusted by adding the divisor when the remainder is non-zero and its sign
1118/// differs from the divisor's.
1119class SModPattern : public SPIRVToLLVMConversion<spirv::SModOp> {
1120public:
1121 using SPIRVToLLVMConversion<spirv::SModOp>::SPIRVToLLVMConversion;
1122
1123 LogicalResult
1124 matchAndRewrite(spirv::SModOp op, OpAdaptor adaptor,
1125 ConversionPatternRewriter &rewriter) const override {
1126 Type srcType = op.getType();
1127 Type dstType = getTypeConverter()->convertType(srcType);
1128 if (!dstType)
1129 return rewriter.notifyMatchFailure(op, "type conversion failed");
1130
1131 Location loc = op.getLoc();
1132 Value lhs = adaptor.getOperand1();
1133 Value rhs = adaptor.getOperand2();
1134 Type i1Type = rewriter.getI1Type();
1135 auto vecSrcType = dyn_cast<VectorType>(srcType);
1136 Type cmpType =
1137 vecSrcType ? VectorType::get(vecSrcType.getShape(), i1Type) : i1Type;
1138
1139 Value rem = LLVM::SRemOp::create(rewriter, loc, dstType, lhs, rhs);
1140 IntegerAttr zeroAttr = rewriter.getIntegerAttr(
1141 cast<IntegerType>(getElementTypeOrSelf(srcType)), 0);
1142 Value zero =
1143 createIntegerConstant(loc, srcType, dstType, rewriter, zeroAttr);
1144
1145 Value remNonZero = LLVM::ICmpOp::create(rewriter, loc, cmpType,
1146 LLVM::ICmpPredicate::ne, rem, zero);
1147 Value remNeg = LLVM::ICmpOp::create(rewriter, loc, cmpType,
1148 LLVM::ICmpPredicate::slt, rem, zero);
1149 Value rhsNeg = LLVM::ICmpOp::create(rewriter, loc, cmpType,
1150 LLVM::ICmpPredicate::slt, rhs, zero);
1151 Value signMismatch =
1152 LLVM::XOrOp::create(rewriter, loc, cmpType, remNeg, rhsNeg);
1153 Value needsAdjust =
1154 LLVM::AndOp::create(rewriter, loc, cmpType, remNonZero, signMismatch);
1155
1156 Value adjusted = LLVM::AddOp::create(rewriter, loc, dstType, rem, rhs);
1157 rewriter.replaceOpWithNewOp<LLVM::SelectOp>(op, dstType, needsAdjust,
1158 adjusted, rem);
1159 return success();
1160 }
1161};
1162
1163/// Converts `spirv.Load` and `spirv.Store` to LLVM dialect.
1164template <typename SPIRVOp>
1165class LoadStorePattern : public SPIRVToLLVMConversion<SPIRVOp> {
1166public:
1167 using SPIRVToLLVMConversion<SPIRVOp>::SPIRVToLLVMConversion;
1168
1169 LogicalResult
1170 matchAndRewrite(SPIRVOp op, typename SPIRVOp::Adaptor adaptor,
1171 ConversionPatternRewriter &rewriter) const override {
1172 if (!op.getMemoryAccess()) {
1173 return replaceWithLoadOrStore(op, adaptor.getOperands(), rewriter,
1174 *this->getTypeConverter(), /*alignment=*/0,
1175 /*isVolatile=*/false,
1176 /*isNonTemporal=*/false);
1177 }
1178 auto memoryAccess = *op.getMemoryAccess();
1179 switch (memoryAccess) {
1180 case spirv::MemoryAccess::Aligned:
1181 case spirv::MemoryAccess::None:
1182 case spirv::MemoryAccess::Nontemporal:
1183 case spirv::MemoryAccess::Volatile: {
1184 unsigned alignment =
1185 memoryAccess == spirv::MemoryAccess::Aligned ? *op.getAlignment() : 0;
1186 bool isNonTemporal = memoryAccess == spirv::MemoryAccess::Nontemporal;
1187 bool isVolatile = memoryAccess == spirv::MemoryAccess::Volatile;
1188 return replaceWithLoadOrStore(op, adaptor.getOperands(), rewriter,
1189 *this->getTypeConverter(), alignment,
1190 isVolatile, isNonTemporal);
1191 }
1192 default:
1193 // There is no support of other memory access attributes.
1194 return failure();
1195 }
1196 }
1197};
1198
1199/// Converts `spirv.Not` and `spirv.LogicalNot` into LLVM dialect.
1200template <typename SPIRVOp>
1201class NotPattern : public SPIRVToLLVMConversion<SPIRVOp> {
1202public:
1203 using SPIRVToLLVMConversion<SPIRVOp>::SPIRVToLLVMConversion;
1204
1205 LogicalResult
1206 matchAndRewrite(SPIRVOp notOp, typename SPIRVOp::Adaptor adaptor,
1207 ConversionPatternRewriter &rewriter) const override {
1208 auto srcType = notOp.getType();
1209 auto dstType = this->getTypeConverter()->convertType(srcType);
1210 if (!dstType)
1211 return rewriter.notifyMatchFailure(notOp, "type conversion failed");
1212
1213 Location loc = notOp.getLoc();
1214 Value mask = createConstantAllBitsSet(loc, srcType, dstType, rewriter);
1215 rewriter.template replaceOpWithNewOp<LLVM::XOrOp>(notOp, dstType,
1216 notOp.getOperand(), mask);
1217 return success();
1218 }
1219};
1220
1221/// A template pattern that erases the given `SPIRVOp`.
1222template <typename SPIRVOp>
1223class ErasePattern : public SPIRVToLLVMConversion<SPIRVOp> {
1224public:
1225 using SPIRVToLLVMConversion<SPIRVOp>::SPIRVToLLVMConversion;
1226
1227 LogicalResult
1228 matchAndRewrite(SPIRVOp op, typename SPIRVOp::Adaptor adaptor,
1229 ConversionPatternRewriter &rewriter) const override {
1230 rewriter.eraseOp(op);
1231 return success();
1232 }
1233};
1234
1235class ReturnPattern : public SPIRVToLLVMConversion<spirv::ReturnOp> {
1236public:
1237 using SPIRVToLLVMConversion<spirv::ReturnOp>::SPIRVToLLVMConversion;
1238
1239 LogicalResult
1240 matchAndRewrite(spirv::ReturnOp returnOp, OpAdaptor adaptor,
1241 ConversionPatternRewriter &rewriter) const override {
1242 rewriter.replaceOpWithNewOp<LLVM::ReturnOp>(returnOp, ArrayRef<Type>(),
1243 ArrayRef<Value>());
1244 return success();
1245 }
1246};
1247
1248class ReturnValuePattern : public SPIRVToLLVMConversion<spirv::ReturnValueOp> {
1249public:
1250 using SPIRVToLLVMConversion<spirv::ReturnValueOp>::SPIRVToLLVMConversion;
1251
1252 LogicalResult
1253 matchAndRewrite(spirv::ReturnValueOp returnValueOp, OpAdaptor adaptor,
1254 ConversionPatternRewriter &rewriter) const override {
1255 rewriter.replaceOpWithNewOp<LLVM::ReturnOp>(returnValueOp, ArrayRef<Type>(),
1256 adaptor.getOperands());
1257 return success();
1258 }
1259};
1260
1261class UnreachablePattern : public SPIRVToLLVMConversion<spirv::UnreachableOp> {
1262public:
1263 using SPIRVToLLVMConversion<spirv::UnreachableOp>::SPIRVToLLVMConversion;
1264
1265 LogicalResult
1266 matchAndRewrite(spirv::UnreachableOp unreachableOp, OpAdaptor adaptor,
1267 ConversionPatternRewriter &rewriter) const override {
1268 rewriter.replaceOpWithNewOp<LLVM::UnreachableOp>(unreachableOp);
1269 return success();
1270 }
1271};
1272
1273static LLVM::LLVMFuncOp lookupOrCreateSPIRVFn(Operation *symbolTable,
1274 StringRef name,
1275 ArrayRef<Type> paramTypes,
1276 Type resultType,
1277 bool convergent = true) {
1278 auto func = dyn_cast_or_null<LLVM::LLVMFuncOp>(
1279 SymbolTable::lookupSymbolIn(symbolTable, name));
1280 if (func)
1281 return func;
1282
1283 OpBuilder b(symbolTable->getRegion(0));
1284 func = LLVM::LLVMFuncOp::create(
1285 b, symbolTable->getLoc(), name,
1286 LLVM::LLVMFunctionType::get(resultType, paramTypes));
1287 func.setCConv(LLVM::cconv::CConv::SPIR_FUNC);
1288 func.setConvergent(convergent);
1289 func.setNoUnwind(true);
1290 func.setWillReturn(true);
1291 return func;
1292}
1293
1294static LLVM::CallOp createSPIRVBuiltinCall(Location loc, OpBuilder &builder,
1295 LLVM::LLVMFuncOp func,
1296 ValueRange args) {
1297 auto call = LLVM::CallOp::create(builder, loc, func, args);
1298 call.setCConv(func.getCConv());
1299 call.setConvergentAttr(func.getConvergentAttr());
1300 call.setNoUnwindAttr(func.getNoUnwindAttr());
1301 call.setWillReturnAttr(func.getWillReturnAttr());
1302 return call;
1303}
1304
1305template <typename BarrierOpTy>
1306class ControlBarrierPattern : public SPIRVToLLVMConversion<BarrierOpTy> {
1307public:
1308 using OpAdaptor = typename SPIRVToLLVMConversion<BarrierOpTy>::OpAdaptor;
1309
1310 using SPIRVToLLVMConversion<BarrierOpTy>::SPIRVToLLVMConversion;
1311
1312 static constexpr StringRef getFuncName();
1313
1314 LogicalResult
1315 matchAndRewrite(BarrierOpTy controlBarrierOp, OpAdaptor adaptor,
1316 ConversionPatternRewriter &rewriter) const override {
1317 constexpr StringRef funcName = getFuncName();
1318 Operation *symbolTable =
1319 controlBarrierOp->template getParentWithTrait<OpTrait::SymbolTable>();
1320
1321 Type i32 = rewriter.getI32Type();
1322
1323 Type voidTy = rewriter.getType<LLVM::LLVMVoidType>();
1324 LLVM::LLVMFuncOp func =
1325 lookupOrCreateSPIRVFn(symbolTable, funcName, {i32, i32, i32}, voidTy);
1326
1327 Location loc = controlBarrierOp->getLoc();
1328 Value execution = LLVM::ConstantOp::create(
1329 rewriter, loc, i32, static_cast<int32_t>(adaptor.getExecutionScope()));
1330 Value memory = LLVM::ConstantOp::create(
1331 rewriter, loc, i32, static_cast<int32_t>(adaptor.getMemoryScope()));
1332 Value semantics = LLVM::ConstantOp::create(
1333 rewriter, loc, i32, static_cast<int32_t>(adaptor.getMemorySemantics()));
1334
1335 auto call = createSPIRVBuiltinCall(loc, rewriter, func,
1336 {execution, memory, semantics});
1337
1338 rewriter.replaceOp(controlBarrierOp, call);
1339 return success();
1340 }
1341};
1342
1343namespace {
1344
1345StringRef getTypeMangling(Type type, bool isSigned) {
1347 .Case([](Float16Type) { return "Dh"; })
1348 .Case([](Float32Type) { return "f"; })
1349 .Case([](Float64Type) { return "d"; })
1350 .Case([isSigned](IntegerType intTy) {
1351 switch (intTy.getWidth()) {
1352 case 1:
1353 return "b";
1354 case 8:
1355 return (isSigned) ? "a" : "c";
1356 case 16:
1357 return (isSigned) ? "s" : "t";
1358 case 32:
1359 return (isSigned) ? "i" : "j";
1360 case 64:
1361 return (isSigned) ? "l" : "m";
1362 default:
1363 llvm_unreachable("Unsupported integer width");
1364 }
1365 })
1366 .DefaultUnreachable("No mangling defined");
1367}
1368
1369template <typename ReduceOp>
1370constexpr StringLiteral getGroupFuncName();
1371
1372template <>
1373constexpr StringLiteral getGroupFuncName<spirv::GroupIAddOp>() {
1374 return "_Z17__spirv_GroupIAddii";
1375}
1376template <>
1377constexpr StringLiteral getGroupFuncName<spirv::GroupFAddOp>() {
1378 return "_Z17__spirv_GroupFAddii";
1379}
1380template <>
1381constexpr StringLiteral getGroupFuncName<spirv::GroupSMinOp>() {
1382 return "_Z17__spirv_GroupSMinii";
1383}
1384template <>
1385constexpr StringLiteral getGroupFuncName<spirv::GroupUMinOp>() {
1386 return "_Z17__spirv_GroupUMinii";
1387}
1388template <>
1389constexpr StringLiteral getGroupFuncName<spirv::GroupFMinOp>() {
1390 return "_Z17__spirv_GroupFMinii";
1391}
1392template <>
1393constexpr StringLiteral getGroupFuncName<spirv::GroupSMaxOp>() {
1394 return "_Z17__spirv_GroupSMaxii";
1395}
1396template <>
1397constexpr StringLiteral getGroupFuncName<spirv::GroupUMaxOp>() {
1398 return "_Z17__spirv_GroupUMaxii";
1399}
1400template <>
1401constexpr StringLiteral getGroupFuncName<spirv::GroupFMaxOp>() {
1402 return "_Z17__spirv_GroupFMaxii";
1403}
1404template <>
1405constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformIAddOp>() {
1406 return "_Z27__spirv_GroupNonUniformIAddii";
1407}
1408template <>
1409constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformFAddOp>() {
1410 return "_Z27__spirv_GroupNonUniformFAddii";
1411}
1412template <>
1413constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformIMulOp>() {
1414 return "_Z27__spirv_GroupNonUniformIMulii";
1415}
1416template <>
1417constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformFMulOp>() {
1418 return "_Z27__spirv_GroupNonUniformFMulii";
1419}
1420template <>
1421constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformSMinOp>() {
1422 return "_Z27__spirv_GroupNonUniformSMinii";
1423}
1424template <>
1425constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformUMinOp>() {
1426 return "_Z27__spirv_GroupNonUniformUMinii";
1427}
1428template <>
1429constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformFMinOp>() {
1430 return "_Z27__spirv_GroupNonUniformFMinii";
1431}
1432template <>
1433constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformSMaxOp>() {
1434 return "_Z27__spirv_GroupNonUniformSMaxii";
1435}
1436template <>
1437constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformUMaxOp>() {
1438 return "_Z27__spirv_GroupNonUniformUMaxii";
1439}
1440template <>
1441constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformFMaxOp>() {
1442 return "_Z27__spirv_GroupNonUniformFMaxii";
1443}
1444template <>
1445constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformBitwiseAndOp>() {
1446 return "_Z33__spirv_GroupNonUniformBitwiseAndii";
1447}
1448template <>
1449constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformBitwiseOrOp>() {
1450 return "_Z32__spirv_GroupNonUniformBitwiseOrii";
1451}
1452template <>
1453constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformBitwiseXorOp>() {
1454 return "_Z33__spirv_GroupNonUniformBitwiseXorii";
1455}
1456template <>
1457constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformLogicalAndOp>() {
1458 return "_Z33__spirv_GroupNonUniformLogicalAndii";
1459}
1460template <>
1461constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformLogicalOrOp>() {
1462 return "_Z32__spirv_GroupNonUniformLogicalOrii";
1463}
1464template <>
1465constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformLogicalXorOp>() {
1466 return "_Z33__spirv_GroupNonUniformLogicalXorii";
1467}
1468} // namespace
1469
1470template <typename ReduceOp, bool Signed = false, bool NonUniform = false>
1471class GroupReducePattern : public SPIRVToLLVMConversion<ReduceOp> {
1472public:
1473 using SPIRVToLLVMConversion<ReduceOp>::SPIRVToLLVMConversion;
1474
1475 LogicalResult
1476 matchAndRewrite(ReduceOp op, typename ReduceOp::Adaptor adaptor,
1477 ConversionPatternRewriter &rewriter) const override {
1478
1479 Type retTy = op.getResult().getType();
1480 if (!retTy.isIntOrFloat()) {
1481 return failure();
1482 }
1483 SmallString<36> funcName = getGroupFuncName<ReduceOp>();
1484 funcName += getTypeMangling(retTy, false);
1485
1486 Type i32Ty = rewriter.getI32Type();
1487 SmallVector<Type> paramTypes{i32Ty, i32Ty, retTy};
1488 if constexpr (NonUniform) {
1489 if (adaptor.getClusterSize()) {
1490 funcName += "j";
1491 paramTypes.push_back(i32Ty);
1492 }
1493 }
1494
1495 Operation *symbolTable =
1496 op->template getParentWithTrait<OpTrait::SymbolTable>();
1497
1498 LLVM::LLVMFuncOp func =
1499 lookupOrCreateSPIRVFn(symbolTable, funcName, paramTypes, retTy);
1500
1501 Location loc = op.getLoc();
1502 Value scope = LLVM::ConstantOp::create(
1503 rewriter, loc, i32Ty,
1504 static_cast<int32_t>(adaptor.getExecutionScope()));
1505 Value groupOp = LLVM::ConstantOp::create(
1506 rewriter, loc, i32Ty,
1507 static_cast<int32_t>(adaptor.getGroupOperation()));
1508 SmallVector<Value> operands{scope, groupOp};
1509 operands.append(adaptor.getOperands().begin(), adaptor.getOperands().end());
1510
1511 auto call = createSPIRVBuiltinCall(loc, rewriter, func, operands);
1512 rewriter.replaceOp(op, call);
1513 return success();
1514 }
1515};
1516
1517template <>
1518constexpr StringRef
1519ControlBarrierPattern<spirv::ControlBarrierOp>::getFuncName() {
1520 return "_Z22__spirv_ControlBarrieriii";
1521}
1522
1523template <>
1524constexpr StringRef
1525ControlBarrierPattern<spirv::INTELControlBarrierArriveOp>::getFuncName() {
1526 return "_Z33__spirv_ControlBarrierArriveINTELiii";
1527}
1528
1529template <>
1530constexpr StringRef
1531ControlBarrierPattern<spirv::INTELControlBarrierWaitOp>::getFuncName() {
1532 return "_Z31__spirv_ControlBarrierWaitINTELiii";
1533}
1534
1535/// Converts `spirv.mlir.loop` to LLVM dialect. All blocks within selection
1536/// should be reachable for conversion to succeed. The structure of the loop in
1537/// LLVM dialect will be the following:
1538///
1539/// +------------------------------------+
1540/// | <code before spirv.mlir.loop> |
1541/// | llvm.br ^header |
1542/// +------------------------------------+
1543/// |
1544/// +----------------+ |
1545/// | | |
1546/// | V V
1547/// | +------------------------------------+
1548/// | | ^header: |
1549/// | | <header code> |
1550/// | | llvm.cond_br %cond, ^body, ^exit |
1551/// | +------------------------------------+
1552/// | |
1553/// | |----------------------+
1554/// | | |
1555/// | V |
1556/// | +------------------------------------+ |
1557/// | | ^body: | |
1558/// | | <body code> | |
1559/// | | llvm.br ^continue | |
1560/// | +------------------------------------+ |
1561/// | | |
1562/// | V |
1563/// | +------------------------------------+ |
1564/// | | ^continue: | |
1565/// | | <continue code> | |
1566/// | | llvm.br ^header | |
1567/// | +------------------------------------+ |
1568/// | | |
1569/// +---------------+ +----------------------+
1570/// |
1571/// V
1572/// +------------------------------------+
1573/// | ^exit: |
1574/// | llvm.br ^remaining |
1575/// +------------------------------------+
1576/// |
1577/// V
1578/// +------------------------------------+
1579/// | ^remaining: |
1580/// | <code after spirv.mlir.loop> |
1581/// +------------------------------------+
1582///
1583class LoopPattern : public SPIRVToLLVMConversion<spirv::LoopOp> {
1584public:
1585 using SPIRVToLLVMConversion<spirv::LoopOp>::SPIRVToLLVMConversion;
1586
1587 LogicalResult
1588 matchAndRewrite(spirv::LoopOp loopOp, OpAdaptor adaptor,
1589 ConversionPatternRewriter &rewriter) const override {
1590 // There is no support of loop control at the moment.
1591 if (loopOp.getLoopControl() != spirv::LoopControl::None)
1592 return failure();
1593
1594 // `spirv.mlir.loop` with empty region is redundant and should be erased.
1595 if (loopOp.getBody().empty()) {
1596 rewriter.eraseOp(loopOp);
1597 return success();
1598 }
1599
1600 Location loc = loopOp.getLoc();
1601
1602 // Split the current block after `spirv.mlir.loop`. The remaining ops will
1603 // be used in `endBlock`.
1604 Block *currentBlock = rewriter.getBlock();
1605 auto position = Block::iterator(loopOp);
1606 Block *endBlock = rewriter.splitBlock(currentBlock, position);
1607
1608 // Remove entry block and create a branch in the current block going to the
1609 // header block.
1610 Block *entryBlock = loopOp.getEntryBlock();
1611 assert(entryBlock->getOperations().size() == 1);
1612 auto brOp = dyn_cast<spirv::BranchOp>(entryBlock->getOperations().front());
1613 if (!brOp)
1614 return failure();
1615 Block *headerBlock = loopOp.getHeaderBlock();
1616 rewriter.setInsertionPointToEnd(currentBlock);
1617 LLVM::BrOp::create(rewriter, loc, brOp.getBlockArguments(), headerBlock);
1618 rewriter.eraseBlock(entryBlock);
1619
1620 // Branch from merge block to end block.
1621 Block *mergeBlock = loopOp.getMergeBlock();
1622 Operation *terminator = mergeBlock->getTerminator();
1623 ValueRange terminatorOperands = terminator->getOperands();
1624 rewriter.setInsertionPointToEnd(mergeBlock);
1625 LLVM::BrOp::create(rewriter, loc, terminatorOperands, endBlock);
1626
1627 rewriter.inlineRegionBefore(loopOp.getBody(), endBlock);
1628 rewriter.replaceOp(loopOp, endBlock->getArguments());
1629 return success();
1630 }
1631};
1632
1633/// Converts `spirv.mlir.selection` with `spirv.BranchConditional` in its header
1634/// block. All blocks within selection should be reachable for conversion to
1635/// succeed.
1636class SelectionPattern : public SPIRVToLLVMConversion<spirv::SelectionOp> {
1637public:
1638 using SPIRVToLLVMConversion<spirv::SelectionOp>::SPIRVToLLVMConversion;
1639
1640 LogicalResult
1641 matchAndRewrite(spirv::SelectionOp op, OpAdaptor adaptor,
1642 ConversionPatternRewriter &rewriter) const override {
1643 // There is no support for `Flatten` or `DontFlatten` selection control at
1644 // the moment. This are just compiler hints and can be performed during the
1645 // optimization passes.
1646 if (op.getSelectionControl() != spirv::SelectionControl::None)
1647 return failure();
1648
1649 // `spirv.mlir.selection` should have at least two blocks: one selection
1650 // header block and one merge block. If no blocks are present, or control
1651 // flow branches straight to merge block (two blocks are present), the op is
1652 // redundant and it is erased.
1653 if (op.getBody().getBlocks().size() <= 2) {
1654 rewriter.eraseOp(op);
1655 return success();
1656 }
1657
1658 Location loc = op.getLoc();
1659
1660 // Split the current block after `spirv.mlir.selection`. The remaining ops
1661 // will be used in `continueBlock`.
1662 auto *currentBlock = rewriter.getInsertionBlock();
1663 rewriter.setInsertionPointAfter(op);
1664 auto position = rewriter.getInsertionPoint();
1665 auto *continueBlock = rewriter.splitBlock(currentBlock, position);
1666
1667 // Add arguments to the continue block for selections that yield values.
1668 for (auto ty : op.getResultTypes()) {
1669 Type dstTy = getTypeConverter()->convertType(ty);
1670 if (!dstTy)
1671 return rewriter.notifyMatchFailure(op, "failed to convert type");
1672 continueBlock->addArgument(dstTy, loc);
1673 }
1674
1675 // Extract conditional branch information from the header block. By SPIR-V
1676 // dialect spec, it should contain `spirv.BranchConditional` or
1677 // `spirv.Switch` op. Note that `spirv.Switch op` is not supported at the
1678 // moment in the SPIR-V dialect. Remove this block when finished.
1679 auto *headerBlock = op.getHeaderBlock();
1680 assert(headerBlock->getOperations().size() == 1);
1681 auto condBrOp = dyn_cast<spirv::BranchConditionalOp>(
1682 headerBlock->getOperations().front());
1683 if (!condBrOp)
1684 return failure();
1685
1686 // Branch from merge block to continue block.
1687 auto *mergeBlock = op.getMergeBlock();
1688 Operation *terminator = mergeBlock->getTerminator();
1689 ValueRange terminatorOperands = terminator->getOperands();
1690 rewriter.setInsertionPointToEnd(mergeBlock);
1691 LLVM::BrOp::create(rewriter, loc, terminatorOperands, continueBlock);
1692
1693 // Link current block to `true` and `false` blocks within the selection.
1694 Block *trueBlock = condBrOp.getTrueBlock();
1695 Block *falseBlock = condBrOp.getFalseBlock();
1696 rewriter.setInsertionPointToEnd(currentBlock);
1697 LLVM::CondBrOp::create(rewriter, loc, condBrOp.getCondition(), trueBlock,
1698 condBrOp.getTrueTargetOperands(), falseBlock,
1699 condBrOp.getFalseTargetOperands());
1700
1701 rewriter.eraseBlock(headerBlock);
1702 rewriter.inlineRegionBefore(op.getBody(), continueBlock);
1703 rewriter.replaceOp(op, continueBlock->getArguments());
1704 return success();
1705 }
1706};
1707
1708/// Converts SPIR-V shift ops to LLVM shift ops. Since LLVM dialect
1709/// puts a restriction on `Shift` and `Base` to have the same bit width,
1710/// `Shift` is zero or sign extended to match this specification. Cases when
1711/// `Shift` bit width > `Base` bit width are considered to be illegal.
1712template <typename SPIRVOp, typename LLVMOp>
1713class ShiftPattern : public SPIRVToLLVMConversion<SPIRVOp> {
1714public:
1715 using SPIRVToLLVMConversion<SPIRVOp>::SPIRVToLLVMConversion;
1716
1717 LogicalResult
1718 matchAndRewrite(SPIRVOp op, typename SPIRVOp::Adaptor adaptor,
1719 ConversionPatternRewriter &rewriter) const override {
1720
1721 auto dstType = this->getTypeConverter()->convertType(op.getType());
1722 if (!dstType)
1723 return rewriter.notifyMatchFailure(op, "type conversion failed");
1724
1725 Type op1Type = op.getOperand1().getType();
1726 Type op2Type = op.getOperand2().getType();
1727
1728 if (op1Type == op2Type) {
1729 rewriter.template replaceOpWithNewOp<LLVMOp>(op, dstType,
1730 adaptor.getOperands());
1731 return success();
1732 }
1733
1734 std::optional<uint64_t> dstTypeWidth =
1736 std::optional<uint64_t> op2TypeWidth =
1738
1739 if (!dstTypeWidth || !op2TypeWidth)
1740 return failure();
1741
1742 Location loc = op.getLoc();
1743 Value extended;
1744 if (op2TypeWidth < dstTypeWidth) {
1745 if (isUnsignedIntegerOrVector(op2Type)) {
1746 extended =
1747 LLVM::ZExtOp::create(rewriter, loc, dstType, adaptor.getOperand2());
1748 } else {
1749 extended =
1750 LLVM::SExtOp::create(rewriter, loc, dstType, adaptor.getOperand2());
1751 }
1752 } else if (op2TypeWidth == dstTypeWidth) {
1753 extended = adaptor.getOperand2();
1754 } else {
1755 return failure();
1756 }
1757
1758 Value result =
1759 LLVMOp::create(rewriter, loc, dstType, adaptor.getOperand1(), extended);
1760 rewriter.replaceOp(op, result);
1761 return success();
1762 }
1763};
1764
1765// `llvm.intr.abs` requires an `is_int_min_poison` immarg that `spirv.GL.SAbs`
1766// does not carry; default to `false` to preserve SPIR-V's well-defined
1767// behavior on INT_MIN.
1768class SAbsPattern : public SPIRVToLLVMConversion<spirv::GLSAbsOp> {
1769public:
1770 using SPIRVToLLVMConversion<spirv::GLSAbsOp>::SPIRVToLLVMConversion;
1771
1772 LogicalResult
1773 matchAndRewrite(spirv::GLSAbsOp op, OpAdaptor adaptor,
1774 ConversionPatternRewriter &rewriter) const override {
1775 Type dstType = getTypeConverter()->convertType(op.getType());
1776 if (!dstType)
1777 return rewriter.notifyMatchFailure(op, "type conversion failed");
1778
1779 rewriter.replaceOpWithNewOp<LLVM::AbsOp>(op, dstType, adaptor.getOperand(),
1780 /*is_int_min_poison=*/false);
1781 return success();
1782 }
1783};
1784
1785/// Converts `spirv.GL.Fract` to `x - floor(x)`.
1786class FractPattern : public SPIRVToLLVMConversion<spirv::GLFractOp> {
1787public:
1788 using SPIRVToLLVMConversion<spirv::GLFractOp>::SPIRVToLLVMConversion;
1789
1790 LogicalResult
1791 matchAndRewrite(spirv::GLFractOp op, OpAdaptor adaptor,
1792 ConversionPatternRewriter &rewriter) const override {
1793 Type dstType = getTypeConverter()->convertType(op.getType());
1794 if (!dstType)
1795 return rewriter.notifyMatchFailure(op, "type conversion failed");
1796
1797 Location loc = op.getLoc();
1798 Value operand = adaptor.getOperand();
1799 Value floored = LLVM::FFloorOp::create(rewriter, loc, dstType, operand);
1800 rewriter.replaceOpWithNewOp<LLVM::FSubOp>(op, dstType, operand, floored);
1801 return success();
1802 }
1803};
1804
1805/// Converts `spirv.GL.FMix` to `x * (1 - a) + y * a` as specified by
1806/// GL.std.450.
1807class GLFMixPattern : public SPIRVToLLVMConversion<spirv::GLFMixOp> {
1808public:
1809 using SPIRVToLLVMConversion<spirv::GLFMixOp>::SPIRVToLLVMConversion;
1810
1811 LogicalResult
1812 matchAndRewrite(spirv::GLFMixOp op, OpAdaptor adaptor,
1813 ConversionPatternRewriter &rewriter) const override {
1814 Type dstType = getTypeConverter()->convertType(op.getType());
1815 if (!dstType)
1816 return rewriter.notifyMatchFailure(op, "type conversion failed");
1817
1818 Location loc = op.getLoc();
1819 Value x = adaptor.getX();
1820 Value y = adaptor.getY();
1821 Value a = adaptor.getA();
1822 Value one = createFPConstant(loc, op.getType(), dstType, rewriter, 1.0);
1823 Value oneMinusA = LLVM::FSubOp::create(rewriter, loc, dstType, one, a);
1824 Value lhs = LLVM::FMulOp::create(rewriter, loc, dstType, x, oneMinusA);
1825 Value rhs = LLVM::FMulOp::create(rewriter, loc, dstType, y, a);
1826 rewriter.replaceOpWithNewOp<LLVM::FAddOp>(op, dstType, lhs, rhs);
1827 return success();
1828 }
1829};
1830
1831/// Converts `spirv.CL.mix` to `fma(a, y - x, x)`. The OpenCL spec defines
1832/// mix as `x + (y - x) * a` and explicitly permits FMA contractions.
1833class CLMixPattern : public SPIRVToLLVMConversion<spirv::CLMixOp> {
1834public:
1835 using SPIRVToLLVMConversion<spirv::CLMixOp>::SPIRVToLLVMConversion;
1836
1837 LogicalResult
1838 matchAndRewrite(spirv::CLMixOp op, OpAdaptor adaptor,
1839 ConversionPatternRewriter &rewriter) const override {
1840 Type dstType = getTypeConverter()->convertType(op.getType());
1841 if (!dstType)
1842 return rewriter.notifyMatchFailure(op, "type conversion failed");
1843
1844 Location loc = op.getLoc();
1845 Value x = adaptor.getX();
1846 Value y = adaptor.getY();
1847 Value a = adaptor.getZ();
1848 Value diff = LLVM::FSubOp::create(rewriter, loc, dstType, y, x);
1849 rewriter.replaceOpWithNewOp<LLVM::FMAOp>(op, dstType, a, diff, x);
1850 return success();
1851 }
1852};
1853
1854// Converts spirv.GL.Radians (scale = pi/180) and spirv.GL.Degrees
1855// (scale = 180/pi) by multiplying the operand by a compile-time constant.
1856template <typename SPIRVOp>
1857class ScalePattern : public SPIRVToLLVMConversion<SPIRVOp> {
1858public:
1859 template <typename... Args>
1860 ScalePattern(double scale, Args &&...args)
1861 : SPIRVToLLVMConversion<SPIRVOp>(std::forward<Args>(args)...),
1862 scale(scale) {}
1863
1864 LogicalResult
1865 matchAndRewrite(SPIRVOp op, typename SPIRVOp::Adaptor adaptor,
1866 ConversionPatternRewriter &rewriter) const override {
1867 Type srcType = op.getType();
1868 Type dstType = this->getTypeConverter()->convertType(srcType);
1869 if (!dstType)
1870 return rewriter.notifyMatchFailure(op, "type conversion failed");
1871
1872 Location loc = op.getLoc();
1873 Value factor = createFPConstant(loc, srcType, dstType, rewriter, scale);
1874 rewriter.replaceOpWithNewOp<LLVM::FMulOp>(op, dstType, adaptor.getOperand(),
1875 factor);
1876 return success();
1877 }
1878
1879private:
1880 double scale;
1881};
1882
1883/// Converts `spirv.GL.FSign`/`spirv.GL.SSign` to a sign(x) sequence that maps
1884/// the operand to -1/0/1 using two comparisons and two selects. The `isFloat`
1885/// flag selects between floating-point and integer comparisons/constants.
1886template <typename SPIRVOp, bool isFloat>
1887class SignPattern : public SPIRVToLLVMConversion<SPIRVOp> {
1888public:
1889 using SPIRVToLLVMConversion<SPIRVOp>::SPIRVToLLVMConversion;
1890
1891 LogicalResult
1892 matchAndRewrite(SPIRVOp op, typename SPIRVOp::Adaptor adaptor,
1893 ConversionPatternRewriter &rewriter) const override {
1894 Type srcType = op.getType();
1895 Type dstType = this->getTypeConverter()->convertType(srcType);
1896 if (!dstType)
1897 return rewriter.notifyMatchFailure(op, "type conversion failed");
1898
1899 Location loc = op.getLoc();
1900 Value operand = adaptor.getOperand();
1901 auto vecSrcType = dyn_cast<VectorType>(srcType);
1902 Type i1Type = rewriter.getI1Type();
1903 Type cmpType =
1904 vecSrcType ? VectorType::get(vecSrcType.getShape(), i1Type) : i1Type;
1905
1906 Value zero, one, minusOne, gt, lt;
1907 if constexpr (isFloat) {
1908 zero = createFPConstant(loc, srcType, dstType, rewriter, 0.0);
1909 one = createFPConstant(loc, srcType, dstType, rewriter, 1.0);
1910 minusOne = createFPConstant(loc, srcType, dstType, rewriter, -1.0);
1911 gt = LLVM::FCmpOp::create(rewriter, loc, cmpType,
1912 LLVM::FCmpPredicate::ogt, operand, zero);
1913 lt = LLVM::FCmpOp::create(rewriter, loc, cmpType,
1914 LLVM::FCmpPredicate::olt, operand, zero);
1915 } else {
1916 auto intElemType = cast<IntegerType>(getElementTypeOrSelf(srcType));
1917 zero = createIntegerConstant(loc, srcType, dstType, rewriter,
1918 rewriter.getIntegerAttr(intElemType, 0));
1919 one = createIntegerConstant(loc, srcType, dstType, rewriter,
1920 rewriter.getIntegerAttr(intElemType, 1));
1921 minusOne = createConstantAllBitsSet(loc, srcType, dstType, rewriter);
1922 gt = LLVM::ICmpOp::create(rewriter, loc, cmpType,
1923 LLVM::ICmpPredicate::sgt, operand, zero);
1924 lt = LLVM::ICmpOp::create(rewriter, loc, cmpType,
1925 LLVM::ICmpPredicate::slt, operand, zero);
1926 }
1927
1928 Value negOrZero =
1929 LLVM::SelectOp::create(rewriter, loc, dstType, lt, minusOne, zero);
1930 rewriter.replaceOpWithNewOp<LLVM::SelectOp>(op, dstType, gt, one,
1931 negOrZero);
1932 return success();
1933 }
1934};
1935
1936class VariablePattern : public SPIRVToLLVMConversion<spirv::VariableOp> {
1937public:
1938 using SPIRVToLLVMConversion<spirv::VariableOp>::SPIRVToLLVMConversion;
1939
1940 LogicalResult
1941 matchAndRewrite(spirv::VariableOp varOp, OpAdaptor adaptor,
1942 ConversionPatternRewriter &rewriter) const override {
1943 auto srcType = varOp.getType();
1944 // Initialization is supported for scalars and vectors only.
1945 auto pointerTo = cast<spirv::PointerType>(srcType).getPointeeType();
1946 auto init = varOp.getInitializer();
1947 if (init && !pointerTo.isIntOrFloat() && !isa<VectorType>(pointerTo))
1948 return failure();
1949
1950 auto dstType = getTypeConverter()->convertType(srcType);
1951 if (!dstType)
1952 return rewriter.notifyMatchFailure(varOp, "type conversion failed");
1953
1954 Location loc = varOp.getLoc();
1955 Value size = createI32ConstantOf(loc, rewriter, 1);
1956 if (!init) {
1957 auto elementType = getTypeConverter()->convertType(pointerTo);
1958 if (!elementType)
1959 return rewriter.notifyMatchFailure(varOp, "type conversion failed");
1960 rewriter.replaceOpWithNewOp<LLVM::AllocaOp>(varOp, dstType, elementType,
1961 size);
1962 return success();
1963 }
1964 auto elementType = getTypeConverter()->convertType(pointerTo);
1965 if (!elementType)
1966 return rewriter.notifyMatchFailure(varOp, "type conversion failed");
1967 Value allocated =
1968 LLVM::AllocaOp::create(rewriter, loc, dstType, elementType, size);
1969 LLVM::StoreOp::create(rewriter, loc, adaptor.getInitializer(), allocated);
1970 rewriter.replaceOp(varOp, allocated);
1971 return success();
1972 }
1973};
1974
1975//===----------------------------------------------------------------------===//
1976// BitcastOp conversion
1977//===----------------------------------------------------------------------===//
1978
1979class BitcastConversionPattern
1980 : public SPIRVToLLVMConversion<spirv::BitcastOp> {
1981public:
1982 using SPIRVToLLVMConversion<spirv::BitcastOp>::SPIRVToLLVMConversion;
1983
1984 LogicalResult
1985 matchAndRewrite(spirv::BitcastOp bitcastOp, OpAdaptor adaptor,
1986 ConversionPatternRewriter &rewriter) const override {
1987 auto dstType = getTypeConverter()->convertType(bitcastOp.getType());
1988 if (!dstType)
1989 return rewriter.notifyMatchFailure(bitcastOp, "type conversion failed");
1990
1991 // LLVM's opaque pointers do not require bitcasts.
1992 if (isa<LLVM::LLVMPointerType>(dstType)) {
1993 rewriter.replaceOp(bitcastOp, adaptor.getOperand());
1994 return success();
1995 }
1996
1997 LLVM::BitcastOp::Properties properties{};
1998 SmallVector<NamedAttribute> discardableAttrs;
1999 if (failed(collectAttrsForConversion<LLVM::BitcastOp>(bitcastOp, properties,
2000 discardableAttrs)))
2001 return failure();
2002 rewriter.replaceOpWithNewOp<LLVM::BitcastOp>(bitcastOp, dstType,
2003 adaptor.getOperands(),
2004 properties, discardableAttrs);
2005 return success();
2006 }
2007};
2008
2009//===----------------------------------------------------------------------===//
2010// FuncOp conversion
2011//===----------------------------------------------------------------------===//
2012
2013class FuncConversionPattern : public SPIRVToLLVMConversion<spirv::FuncOp> {
2014public:
2015 using SPIRVToLLVMConversion<spirv::FuncOp>::SPIRVToLLVMConversion;
2016
2017 LogicalResult
2018 matchAndRewrite(spirv::FuncOp funcOp, OpAdaptor adaptor,
2019 ConversionPatternRewriter &rewriter) const override {
2020
2021 // Convert function signature. At the moment LLVMType converter is enough
2022 // for currently supported types.
2023 auto funcType = funcOp.getFunctionType();
2024 TypeConverter::SignatureConversion signatureConverter(
2025 funcType.getNumInputs());
2026 auto llvmType = static_cast<const LLVMTypeConverter *>(getTypeConverter())
2027 ->convertFunctionSignature(
2028 funcType, /*isVariadic=*/false,
2029 /*useBarePtrCallConv=*/false, signatureConverter);
2030 if (!llvmType)
2031 return failure();
2032
2033 // Create a new `LLVMFuncOp`
2034 Location loc = funcOp.getLoc();
2035 StringRef name = funcOp.getName();
2036 auto newFuncOp = LLVM::LLVMFuncOp::create(rewriter, loc, name, llvmType);
2037
2038 // Convert SPIR-V Function Control to equivalent LLVM function attribute
2039 MLIRContext *context = funcOp.getContext();
2040 switch (funcOp.getFunctionControl()) {
2041 case spirv::FunctionControl::Inline:
2042 newFuncOp.setAlwaysInline(true);
2043 break;
2044 case spirv::FunctionControl::DontInline:
2045 newFuncOp.setNoInline(true);
2046 break;
2047
2048#define DISPATCH(functionControl, llvmAttr) \
2049 case functionControl: \
2050 newFuncOp->setDiscardableAttr("passthrough", \
2051 ArrayAttr::get(context, {llvmAttr})); \
2052 break;
2053
2054 DISPATCH(spirv::FunctionControl::Pure,
2055 StringAttr::get(context, "readonly"));
2056 DISPATCH(spirv::FunctionControl::Const,
2057 StringAttr::get(context, "readnone"));
2058
2059#undef DISPATCH
2060
2061 // Default: if `spirv::FunctionControl::None`, then no attributes are
2062 // needed.
2063 default:
2064 break;
2065 }
2066
2067 rewriter.inlineRegionBefore(funcOp.getBody(), newFuncOp.getBody(),
2068 newFuncOp.end());
2069 if (failed(rewriter.convertRegionTypes(
2070 &newFuncOp.getBody(), *getTypeConverter(), &signatureConverter))) {
2071 return failure();
2072 }
2073 rewriter.eraseOp(funcOp);
2074 return success();
2075 }
2076};
2077
2078//===----------------------------------------------------------------------===//
2079// ModuleOp conversion
2080//===----------------------------------------------------------------------===//
2081
2082class ModuleConversionPattern : public SPIRVToLLVMConversion<spirv::ModuleOp> {
2083public:
2084 using SPIRVToLLVMConversion<spirv::ModuleOp>::SPIRVToLLVMConversion;
2085
2086 LogicalResult
2087 matchAndRewrite(spirv::ModuleOp spvModuleOp, OpAdaptor adaptor,
2088 ConversionPatternRewriter &rewriter) const override {
2089
2090 auto newModuleOp =
2091 ModuleOp::create(rewriter, spvModuleOp.getLoc(), spvModuleOp.getName());
2092 rewriter.inlineRegionBefore(spvModuleOp.getRegion(), newModuleOp.getBody());
2093
2094 // Remove the terminator block that was automatically added by builder
2095 rewriter.eraseBlock(&newModuleOp.getBodyRegion().back());
2096 rewriter.eraseOp(spvModuleOp);
2097 return success();
2098 }
2099};
2100
2101//===----------------------------------------------------------------------===//
2102// VectorShuffleOp conversion
2103//===----------------------------------------------------------------------===//
2104
2105class VectorShufflePattern
2106 : public SPIRVToLLVMConversion<spirv::VectorShuffleOp> {
2107public:
2108 using SPIRVToLLVMConversion<spirv::VectorShuffleOp>::SPIRVToLLVMConversion;
2109 LogicalResult
2110 matchAndRewrite(spirv::VectorShuffleOp op, OpAdaptor adaptor,
2111 ConversionPatternRewriter &rewriter) const override {
2112 Location loc = op.getLoc();
2113 auto components = adaptor.getComponents();
2114 auto vector1 = adaptor.getVector1();
2115 auto vector2 = adaptor.getVector2();
2116 int vector1Size = cast<VectorType>(vector1.getType()).getNumElements();
2117 int vector2Size = cast<VectorType>(vector2.getType()).getNumElements();
2118 if (vector1Size == vector2Size) {
2119 rewriter.replaceOpWithNewOp<LLVM::ShuffleVectorOp>(
2120 op, vector1, vector2,
2121 LLVM::convertArrayToIndices<int32_t>(components));
2122 return success();
2123 }
2124
2125 auto dstType = getTypeConverter()->convertType(op.getType());
2126 if (!dstType)
2127 return rewriter.notifyMatchFailure(op, "type conversion failed");
2128 auto scalarType = cast<VectorType>(dstType).getElementType();
2129 auto componentsArray = components.getValue();
2130 auto *context = rewriter.getContext();
2131 auto llvmI32Type = IntegerType::get(context, 32);
2132 Value targetOp = LLVM::PoisonOp::create(rewriter, loc, dstType);
2133 for (unsigned i = 0; i < componentsArray.size(); i++) {
2134 if (!isa<IntegerAttr>(componentsArray[i]))
2135 return op.emitError("unable to support non-constant component");
2136
2137 int indexVal = cast<IntegerAttr>(componentsArray[i]).getInt();
2138 if (indexVal == -1)
2139 continue;
2140
2141 int offsetVal = 0;
2142 Value baseVector = vector1;
2143 if (indexVal >= vector1Size) {
2144 offsetVal = vector1Size;
2145 baseVector = vector2;
2146 }
2147
2148 Value dstIndex = LLVM::ConstantOp::create(
2149 rewriter, loc, llvmI32Type,
2150 rewriter.getIntegerAttr(rewriter.getI32Type(), i));
2151 Value index = LLVM::ConstantOp::create(
2152 rewriter, loc, llvmI32Type,
2153 rewriter.getIntegerAttr(rewriter.getI32Type(), indexVal - offsetVal));
2154
2155 auto extractOp = LLVM::ExtractElementOp::create(rewriter, loc, scalarType,
2156 baseVector, index);
2157 targetOp = LLVM::InsertElementOp::create(rewriter, loc, dstType, targetOp,
2158 extractOp, dstIndex);
2159 }
2160 rewriter.replaceOp(op, targetOp);
2161 return success();
2162 }
2163};
2164} // namespace
2165
2166//===----------------------------------------------------------------------===//
2167// Pattern population
2168//===----------------------------------------------------------------------===//
2169
2171 spirv::ClientAPI clientAPI) {
2172 typeConverter.addConversion([&](spirv::ArrayType type) {
2173 return convertArrayType(type, typeConverter);
2174 });
2175 typeConverter.addConversion([&, clientAPI](spirv::PointerType type) {
2176 return convertPointerType(type, typeConverter, clientAPI);
2177 });
2178 typeConverter.addConversion([&](spirv::RuntimeArrayType type) {
2179 return convertRuntimeArrayType(type, typeConverter);
2180 });
2181 typeConverter.addConversion([&](spirv::StructType type) {
2182 return convertStructType(type, typeConverter);
2183 });
2184}
2185
2187 const LLVMTypeConverter &typeConverter, RewritePatternSet &patterns,
2188 spirv::ClientAPI clientAPI) {
2189 patterns.add<
2190 // Arithmetic ops
2191 DirectConversionPattern<spirv::IAddOp, LLVM::AddOp>,
2192 DirectConversionPattern<spirv::IMulOp, LLVM::MulOp>,
2193 DirectConversionPattern<spirv::ISubOp, LLVM::SubOp>,
2194 DirectConversionPattern<spirv::FAddOp, LLVM::FAddOp>,
2195 DirectConversionPattern<spirv::FDivOp, LLVM::FDivOp>,
2196 DirectConversionPattern<spirv::FMulOp, LLVM::FMulOp>,
2197 DirectConversionPattern<spirv::FNegateOp, LLVM::FNegOp>,
2198 DirectConversionPattern<spirv::FRemOp, LLVM::FRemOp>,
2199 DirectConversionPattern<spirv::FSubOp, LLVM::FSubOp>,
2200 DirectConversionPattern<spirv::SDivOp, LLVM::SDivOp>,
2201 DirectConversionPattern<spirv::SRemOp, LLVM::SRemOp>,
2202 DirectConversionPattern<spirv::UDivOp, LLVM::UDivOp>,
2203 DirectConversionPattern<spirv::UModOp, LLVM::URemOp>, FModPattern,
2204 SModPattern, VectorTimesScalarPattern, SNegatePattern,
2205 ArithmeticWithOverflowPattern<spirv::IAddCarryOp,
2206 LLVM::UAddWithOverflowOp>,
2207 ArithmeticWithOverflowPattern<spirv::ISubBorrowOp,
2208 LLVM::USubWithOverflowOp>,
2209
2210 // Bitwise ops
2211 BitFieldInsertPattern, BitFieldUExtractPattern, BitFieldSExtractPattern,
2212 DirectConversionPattern<spirv::BitCountOp, LLVM::CtPopOp>,
2213 DirectConversionPattern<spirv::BitReverseOp, LLVM::BitReverseOp>,
2214 DirectConversionPattern<spirv::BitwiseAndOp, LLVM::AndOp>,
2215 DirectConversionPattern<spirv::BitwiseOrOp, LLVM::OrOp>,
2216 DirectConversionPattern<spirv::BitwiseXorOp, LLVM::XOrOp>,
2217 NotPattern<spirv::NotOp>,
2218
2219 // Cast ops
2220 BitcastConversionPattern,
2221 DirectConversionPattern<spirv::ConvertFToSOp, LLVM::FPToSIOp>,
2222 DirectConversionPattern<spirv::ConvertFToUOp, LLVM::FPToUIOp>,
2223 DirectConversionPattern<spirv::ConvertSToFOp, LLVM::SIToFPOp>,
2224 DirectConversionPattern<spirv::ConvertUToFOp, LLVM::UIToFPOp>,
2225 IndirectCastPattern<spirv::FConvertOp, LLVM::FPExtOp, LLVM::FPTruncOp>,
2226 IndirectCastPattern<spirv::SConvertOp, LLVM::SExtOp, LLVM::TruncOp>,
2227 IndirectCastPattern<spirv::UConvertOp, LLVM::ZExtOp, LLVM::TruncOp>,
2228 DirectConversionPattern<spirv::ConvertPtrToUOp, LLVM::PtrToIntOp>,
2229 DirectConversionPattern<spirv::ConvertUToPtrOp, LLVM::IntToPtrOp>,
2230 DirectConversionPattern<spirv::PtrCastToGenericOp, LLVM::AddrSpaceCastOp>,
2231 DirectConversionPattern<spirv::GenericCastToPtrOp, LLVM::AddrSpaceCastOp>,
2232 DirectConversionPattern<spirv::GenericCastToPtrExplicitOp,
2233 LLVM::AddrSpaceCastOp>,
2234
2235 // Comparison ops
2236 IComparePattern<spirv::IEqualOp, LLVM::ICmpPredicate::eq>,
2237 IComparePattern<spirv::INotEqualOp, LLVM::ICmpPredicate::ne>,
2238 FComparePattern<spirv::FOrdEqualOp, LLVM::FCmpPredicate::oeq>,
2239 FComparePattern<spirv::FOrdGreaterThanOp, LLVM::FCmpPredicate::ogt>,
2240 FComparePattern<spirv::FOrdGreaterThanEqualOp, LLVM::FCmpPredicate::oge>,
2241 FComparePattern<spirv::FOrdLessThanEqualOp, LLVM::FCmpPredicate::ole>,
2242 FComparePattern<spirv::FOrdLessThanOp, LLVM::FCmpPredicate::olt>,
2243 FComparePattern<spirv::FOrdNotEqualOp, LLVM::FCmpPredicate::one>,
2244 FComparePattern<spirv::FUnordEqualOp, LLVM::FCmpPredicate::ueq>,
2245 FComparePattern<spirv::FUnordGreaterThanOp, LLVM::FCmpPredicate::ugt>,
2246 FComparePattern<spirv::FUnordGreaterThanEqualOp,
2247 LLVM::FCmpPredicate::uge>,
2248 FComparePattern<spirv::FUnordLessThanEqualOp, LLVM::FCmpPredicate::ule>,
2249 FComparePattern<spirv::FUnordLessThanOp, LLVM::FCmpPredicate::ult>,
2250 FComparePattern<spirv::FUnordNotEqualOp, LLVM::FCmpPredicate::une>,
2251 FComparePattern<spirv::OrderedOp, LLVM::FCmpPredicate::ord>,
2252 FComparePattern<spirv::UnorderedOp, LLVM::FCmpPredicate::uno>,
2253 IComparePattern<spirv::SGreaterThanOp, LLVM::ICmpPredicate::sgt>,
2254 IComparePattern<spirv::SGreaterThanEqualOp, LLVM::ICmpPredicate::sge>,
2255 IComparePattern<spirv::SLessThanEqualOp, LLVM::ICmpPredicate::sle>,
2256 IComparePattern<spirv::SLessThanOp, LLVM::ICmpPredicate::slt>,
2257 IComparePattern<spirv::UGreaterThanOp, LLVM::ICmpPredicate::ugt>,
2258 IComparePattern<spirv::UGreaterThanEqualOp, LLVM::ICmpPredicate::uge>,
2259 IComparePattern<spirv::ULessThanEqualOp, LLVM::ICmpPredicate::ule>,
2260 IComparePattern<spirv::ULessThanOp, LLVM::ICmpPredicate::ult>,
2261
2262 // Constant op
2263 ConstantScalarAndVectorPattern,
2264
2265 // Control Flow ops
2266 BranchConversionPattern, BranchConditionalConversionPattern,
2267 FunctionCallPattern, LoopPattern, SelectionPattern,
2268 ErasePattern<spirv::MergeOp>,
2269
2270 // Entry points and execution mode are handled separately.
2271 ErasePattern<spirv::EntryPointOp>, ExecutionModePattern,
2272
2273 // GLSL extended instruction set ops
2274 DirectConversionPattern<spirv::GLCeilOp, LLVM::FCeilOp>,
2275 DirectConversionPattern<spirv::GLCosOp, LLVM::CosOp>,
2276 DirectConversionPattern<spirv::GLExpOp, LLVM::ExpOp>,
2277 DirectConversionPattern<spirv::GLExp2Op, LLVM::Exp2Op>,
2278 DirectConversionPattern<spirv::GLFAbsOp, LLVM::FAbsOp>,
2279 DirectConversionPattern<spirv::GLFloorOp, LLVM::FFloorOp>,
2280 DirectConversionPattern<spirv::GLFmaOp, LLVM::FMAOp>,
2281 ClampPattern<spirv::GLFClampOp, LLVM::MinNumOp, LLVM::MaxNumOp>,
2282 ClampPattern<spirv::GLSClampOp, LLVM::SMinOp, LLVM::SMaxOp>,
2283 ClampPattern<spirv::GLUClampOp, LLVM::UMinOp, LLVM::UMaxOp>,
2284 DirectConversionPattern<spirv::GLFMaxOp, LLVM::MaxNumOp>,
2285 DirectConversionPattern<spirv::GLFMinOp, LLVM::MinNumOp>,
2286 DirectConversionPattern<spirv::GLNMaxOp, LLVM::MaxNumOp>,
2287 DirectConversionPattern<spirv::GLNMinOp, LLVM::MinNumOp>,
2288 DirectConversionPattern<spirv::GLLogOp, LLVM::LogOp>,
2289 DirectConversionPattern<spirv::GLLog2Op, LLVM::Log2Op>,
2290 DirectConversionPattern<spirv::GLPowOp, LLVM::PowOp>,
2291 DirectConversionPattern<spirv::GLRoundOp, LLVM::RoundOp>,
2292 DirectConversionPattern<spirv::GLRoundEvenOp, LLVM::RoundEvenOp>,
2293 DirectConversionPattern<spirv::GLSinOp, LLVM::SinOp>,
2294 DirectConversionPattern<spirv::GLSinhOp, LLVM::SinhOp>,
2295 DirectConversionPattern<spirv::GLCoshOp, LLVM::CoshOp>,
2296 DirectConversionPattern<spirv::GLSMaxOp, LLVM::SMaxOp>,
2297 DirectConversionPattern<spirv::GLSMinOp, LLVM::SMinOp>,
2298 DirectConversionPattern<spirv::GLSqrtOp, LLVM::SqrtOp>,
2299 DirectConversionPattern<spirv::GLUMaxOp, LLVM::UMaxOp>,
2300 DirectConversionPattern<spirv::GLUMinOp, LLVM::UMinOp>,
2301 DirectConversionPattern<spirv::GLTruncOp, LLVM::FTruncOp>,
2302 DirectConversionPattern<spirv::GLAsinOp, LLVM::ASinOp>,
2303 DirectConversionPattern<spirv::GLAcosOp, LLVM::ACosOp>,
2304 DirectConversionPattern<spirv::GLAtanOp, LLVM::ATanOp>,
2305 DirectConversionPattern<spirv::GLTanOp, LLVM::TanOp>,
2306 DirectConversionPattern<spirv::GLTanhOp, LLVM::TanhOp>,
2307 InverseSqrtPattern, SAbsPattern, FractPattern,
2308 SignPattern<spirv::GLFSignOp, /*isFloat=*/true>,
2309 SignPattern<spirv::GLSSignOp, /*isFloat=*/false>, GLFMixPattern,
2310
2311 // OpenCL extended instruction set ops
2312 DirectConversionPattern<spirv::CLCeilOp, LLVM::FCeilOp>,
2313 DirectConversionPattern<spirv::CLCosOp, LLVM::CosOp>,
2314 DirectConversionPattern<spirv::CLExpOp, LLVM::ExpOp>,
2315 DirectConversionPattern<spirv::CLExp2Op, LLVM::Exp2Op>,
2316 DirectConversionPattern<spirv::CLExp10Op, LLVM::Exp10Op>,
2317 DirectConversionPattern<spirv::CLFAbsOp, LLVM::FAbsOp>,
2318 DirectConversionPattern<spirv::CLFloorOp, LLVM::FFloorOp>,
2319 DirectConversionPattern<spirv::CLFmaOp, LLVM::FMAOp>,
2320 DirectConversionPattern<spirv::CLFMaxOp, LLVM::MaxNumOp>,
2321 DirectConversionPattern<spirv::CLFMinOp, LLVM::MinNumOp>,
2322 DirectConversionPattern<spirv::CLLogOp, LLVM::LogOp>,
2323 DirectConversionPattern<spirv::CLLog2Op, LLVM::Log2Op>,
2324 DirectConversionPattern<spirv::CLLog10Op, LLVM::Log10Op>,
2325 DirectConversionPattern<spirv::CLPowOp, LLVM::PowOp>,
2326 DirectConversionPattern<spirv::CLRintOp, LLVM::RintOp>,
2327 DirectConversionPattern<spirv::CLRoundOp, LLVM::RoundOp>,
2328 DirectConversionPattern<spirv::CLSinOp, LLVM::SinOp>,
2329 DirectConversionPattern<spirv::CLSinhOp, LLVM::SinhOp>,
2330 DirectConversionPattern<spirv::CLCoshOp, LLVM::CoshOp>,
2331 DirectConversionPattern<spirv::CLTanOp, LLVM::TanOp>,
2332 DirectConversionPattern<spirv::CLTanhOp, LLVM::TanhOp>,
2333 DirectConversionPattern<spirv::CLAsinOp, LLVM::ASinOp>,
2334 DirectConversionPattern<spirv::CLAcosOp, LLVM::ACosOp>,
2335 DirectConversionPattern<spirv::CLAtanOp, LLVM::ATanOp>,
2336 DirectConversionPattern<spirv::CLAtan2Op, LLVM::ATan2Op>,
2337 DirectConversionPattern<spirv::CLSqrtOp, LLVM::SqrtOp>,
2338 DirectConversionPattern<spirv::CLTruncOp, LLVM::FTruncOp>,
2339 DirectConversionPattern<spirv::CLCopysignOp, LLVM::CopySignOp>,
2340 DirectConversionPattern<spirv::CLFmodOp, LLVM::FRemOp>,
2341 DirectConversionPattern<spirv::CLSMaxOp, LLVM::SMaxOp>,
2342 DirectConversionPattern<spirv::CLSMinOp, LLVM::SMinOp>,
2343 DirectConversionPattern<spirv::CLUMaxOp, LLVM::UMaxOp>,
2344 DirectConversionPattern<spirv::CLUMinOp, LLVM::UMinOp>, CLMixPattern,
2345
2346 // Logical ops
2347 DirectConversionPattern<spirv::LogicalAndOp, LLVM::AndOp>,
2348 DirectConversionPattern<spirv::LogicalOrOp, LLVM::OrOp>,
2349 IComparePattern<spirv::LogicalEqualOp, LLVM::ICmpPredicate::eq>,
2350 IComparePattern<spirv::LogicalNotEqualOp, LLVM::ICmpPredicate::ne>,
2351 NotPattern<spirv::LogicalNotOp>,
2352
2353 // Memory ops
2354 AccessChainPattern, AddressOfPattern, LoadStorePattern<spirv::LoadOp>,
2355 LoadStorePattern<spirv::StoreOp>, VariablePattern,
2356
2357 // Miscellaneous ops
2358 CompositeExtractPattern, CompositeInsertPattern,
2359 DirectConversionPattern<spirv::SelectOp, LLVM::SelectOp>,
2360 DirectConversionPattern<spirv::UndefOp, LLVM::UndefOp>,
2361 VectorShufflePattern,
2362
2363 // Shift ops
2364 ShiftPattern<spirv::ShiftRightArithmeticOp, LLVM::AShrOp>,
2365 ShiftPattern<spirv::ShiftRightLogicalOp, LLVM::LShrOp>,
2366 ShiftPattern<spirv::ShiftLeftLogicalOp, LLVM::ShlOp>,
2367
2368 // Return ops
2369 ReturnPattern, ReturnValuePattern,
2370
2371 // Unreachable op
2372 UnreachablePattern,
2373
2374 // Barrier ops
2375 ControlBarrierPattern<spirv::ControlBarrierOp>,
2376 ControlBarrierPattern<spirv::INTELControlBarrierArriveOp>,
2377 ControlBarrierPattern<spirv::INTELControlBarrierWaitOp>,
2378
2379 // Group reduction operations
2380 GroupReducePattern<spirv::GroupIAddOp>,
2381 GroupReducePattern<spirv::GroupFAddOp>,
2382 GroupReducePattern<spirv::GroupFMinOp>,
2383 GroupReducePattern<spirv::GroupUMinOp>,
2384 GroupReducePattern<spirv::GroupSMinOp, /*Signed=*/true>,
2385 GroupReducePattern<spirv::GroupFMaxOp>,
2386 GroupReducePattern<spirv::GroupUMaxOp>,
2387 GroupReducePattern<spirv::GroupSMaxOp, /*Signed=*/true>,
2388 GroupReducePattern<spirv::GroupNonUniformIAddOp, /*Signed=*/false,
2389 /*NonUniform=*/true>,
2390 GroupReducePattern<spirv::GroupNonUniformFAddOp, /*Signed=*/false,
2391 /*NonUniform=*/true>,
2392 GroupReducePattern<spirv::GroupNonUniformIMulOp, /*Signed=*/false,
2393 /*NonUniform=*/true>,
2394 GroupReducePattern<spirv::GroupNonUniformFMulOp, /*Signed=*/false,
2395 /*NonUniform=*/true>,
2396 GroupReducePattern<spirv::GroupNonUniformSMinOp, /*Signed=*/true,
2397 /*NonUniform=*/true>,
2398 GroupReducePattern<spirv::GroupNonUniformUMinOp, /*Signed=*/false,
2399 /*NonUniform=*/true>,
2400 GroupReducePattern<spirv::GroupNonUniformFMinOp, /*Signed=*/false,
2401 /*NonUniform=*/true>,
2402 GroupReducePattern<spirv::GroupNonUniformSMaxOp, /*Signed=*/true,
2403 /*NonUniform=*/true>,
2404 GroupReducePattern<spirv::GroupNonUniformUMaxOp, /*Signed=*/false,
2405 /*NonUniform=*/true>,
2406 GroupReducePattern<spirv::GroupNonUniformFMaxOp, /*Signed=*/false,
2407 /*NonUniform=*/true>,
2408 GroupReducePattern<spirv::GroupNonUniformBitwiseAndOp, /*Signed=*/false,
2409 /*NonUniform=*/true>,
2410 GroupReducePattern<spirv::GroupNonUniformBitwiseOrOp, /*Signed=*/false,
2411 /*NonUniform=*/true>,
2412 GroupReducePattern<spirv::GroupNonUniformBitwiseXorOp, /*Signed=*/false,
2413 /*NonUniform=*/true>,
2414 GroupReducePattern<spirv::GroupNonUniformLogicalAndOp, /*Signed=*/false,
2415 /*NonUniform=*/true>,
2416 GroupReducePattern<spirv::GroupNonUniformLogicalOrOp, /*Signed=*/false,
2417 /*NonUniform=*/true>,
2418 GroupReducePattern<spirv::GroupNonUniformLogicalXorOp, /*Signed=*/false,
2419 /*NonUniform=*/true>>(patterns.getContext(),
2420 typeConverter);
2421
2422 patterns.add<GlobalVariablePattern>(clientAPI, patterns.getContext(),
2423 typeConverter);
2424 // pi / 180
2425 patterns.add<ScalePattern<spirv::GLRadiansOp>>(
2426 0.017453292519943295, patterns.getContext(), typeConverter);
2427 // 180 / pi
2428 patterns.add<ScalePattern<spirv::GLDegreesOp>>(
2429 57.29577951308232, patterns.getContext(), typeConverter);
2430}
2431
2433 const LLVMTypeConverter &typeConverter, RewritePatternSet &patterns) {
2434 patterns.add<FuncConversionPattern>(patterns.getContext(), typeConverter);
2435}
2436
2438 const LLVMTypeConverter &typeConverter, RewritePatternSet &patterns) {
2439 patterns.add<ModuleConversionPattern>(patterns.getContext(), typeConverter);
2440}
2441
2442//===----------------------------------------------------------------------===//
2443// Pre-conversion hooks
2444//===----------------------------------------------------------------------===//
2445
2446/// Hook for descriptor set and binding number encoding.
2447void mlir::encodeBindAttribute(ModuleOp module) {
2448 auto spvModules = module.getOps<spirv::ModuleOp>();
2449 for (auto spvModule : spvModules) {
2450 spvModule.walk([&](spirv::GlobalVariableOp op) {
2451 IntegerAttr descriptorSet = op.getDescriptorSetAttr();
2452 IntegerAttr binding = op.getBindingAttr();
2453 // For every global variable in the module, get the ones with descriptor
2454 // set and binding numbers.
2455 if (descriptorSet && binding) {
2456 // Encode these numbers into the variable's symbolic name. If the
2457 // SPIR-V module has a name, add it at the beginning.
2458 auto moduleAndName =
2459 spvModule.getName().has_value()
2460 ? spvModule.getName()->str() + "_" + op.getSymName().str()
2461 : op.getSymName().str();
2462 std::string name =
2463 llvm::formatv("{0}_descriptor_set{1}_binding{2}", moduleAndName,
2464 std::to_string(descriptorSet.getInt()),
2465 std::to_string(binding.getInt()));
2466 auto nameAttr = StringAttr::get(op->getContext(), name);
2467
2468 // Replace all symbol uses and set the new symbol name. Finally, remove
2469 // descriptor set and binding attributes.
2470 if (failed(SymbolTable::replaceAllSymbolUses(op, nameAttr, spvModule)))
2471 op.emitError("unable to replace all symbol uses for ") << name;
2472 SymbolTable::setSymbolName(op, nameAttr);
2473 op.removeDescriptorSetAttr();
2474 op.removeBindingAttr();
2475 }
2476 });
2477 }
2478}
return success()
static LLVM::CallOp createSPIRVBuiltinCall(Location loc, ConversionPatternRewriter &rewriter, LLVM::LLVMFuncOp func, ValueRange args)
static LLVM::LLVMFuncOp lookupOrCreateSPIRVFn(Operation *symbolTable, StringRef name, ArrayRef< Type > paramTypes, Type resultType, bool isMemNone, bool isConvergent)
lhs
b
Return true if permutation is a valid permutation of the outer_dims_perm (case OuterOrInnerPerm::Oute...
ArrayAttr()
b getContext())
static Value max(ImplicitLocOpBuilder &builder, Value value, Value bound)
static Type getElementType(Type type, ArrayRef< int32_t > indices, function_ref< InFlightDiagnostic(StringRef)> emitErrorFn)
Walks the given type hierarchy with the given indices, potentially down to component granularity,...
Definition SPIRVOps.cpp:229
static Value optionallyTruncateOrExtend(Location loc, Value value, Type llvmType, PatternRewriter &rewriter)
Utility function for bitfield ops:
static Value createFPConstant(Location loc, Type srcType, Type dstType, PatternRewriter &rewriter, double value)
Creates llvm.mlir.constant with a floating-point scalar or vector value.
static Value createI32ConstantOf(Location loc, PatternRewriter &rewriter, unsigned value)
Creates LLVM dialect constant with the given value.
static Type convertPointerType(spirv::PointerType type, const TypeConverter &converter, spirv::ClientAPI clientAPI)
Converts SPIR-V pointer type to LLVM pointer.
static Value processCountOrOffset(Location loc, Value value, Type srcType, Type dstType, const TypeConverter &converter, ConversionPatternRewriter &rewriter)
Utility function for bitfield ops: BitFieldInsert, BitFieldSExtract and BitFieldUExtract.
static unsigned getBitWidth(Type type)
Returns the bit width of integer, float or vector of float or integer values.
static LogicalResult replaceWithLoadOrStore(Operation *op, ValueRange operands, ConversionPatternRewriter &rewriter, const TypeConverter &typeConverter, unsigned alignment, bool isVolatile, bool isNonTemporal)
Utility for spirv.Load and spirv.Store conversion.
static Type convertStructTypePacked(spirv::StructType type, const TypeConverter &converter)
Converts SPIR-V struct with no offset to packed LLVM struct.
static bool isSignedIntegerOrVector(Type type)
Returns true if the given type is a signed integer or vector type.
static bool isUnsignedIntegerOrVector(Type type)
Returns true if the given type is an unsigned integer or vector type.
static std::optional< Type > convertRuntimeArrayType(spirv::RuntimeArrayType type, TypeConverter &converter)
Converts SPIR-V runtime array to LLVM array.
static Value optionallyBroadcast(Location loc, Value value, Type srcType, const TypeConverter &typeConverter, ConversionPatternRewriter &rewriter)
Broadcasts the value. If srcType is a scalar, the value remains unchanged.
static Value createConstantAllBitsSet(Location loc, Type srcType, Type dstType, PatternRewriter &rewriter)
Creates llvm.mlir.constant with all bits set for the given type.
static unsigned getLLVMTypeBitWidth(Type type)
Returns the bit width of LLVMType integer or vector.
static std::optional< uint64_t > getIntegerOrVectorElementWidth(Type type)
Returns the width of an integer or of the element type of an integer vector, if applicable.
#define DISPATCH(functionControl, llvmAttr)
static Type convertStructTypeWithOffset(spirv::StructType type, const TypeConverter &converter)
Converts SPIR-V struct with a regular (according to VulkanLayoutUtils) offset to LLVM struct.
static Type convertStructType(spirv::StructType type, const TypeConverter &converter)
Converts SPIR-V struct to LLVM struct.
static Value broadcast(Location loc, Value toBroadcast, unsigned numElements, const TypeConverter &typeConverter, ConversionPatternRewriter &rewriter)
Broadcasts the value to vector with numElements number of elements.
static Value createIntegerConstant(Location loc, Type srcType, Type dstType, PatternRewriter &rewriter, IntegerAttr scalarAttr)
Creates llvm.mlir.constant with a scalar or vector integer value, broadcasting scalarAttr across the ...
static std::optional< Type > convertArrayType(spirv::ArrayType type, TypeConverter &converter)
Converts SPIR-V array type to LLVM array.
#define div(a, b)
#define rem(a, b)
OpListType::iterator iterator
Definition Block.h:164
OpListType & getOperations()
Definition Block.h:161
Operation * getTerminator()
Get the terminator operation of this block.
Definition Block.cpp:249
BlockArgListType getArguments()
Definition Block.h:111
IntegerAttr getIntegerAttr(Type type, int64_t value)
Definition Builders.cpp:237
FloatAttr getFloatAttr(Type type, double value)
Definition Builders.cpp:263
IntegerType getI32Type()
Definition Builders.cpp:71
MLIRContext * getContext() const
Definition Builders.h:56
static DenseElementsAttr get(ShapedType type, ArrayRef< Attribute > values)
Constructs a dense elements attribute from an array of element values.
Conversion from types to the LLVM IR dialect.
This class defines the main interface for locations in MLIR and acts as a non-nullable wrapper around...
Definition Location.h:76
NamedAttrList is array of NamedAttributes that tracks whether it is sorted and does some basic work t...
NamedAttribute represents a combination of a name and an Attribute value.
Definition Attributes.h:164
This class helps build Operations.
Definition Builders.h:210
Operation is the basic unit of execution within MLIR.
Definition Operation.h:87
Region & getRegion(unsigned index)
Returns the region held by this operation at position 'index'.
Definition Operation.h:738
Location getLoc()
The source location the operation was defined or derived from.
Definition Operation.h:240
operand_range getOperands()
Returns an iterator on the underlying Value's.
Definition Operation.h:403
A special type of RewriterBase that coordinates the application of a rewrite pattern on the current I...
MLIRContext * getContext() const
RewritePatternSet & add(ConstructorArg &&arg, ConstructorArgs &&...args)
Add an instance of each of the pattern types 'Ts' to the pattern list with the given arguments.
static LogicalResult replaceAllSymbolUses(StringAttr oldSymbol, StringAttr newSymbol, Operation *from)
Attempt to replace all uses of the given symbol 'oldSymbol' with the provided symbol 'newSymbol' that...
static Operation * lookupSymbolIn(Operation *op, StringAttr symbol)
Returns the operation registered with the given symbol name with the regions of 'symbolTableOp'.
static void setSymbolName(Operation *symbol, StringAttr name)
Sets the name of the given symbol operation.
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
Definition Types.h:74
bool isSignedInteger() const
Return true if this is a signed integer type (with the specified width).
Definition Types.cpp:78
bool isUnsignedInteger() const
Return true if this is an unsigned integer type (with the specified width).
Definition Types.cpp:90
bool isIntOrFloat() const
Return true if this is an integer (of any signedness) or a float type.
Definition Types.cpp:118
unsigned getIntOrFloatBitWidth() const
Return the bit width of an integer or a float type, assert failure on other types.
Definition Types.cpp:124
This class provides an abstraction over the different types of ranges over Values.
Definition ValueRange.h:389
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
Definition Value.h:96
Type getType() const
Return the type of this value.
Definition Value.h:105
static spirv::StructType decorateType(spirv::StructType structType)
Returns a new StructType with layout decoration.
static DenseArrayAttrImpl get(MLIRContext *context, ArrayRef< int32_t > content)
Type getElementType() const
unsigned getArrayStride() const
Returns the array stride in bytes.
unsigned getNumElements() const
StorageClass getStorageClass() const
unsigned getArrayStride() const
Returns the array stride in bytes.
SPIR-V struct type.
Definition SPIRVTypes.h:274
void getMemberDecorations(SmallVectorImpl< StructType::MemberDecorationInfo > &memberDecorations) const
TypeRange getElementTypes() const
bool isCompatibleType(Type type)
Returns true if the given type is compatible with the LLVM dialect.
detail::InFlightRemark failed(Location loc, RemarkOpts opts)
Report an optimization remark that failed.
Definition Remarks.h:732
Include the generated interface declarations.
unsigned storageClassToAddressSpace(spirv::ClientAPI clientAPI, spirv::StorageClass storageClass)
void populateSPIRVToLLVMTypeConversion(LLVMTypeConverter &typeConverter, spirv::ClientAPI clientAPIForAddressSpaceMapping=spirv::ClientAPI::Unknown)
Populates type conversions with additional SPIR-V types.
void populateSPIRVToLLVMFunctionConversionPatterns(const LLVMTypeConverter &typeConverter, RewritePatternSet &patterns)
Populates the given list with patterns for function conversion from SPIR-V to LLVM.
Type getElementTypeOrSelf(Type type)
Return the element type or return the type itself.
detail::DenseArrayAttrImpl< int32_t > DenseI32ArrayAttr
void populateSPIRVToLLVMConversionPatterns(const LLVMTypeConverter &typeConverter, RewritePatternSet &patterns, spirv::ClientAPI clientAPIForAddressSpaceMapping=spirv::ClientAPI::Unknown)
Populates the given list with patterns that convert from SPIR-V to LLVM.
void encodeBindAttribute(ModuleOp module)
Encodes global variable's descriptor set and binding into its name if they both exist.
void populateSPIRVToLLVMModuleConversionPatterns(const LLVMTypeConverter &typeConverter, RewritePatternSet &patterns)
Populates the given patterns for module conversion from SPIR-V to LLVM.
Structure used by default as a "marker" when no "Properties" are set on an Operation.