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