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