MLIR 23.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 `spirv.ExecutionMode` into a global struct constant that holds
653/// execution mode information.
654class ExecutionModePattern
655 : public SPIRVToLLVMConversion<spirv::ExecutionModeOp> {
656public:
657 using SPIRVToLLVMConversion<spirv::ExecutionModeOp>::SPIRVToLLVMConversion;
658
659 LogicalResult
660 matchAndRewrite(spirv::ExecutionModeOp op, OpAdaptor adaptor,
661 ConversionPatternRewriter &rewriter) const override {
662 // First, create the global struct's name that would be associated with
663 // this entry point's execution mode. We set it to be:
664 // __spv__{SPIR-V module name}_{function name}_execution_mode_info_{mode}
665 ModuleOp module = op->getParentOfType<ModuleOp>();
666 spirv::ExecutionModeAttr executionModeAttr = op.getExecutionModeAttr();
667 std::string moduleName;
668 if (module.getName().has_value())
669 moduleName = "_" + module.getName()->str();
670 else
671 moduleName = "";
672 std::string executionModeInfoName = llvm::formatv(
673 "__spv_{0}_{1}_execution_mode_info_{2}", moduleName, op.getFn().str(),
674 static_cast<uint32_t>(executionModeAttr.getValue()));
675
676 MLIRContext *context = rewriter.getContext();
677 OpBuilder::InsertionGuard guard(rewriter);
678 rewriter.setInsertionPointToStart(module.getBody());
679
680 // Create a struct type, corresponding to the C struct below.
681 // struct {
682 // int32_t executionMode;
683 // int32_t values[]; // optional values
684 // };
685 auto llvmI32Type = IntegerType::get(context, 32);
686 SmallVector<Type, 2> fields;
687 fields.push_back(llvmI32Type);
688 ArrayAttr values = op.getValues();
689 if (!values.empty()) {
690 auto arrayType = LLVM::LLVMArrayType::get(llvmI32Type, values.size());
691 fields.push_back(arrayType);
692 }
693 auto structType = LLVM::LLVMStructType::getLiteral(context, fields);
694
695 // Create `llvm.mlir.global` with initializer region containing one block.
696 auto global = LLVM::GlobalOp::create(
697 rewriter, UnknownLoc::get(context), structType, /*isConstant=*/true,
698 LLVM::Linkage::External, executionModeInfoName, Attribute(),
699 /*alignment=*/0);
700 Location loc = global.getLoc();
701 Region &region = global.getInitializerRegion();
702 Block *block = rewriter.createBlock(&region);
703
704 // Initialize the struct and set the execution mode value.
705 rewriter.setInsertionPointToStart(block);
706 Value structValue = LLVM::PoisonOp::create(rewriter, loc, structType);
707 Value executionMode = LLVM::ConstantOp::create(
708 rewriter, loc, llvmI32Type,
709 rewriter.getI32IntegerAttr(
710 static_cast<uint32_t>(executionModeAttr.getValue())));
711 SmallVector<int64_t> position{0};
712 structValue = LLVM::InsertValueOp::create(rewriter, loc, structValue,
713 executionMode, position);
714
715 // Insert extra operands if they exist into execution mode info struct.
716 for (unsigned i = 0, e = values.size(); i < e; ++i) {
717 auto attr = values.getValue()[i];
718 Value entry = LLVM::ConstantOp::create(rewriter, loc, llvmI32Type, attr);
719 structValue = LLVM::InsertValueOp::create(
720 rewriter, loc, structValue, entry, ArrayRef<int64_t>({1, i}));
721 }
722 LLVM::ReturnOp::create(rewriter, loc, ArrayRef<Value>({structValue}));
723 rewriter.eraseOp(op);
724 return success();
725 }
726};
727
728/// Converts `spirv.GlobalVariable` to `llvm.mlir.global`. Note that SPIR-V
729/// global returns a pointer, whereas in LLVM dialect the global holds an actual
730/// value. This difference is handled by `spirv.mlir.addressof` and
731/// `llvm.mlir.addressof`ops that both return a pointer.
732class GlobalVariablePattern
733 : public SPIRVToLLVMConversion<spirv::GlobalVariableOp> {
734public:
735 template <typename... Args>
736 GlobalVariablePattern(spirv::ClientAPI clientAPI, Args &&...args)
737 : SPIRVToLLVMConversion<spirv::GlobalVariableOp>(
738 std::forward<Args>(args)...),
739 clientAPI(clientAPI) {}
740
741 LogicalResult
742 matchAndRewrite(spirv::GlobalVariableOp op, OpAdaptor adaptor,
743 ConversionPatternRewriter &rewriter) const override {
744 // Currently, there is no support of initialization with a constant value in
745 // SPIR-V dialect. Specialization constants are not considered as well.
746 if (op.getInitializer())
747 return failure();
748
749 auto srcType = cast<spirv::PointerType>(op.getType());
750 auto dstType = getTypeConverter()->convertType(srcType.getPointeeType());
751 if (!dstType)
752 return rewriter.notifyMatchFailure(op, "type conversion failed");
753
754 // Limit conversion to the current invocation only or `StorageBuffer`
755 // required by SPIR-V runner.
756 // This is okay because multiple invocations are not supported yet.
757 auto storageClass = srcType.getStorageClass();
758 switch (storageClass) {
759 case spirv::StorageClass::Input:
760 case spirv::StorageClass::Private:
761 case spirv::StorageClass::Output:
762 case spirv::StorageClass::StorageBuffer:
763 case spirv::StorageClass::UniformConstant:
764 break;
765 default:
766 return failure();
767 }
768
769 // LLVM dialect spec: "If the global value is a constant, storing into it is
770 // not allowed.". This corresponds to SPIR-V 'Input' and 'UniformConstant'
771 // storage class that is read-only.
772 bool isConstant = (storageClass == spirv::StorageClass::Input) ||
773 (storageClass == spirv::StorageClass::UniformConstant);
774 // SPIR-V spec: "By default, functions and global variables are private to a
775 // module and cannot be accessed by other modules. However, a module may be
776 // written to export or import functions and global (module scope)
777 // variables.". Therefore, map 'Private' storage class to private linkage,
778 // 'Input' and 'Output' to external linkage.
779 auto linkage = storageClass == spirv::StorageClass::Private
780 ? LLVM::Linkage::Private
781 : LLVM::Linkage::External;
782 StringAttr locationAttrName = op.getLocationAttrName();
783 IntegerAttr locationAttr = op.getLocationAttr();
784 auto newGlobalOp = rewriter.replaceOpWithNewOp<LLVM::GlobalOp>(
785 op, dstType, isConstant, linkage, op.getSymName(), Attribute(),
786 /*alignment=*/0, storageClassToAddressSpace(clientAPI, storageClass));
787
788 // Attach location attribute if applicable
789 if (locationAttr)
790 newGlobalOp->setAttr(locationAttrName, locationAttr);
791
792 return success();
793 }
794
795private:
796 spirv::ClientAPI clientAPI;
797};
798
799/// Converts SPIR-V cast ops that do not have straightforward LLVM
800/// equivalent in LLVM dialect.
801template <typename SPIRVOp, typename LLVMExtOp, typename LLVMTruncOp>
802class IndirectCastPattern : public SPIRVToLLVMConversion<SPIRVOp> {
803public:
804 using SPIRVToLLVMConversion<SPIRVOp>::SPIRVToLLVMConversion;
805
806 LogicalResult
807 matchAndRewrite(SPIRVOp op, typename SPIRVOp::Adaptor adaptor,
808 ConversionPatternRewriter &rewriter) const override {
809
810 Type fromType = op.getOperand().getType();
811 Type toType = op.getType();
812
813 auto dstType = this->getTypeConverter()->convertType(toType);
814 if (!dstType)
815 return rewriter.notifyMatchFailure(op, "type conversion failed");
816
817 if (getBitWidth(fromType) < getBitWidth(toType)) {
818 rewriter.template replaceOpWithNewOp<LLVMExtOp>(op, dstType,
819 adaptor.getOperands());
820 return success();
821 }
822 if (getBitWidth(fromType) > getBitWidth(toType)) {
823 rewriter.template replaceOpWithNewOp<LLVMTruncOp>(op, dstType,
824 adaptor.getOperands());
825 return success();
826 }
827 return failure();
828 }
829};
830
831class FunctionCallPattern
832 : public SPIRVToLLVMConversion<spirv::FunctionCallOp> {
833public:
834 using SPIRVToLLVMConversion<spirv::FunctionCallOp>::SPIRVToLLVMConversion;
835
836 LogicalResult
837 matchAndRewrite(spirv::FunctionCallOp callOp, OpAdaptor adaptor,
838 ConversionPatternRewriter &rewriter) const override {
839 if (callOp.getNumResults() == 0) {
840 auto newOp = rewriter.replaceOpWithNewOp<LLVM::CallOp>(
841 callOp, TypeRange(), adaptor.getOperands(), callOp->getAttrs());
842 newOp.getProperties().operandSegmentSizes = {
843 static_cast<int32_t>(adaptor.getOperands().size()), 0};
844 newOp.getProperties().op_bundle_sizes = rewriter.getDenseI32ArrayAttr({});
845 return success();
846 }
847
848 // Function returns a single result.
849 auto dstType = getTypeConverter()->convertType(callOp.getType(0));
850 if (!dstType)
851 return rewriter.notifyMatchFailure(callOp, "type conversion failed");
852 auto newOp = rewriter.replaceOpWithNewOp<LLVM::CallOp>(
853 callOp, dstType, adaptor.getOperands(), callOp->getAttrs());
854 newOp.getProperties().operandSegmentSizes = {
855 static_cast<int32_t>(adaptor.getOperands().size()), 0};
856 newOp.getProperties().op_bundle_sizes = rewriter.getDenseI32ArrayAttr({});
857 return success();
858 }
859};
860
861/// Converts SPIR-V floating-point comparisons to llvm.fcmp "predicate"
862template <typename SPIRVOp, LLVM::FCmpPredicate predicate>
863class FComparePattern : public SPIRVToLLVMConversion<SPIRVOp> {
864public:
865 using SPIRVToLLVMConversion<SPIRVOp>::SPIRVToLLVMConversion;
866
867 LogicalResult
868 matchAndRewrite(SPIRVOp op, typename SPIRVOp::Adaptor adaptor,
869 ConversionPatternRewriter &rewriter) const override {
870
871 auto dstType = this->getTypeConverter()->convertType(op.getType());
872 if (!dstType)
873 return rewriter.notifyMatchFailure(op, "type conversion failed");
874
875 rewriter.template replaceOpWithNewOp<LLVM::FCmpOp>(
876 op, dstType, predicate, op.getOperand1(), op.getOperand2());
877 return success();
878 }
879};
880
881/// Converts SPIR-V integer comparisons to llvm.icmp "predicate"
882template <typename SPIRVOp, LLVM::ICmpPredicate predicate>
883class IComparePattern : public SPIRVToLLVMConversion<SPIRVOp> {
884public:
885 using SPIRVToLLVMConversion<SPIRVOp>::SPIRVToLLVMConversion;
886
887 LogicalResult
888 matchAndRewrite(SPIRVOp op, typename SPIRVOp::Adaptor adaptor,
889 ConversionPatternRewriter &rewriter) const override {
890
891 auto dstType = this->getTypeConverter()->convertType(op.getType());
892 if (!dstType)
893 return rewriter.notifyMatchFailure(op, "type conversion failed");
894
895 rewriter.template replaceOpWithNewOp<LLVM::ICmpOp>(
896 op, dstType, predicate, op.getOperand1(), op.getOperand2());
897 return success();
898 }
899};
900
901class InverseSqrtPattern
902 : public SPIRVToLLVMConversion<spirv::GLInverseSqrtOp> {
903public:
904 using SPIRVToLLVMConversion<spirv::GLInverseSqrtOp>::SPIRVToLLVMConversion;
905
906 LogicalResult
907 matchAndRewrite(spirv::GLInverseSqrtOp op, OpAdaptor adaptor,
908 ConversionPatternRewriter &rewriter) const override {
909 auto srcType = op.getType();
910 auto dstType = getTypeConverter()->convertType(srcType);
911 if (!dstType)
912 return rewriter.notifyMatchFailure(op, "type conversion failed");
913
914 Location loc = op.getLoc();
915 Value one = createFPConstant(loc, srcType, dstType, rewriter, 1.0);
916 Value sqrt = LLVM::SqrtOp::create(rewriter, loc, dstType, op.getOperand());
917 rewriter.replaceOpWithNewOp<LLVM::FDivOp>(op, dstType, one, sqrt);
918 return success();
919 }
920};
921
922/// Converts `spirv.SNegate` to `0 - x`.
923class SNegatePattern : public SPIRVToLLVMConversion<spirv::SNegateOp> {
924public:
925 using SPIRVToLLVMConversion<spirv::SNegateOp>::SPIRVToLLVMConversion;
926
927 LogicalResult
928 matchAndRewrite(spirv::SNegateOp op, OpAdaptor adaptor,
929 ConversionPatternRewriter &rewriter) const override {
930 Type srcType = op.getType();
931 Type dstType = getTypeConverter()->convertType(srcType);
932 if (!dstType)
933 return rewriter.notifyMatchFailure(op, "type conversion failed");
934
935 Location loc = op.getLoc();
936 IntegerAttr zeroAttr = rewriter.getIntegerAttr(
937 cast<IntegerType>(getElementTypeOrSelf(srcType)), 0);
938 Value zero =
939 createIntegerConstant(loc, srcType, dstType, rewriter, zeroAttr);
940 rewriter.replaceOpWithNewOp<LLVM::SubOp>(op, dstType, zero,
941 adaptor.getOperand());
942 return success();
943 }
944};
945
946/// Converts the GLSL clamp ops (FClamp, SClamp, UClamp) into a nested
947/// min/max sequence, following the op semantics `min(max(x, minVal), maxVal)`.
948template <typename SPIRVOp, typename LLVMMinOp, typename LLVMMaxOp>
949class ClampPattern : public SPIRVToLLVMConversion<SPIRVOp> {
950public:
951 using SPIRVToLLVMConversion<SPIRVOp>::SPIRVToLLVMConversion;
952
953 LogicalResult
954 matchAndRewrite(SPIRVOp op, typename SPIRVOp::Adaptor adaptor,
955 ConversionPatternRewriter &rewriter) const override {
956 Type dstType = this->getTypeConverter()->convertType(op.getType());
957 if (!dstType)
958 return rewriter.notifyMatchFailure(op, "type conversion failed");
959
960 Location loc = op.getLoc();
961 Value max = LLVMMaxOp::create(rewriter, loc, dstType, adaptor.getX(),
962 adaptor.getY());
963 rewriter.template replaceOpWithNewOp<LLVMMinOp>(op, dstType, max,
964 adaptor.getZ());
965 return success();
966 }
967};
968
969/// Converts `spirv.Load` and `spirv.Store` to LLVM dialect.
970template <typename SPIRVOp>
971class LoadStorePattern : public SPIRVToLLVMConversion<SPIRVOp> {
972public:
973 using SPIRVToLLVMConversion<SPIRVOp>::SPIRVToLLVMConversion;
974
975 LogicalResult
976 matchAndRewrite(SPIRVOp op, typename SPIRVOp::Adaptor adaptor,
977 ConversionPatternRewriter &rewriter) const override {
978 if (!op.getMemoryAccess()) {
979 return replaceWithLoadOrStore(op, adaptor.getOperands(), rewriter,
980 *this->getTypeConverter(), /*alignment=*/0,
981 /*isVolatile=*/false,
982 /*isNonTemporal=*/false);
983 }
984 auto memoryAccess = *op.getMemoryAccess();
985 switch (memoryAccess) {
986 case spirv::MemoryAccess::Aligned:
987 case spirv::MemoryAccess::None:
988 case spirv::MemoryAccess::Nontemporal:
989 case spirv::MemoryAccess::Volatile: {
990 unsigned alignment =
991 memoryAccess == spirv::MemoryAccess::Aligned ? *op.getAlignment() : 0;
992 bool isNonTemporal = memoryAccess == spirv::MemoryAccess::Nontemporal;
993 bool isVolatile = memoryAccess == spirv::MemoryAccess::Volatile;
994 return replaceWithLoadOrStore(op, adaptor.getOperands(), rewriter,
995 *this->getTypeConverter(), alignment,
996 isVolatile, isNonTemporal);
997 }
998 default:
999 // There is no support of other memory access attributes.
1000 return failure();
1001 }
1002 }
1003};
1004
1005/// Converts `spirv.Not` and `spirv.LogicalNot` into LLVM dialect.
1006template <typename SPIRVOp>
1007class NotPattern : public SPIRVToLLVMConversion<SPIRVOp> {
1008public:
1009 using SPIRVToLLVMConversion<SPIRVOp>::SPIRVToLLVMConversion;
1010
1011 LogicalResult
1012 matchAndRewrite(SPIRVOp notOp, typename SPIRVOp::Adaptor adaptor,
1013 ConversionPatternRewriter &rewriter) const override {
1014 auto srcType = notOp.getType();
1015 auto dstType = this->getTypeConverter()->convertType(srcType);
1016 if (!dstType)
1017 return rewriter.notifyMatchFailure(notOp, "type conversion failed");
1018
1019 Location loc = notOp.getLoc();
1020 Value mask = createConstantAllBitsSet(loc, srcType, dstType, rewriter);
1021 rewriter.template replaceOpWithNewOp<LLVM::XOrOp>(notOp, dstType,
1022 notOp.getOperand(), mask);
1023 return success();
1024 }
1025};
1026
1027/// A template pattern that erases the given `SPIRVOp`.
1028template <typename SPIRVOp>
1029class ErasePattern : public SPIRVToLLVMConversion<SPIRVOp> {
1030public:
1031 using SPIRVToLLVMConversion<SPIRVOp>::SPIRVToLLVMConversion;
1032
1033 LogicalResult
1034 matchAndRewrite(SPIRVOp op, typename SPIRVOp::Adaptor adaptor,
1035 ConversionPatternRewriter &rewriter) const override {
1036 rewriter.eraseOp(op);
1037 return success();
1038 }
1039};
1040
1041class ReturnPattern : public SPIRVToLLVMConversion<spirv::ReturnOp> {
1042public:
1043 using SPIRVToLLVMConversion<spirv::ReturnOp>::SPIRVToLLVMConversion;
1044
1045 LogicalResult
1046 matchAndRewrite(spirv::ReturnOp returnOp, OpAdaptor adaptor,
1047 ConversionPatternRewriter &rewriter) const override {
1048 rewriter.replaceOpWithNewOp<LLVM::ReturnOp>(returnOp, ArrayRef<Type>(),
1049 ArrayRef<Value>());
1050 return success();
1051 }
1052};
1053
1054class ReturnValuePattern : public SPIRVToLLVMConversion<spirv::ReturnValueOp> {
1055public:
1056 using SPIRVToLLVMConversion<spirv::ReturnValueOp>::SPIRVToLLVMConversion;
1057
1058 LogicalResult
1059 matchAndRewrite(spirv::ReturnValueOp returnValueOp, OpAdaptor adaptor,
1060 ConversionPatternRewriter &rewriter) const override {
1061 rewriter.replaceOpWithNewOp<LLVM::ReturnOp>(returnValueOp, ArrayRef<Type>(),
1062 adaptor.getOperands());
1063 return success();
1064 }
1065};
1066
1067class UnreachablePattern : public SPIRVToLLVMConversion<spirv::UnreachableOp> {
1068public:
1069 using SPIRVToLLVMConversion<spirv::UnreachableOp>::SPIRVToLLVMConversion;
1070
1071 LogicalResult
1072 matchAndRewrite(spirv::UnreachableOp unreachableOp, OpAdaptor adaptor,
1073 ConversionPatternRewriter &rewriter) const override {
1074 rewriter.replaceOpWithNewOp<LLVM::UnreachableOp>(unreachableOp);
1075 return success();
1076 }
1077};
1078
1079static LLVM::LLVMFuncOp lookupOrCreateSPIRVFn(Operation *symbolTable,
1080 StringRef name,
1081 ArrayRef<Type> paramTypes,
1082 Type resultType,
1083 bool convergent = true) {
1084 auto func = dyn_cast_or_null<LLVM::LLVMFuncOp>(
1085 SymbolTable::lookupSymbolIn(symbolTable, name));
1086 if (func)
1087 return func;
1088
1089 OpBuilder b(symbolTable->getRegion(0));
1090 func = LLVM::LLVMFuncOp::create(
1091 b, symbolTable->getLoc(), name,
1092 LLVM::LLVMFunctionType::get(resultType, paramTypes));
1093 func.setCConv(LLVM::cconv::CConv::SPIR_FUNC);
1094 func.setConvergent(convergent);
1095 func.setNoUnwind(true);
1096 func.setWillReturn(true);
1097 return func;
1098}
1099
1100static LLVM::CallOp createSPIRVBuiltinCall(Location loc, OpBuilder &builder,
1101 LLVM::LLVMFuncOp func,
1102 ValueRange args) {
1103 auto call = LLVM::CallOp::create(builder, loc, func, args);
1104 call.setCConv(func.getCConv());
1105 call.setConvergentAttr(func.getConvergentAttr());
1106 call.setNoUnwindAttr(func.getNoUnwindAttr());
1107 call.setWillReturnAttr(func.getWillReturnAttr());
1108 return call;
1109}
1110
1111template <typename BarrierOpTy>
1112class ControlBarrierPattern : public SPIRVToLLVMConversion<BarrierOpTy> {
1113public:
1114 using OpAdaptor = typename SPIRVToLLVMConversion<BarrierOpTy>::OpAdaptor;
1115
1116 using SPIRVToLLVMConversion<BarrierOpTy>::SPIRVToLLVMConversion;
1117
1118 static constexpr StringRef getFuncName();
1119
1120 LogicalResult
1121 matchAndRewrite(BarrierOpTy controlBarrierOp, OpAdaptor adaptor,
1122 ConversionPatternRewriter &rewriter) const override {
1123 constexpr StringRef funcName = getFuncName();
1124 Operation *symbolTable =
1125 controlBarrierOp->template getParentWithTrait<OpTrait::SymbolTable>();
1126
1127 Type i32 = rewriter.getI32Type();
1128
1129 Type voidTy = rewriter.getType<LLVM::LLVMVoidType>();
1130 LLVM::LLVMFuncOp func =
1131 lookupOrCreateSPIRVFn(symbolTable, funcName, {i32, i32, i32}, voidTy);
1132
1133 Location loc = controlBarrierOp->getLoc();
1134 Value execution = LLVM::ConstantOp::create(
1135 rewriter, loc, i32, static_cast<int32_t>(adaptor.getExecutionScope()));
1136 Value memory = LLVM::ConstantOp::create(
1137 rewriter, loc, i32, static_cast<int32_t>(adaptor.getMemoryScope()));
1138 Value semantics = LLVM::ConstantOp::create(
1139 rewriter, loc, i32, static_cast<int32_t>(adaptor.getMemorySemantics()));
1140
1141 auto call = createSPIRVBuiltinCall(loc, rewriter, func,
1142 {execution, memory, semantics});
1143
1144 rewriter.replaceOp(controlBarrierOp, call);
1145 return success();
1146 }
1147};
1148
1149namespace {
1150
1151StringRef getTypeMangling(Type type, bool isSigned) {
1153 .Case([](Float16Type) { return "Dh"; })
1154 .Case([](Float32Type) { return "f"; })
1155 .Case([](Float64Type) { return "d"; })
1156 .Case([isSigned](IntegerType intTy) {
1157 switch (intTy.getWidth()) {
1158 case 1:
1159 return "b";
1160 case 8:
1161 return (isSigned) ? "a" : "c";
1162 case 16:
1163 return (isSigned) ? "s" : "t";
1164 case 32:
1165 return (isSigned) ? "i" : "j";
1166 case 64:
1167 return (isSigned) ? "l" : "m";
1168 default:
1169 llvm_unreachable("Unsupported integer width");
1170 }
1171 })
1172 .DefaultUnreachable("No mangling defined");
1173}
1174
1175template <typename ReduceOp>
1176constexpr StringLiteral getGroupFuncName();
1177
1178template <>
1179constexpr StringLiteral getGroupFuncName<spirv::GroupIAddOp>() {
1180 return "_Z17__spirv_GroupIAddii";
1181}
1182template <>
1183constexpr StringLiteral getGroupFuncName<spirv::GroupFAddOp>() {
1184 return "_Z17__spirv_GroupFAddii";
1185}
1186template <>
1187constexpr StringLiteral getGroupFuncName<spirv::GroupSMinOp>() {
1188 return "_Z17__spirv_GroupSMinii";
1189}
1190template <>
1191constexpr StringLiteral getGroupFuncName<spirv::GroupUMinOp>() {
1192 return "_Z17__spirv_GroupUMinii";
1193}
1194template <>
1195constexpr StringLiteral getGroupFuncName<spirv::GroupFMinOp>() {
1196 return "_Z17__spirv_GroupFMinii";
1197}
1198template <>
1199constexpr StringLiteral getGroupFuncName<spirv::GroupSMaxOp>() {
1200 return "_Z17__spirv_GroupSMaxii";
1201}
1202template <>
1203constexpr StringLiteral getGroupFuncName<spirv::GroupUMaxOp>() {
1204 return "_Z17__spirv_GroupUMaxii";
1205}
1206template <>
1207constexpr StringLiteral getGroupFuncName<spirv::GroupFMaxOp>() {
1208 return "_Z17__spirv_GroupFMaxii";
1209}
1210template <>
1211constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformIAddOp>() {
1212 return "_Z27__spirv_GroupNonUniformIAddii";
1213}
1214template <>
1215constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformFAddOp>() {
1216 return "_Z27__spirv_GroupNonUniformFAddii";
1217}
1218template <>
1219constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformIMulOp>() {
1220 return "_Z27__spirv_GroupNonUniformIMulii";
1221}
1222template <>
1223constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformFMulOp>() {
1224 return "_Z27__spirv_GroupNonUniformFMulii";
1225}
1226template <>
1227constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformSMinOp>() {
1228 return "_Z27__spirv_GroupNonUniformSMinii";
1229}
1230template <>
1231constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformUMinOp>() {
1232 return "_Z27__spirv_GroupNonUniformUMinii";
1233}
1234template <>
1235constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformFMinOp>() {
1236 return "_Z27__spirv_GroupNonUniformFMinii";
1237}
1238template <>
1239constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformSMaxOp>() {
1240 return "_Z27__spirv_GroupNonUniformSMaxii";
1241}
1242template <>
1243constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformUMaxOp>() {
1244 return "_Z27__spirv_GroupNonUniformUMaxii";
1245}
1246template <>
1247constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformFMaxOp>() {
1248 return "_Z27__spirv_GroupNonUniformFMaxii";
1249}
1250template <>
1251constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformBitwiseAndOp>() {
1252 return "_Z33__spirv_GroupNonUniformBitwiseAndii";
1253}
1254template <>
1255constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformBitwiseOrOp>() {
1256 return "_Z32__spirv_GroupNonUniformBitwiseOrii";
1257}
1258template <>
1259constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformBitwiseXorOp>() {
1260 return "_Z33__spirv_GroupNonUniformBitwiseXorii";
1261}
1262template <>
1263constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformLogicalAndOp>() {
1264 return "_Z33__spirv_GroupNonUniformLogicalAndii";
1265}
1266template <>
1267constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformLogicalOrOp>() {
1268 return "_Z32__spirv_GroupNonUniformLogicalOrii";
1269}
1270template <>
1271constexpr StringLiteral getGroupFuncName<spirv::GroupNonUniformLogicalXorOp>() {
1272 return "_Z33__spirv_GroupNonUniformLogicalXorii";
1273}
1274} // namespace
1275
1276template <typename ReduceOp, bool Signed = false, bool NonUniform = false>
1277class GroupReducePattern : public SPIRVToLLVMConversion<ReduceOp> {
1278public:
1279 using SPIRVToLLVMConversion<ReduceOp>::SPIRVToLLVMConversion;
1280
1281 LogicalResult
1282 matchAndRewrite(ReduceOp op, typename ReduceOp::Adaptor adaptor,
1283 ConversionPatternRewriter &rewriter) const override {
1284
1285 Type retTy = op.getResult().getType();
1286 if (!retTy.isIntOrFloat()) {
1287 return failure();
1288 }
1289 SmallString<36> funcName = getGroupFuncName<ReduceOp>();
1290 funcName += getTypeMangling(retTy, false);
1291
1292 Type i32Ty = rewriter.getI32Type();
1293 SmallVector<Type> paramTypes{i32Ty, i32Ty, retTy};
1294 if constexpr (NonUniform) {
1295 if (adaptor.getClusterSize()) {
1296 funcName += "j";
1297 paramTypes.push_back(i32Ty);
1298 }
1299 }
1300
1301 Operation *symbolTable =
1302 op->template getParentWithTrait<OpTrait::SymbolTable>();
1303
1304 LLVM::LLVMFuncOp func =
1305 lookupOrCreateSPIRVFn(symbolTable, funcName, paramTypes, retTy);
1306
1307 Location loc = op.getLoc();
1308 Value scope = LLVM::ConstantOp::create(
1309 rewriter, loc, i32Ty,
1310 static_cast<int32_t>(adaptor.getExecutionScope()));
1311 Value groupOp = LLVM::ConstantOp::create(
1312 rewriter, loc, i32Ty,
1313 static_cast<int32_t>(adaptor.getGroupOperation()));
1314 SmallVector<Value> operands{scope, groupOp};
1315 operands.append(adaptor.getOperands().begin(), adaptor.getOperands().end());
1316
1317 auto call = createSPIRVBuiltinCall(loc, rewriter, func, operands);
1318 rewriter.replaceOp(op, call);
1319 return success();
1320 }
1321};
1322
1323template <>
1324constexpr StringRef
1325ControlBarrierPattern<spirv::ControlBarrierOp>::getFuncName() {
1326 return "_Z22__spirv_ControlBarrieriii";
1327}
1328
1329template <>
1330constexpr StringRef
1331ControlBarrierPattern<spirv::INTELControlBarrierArriveOp>::getFuncName() {
1332 return "_Z33__spirv_ControlBarrierArriveINTELiii";
1333}
1334
1335template <>
1336constexpr StringRef
1337ControlBarrierPattern<spirv::INTELControlBarrierWaitOp>::getFuncName() {
1338 return "_Z31__spirv_ControlBarrierWaitINTELiii";
1339}
1340
1341/// Converts `spirv.mlir.loop` to LLVM dialect. All blocks within selection
1342/// should be reachable for conversion to succeed. The structure of the loop in
1343/// LLVM dialect will be the following:
1344///
1345/// +------------------------------------+
1346/// | <code before spirv.mlir.loop> |
1347/// | llvm.br ^header |
1348/// +------------------------------------+
1349/// |
1350/// +----------------+ |
1351/// | | |
1352/// | V V
1353/// | +------------------------------------+
1354/// | | ^header: |
1355/// | | <header code> |
1356/// | | llvm.cond_br %cond, ^body, ^exit |
1357/// | +------------------------------------+
1358/// | |
1359/// | |----------------------+
1360/// | | |
1361/// | V |
1362/// | +------------------------------------+ |
1363/// | | ^body: | |
1364/// | | <body code> | |
1365/// | | llvm.br ^continue | |
1366/// | +------------------------------------+ |
1367/// | | |
1368/// | V |
1369/// | +------------------------------------+ |
1370/// | | ^continue: | |
1371/// | | <continue code> | |
1372/// | | llvm.br ^header | |
1373/// | +------------------------------------+ |
1374/// | | |
1375/// +---------------+ +----------------------+
1376/// |
1377/// V
1378/// +------------------------------------+
1379/// | ^exit: |
1380/// | llvm.br ^remaining |
1381/// +------------------------------------+
1382/// |
1383/// V
1384/// +------------------------------------+
1385/// | ^remaining: |
1386/// | <code after spirv.mlir.loop> |
1387/// +------------------------------------+
1388///
1389class LoopPattern : public SPIRVToLLVMConversion<spirv::LoopOp> {
1390public:
1391 using SPIRVToLLVMConversion<spirv::LoopOp>::SPIRVToLLVMConversion;
1392
1393 LogicalResult
1394 matchAndRewrite(spirv::LoopOp loopOp, OpAdaptor adaptor,
1395 ConversionPatternRewriter &rewriter) const override {
1396 // There is no support of loop control at the moment.
1397 if (loopOp.getLoopControl() != spirv::LoopControl::None)
1398 return failure();
1399
1400 // `spirv.mlir.loop` with empty region is redundant and should be erased.
1401 if (loopOp.getBody().empty()) {
1402 rewriter.eraseOp(loopOp);
1403 return success();
1404 }
1405
1406 Location loc = loopOp.getLoc();
1407
1408 // Split the current block after `spirv.mlir.loop`. The remaining ops will
1409 // be used in `endBlock`.
1410 Block *currentBlock = rewriter.getBlock();
1411 auto position = Block::iterator(loopOp);
1412 Block *endBlock = rewriter.splitBlock(currentBlock, position);
1413
1414 // Remove entry block and create a branch in the current block going to the
1415 // header block.
1416 Block *entryBlock = loopOp.getEntryBlock();
1417 assert(entryBlock->getOperations().size() == 1);
1418 auto brOp = dyn_cast<spirv::BranchOp>(entryBlock->getOperations().front());
1419 if (!brOp)
1420 return failure();
1421 Block *headerBlock = loopOp.getHeaderBlock();
1422 rewriter.setInsertionPointToEnd(currentBlock);
1423 LLVM::BrOp::create(rewriter, loc, brOp.getBlockArguments(), headerBlock);
1424 rewriter.eraseBlock(entryBlock);
1425
1426 // Branch from merge block to end block.
1427 Block *mergeBlock = loopOp.getMergeBlock();
1428 Operation *terminator = mergeBlock->getTerminator();
1429 ValueRange terminatorOperands = terminator->getOperands();
1430 rewriter.setInsertionPointToEnd(mergeBlock);
1431 LLVM::BrOp::create(rewriter, loc, terminatorOperands, endBlock);
1432
1433 rewriter.inlineRegionBefore(loopOp.getBody(), endBlock);
1434 rewriter.replaceOp(loopOp, endBlock->getArguments());
1435 return success();
1436 }
1437};
1438
1439/// Converts `spirv.mlir.selection` with `spirv.BranchConditional` in its header
1440/// block. All blocks within selection should be reachable for conversion to
1441/// succeed.
1442class SelectionPattern : public SPIRVToLLVMConversion<spirv::SelectionOp> {
1443public:
1444 using SPIRVToLLVMConversion<spirv::SelectionOp>::SPIRVToLLVMConversion;
1445
1446 LogicalResult
1447 matchAndRewrite(spirv::SelectionOp op, OpAdaptor adaptor,
1448 ConversionPatternRewriter &rewriter) const override {
1449 // There is no support for `Flatten` or `DontFlatten` selection control at
1450 // the moment. This are just compiler hints and can be performed during the
1451 // optimization passes.
1452 if (op.getSelectionControl() != spirv::SelectionControl::None)
1453 return failure();
1454
1455 // `spirv.mlir.selection` should have at least two blocks: one selection
1456 // header block and one merge block. If no blocks are present, or control
1457 // flow branches straight to merge block (two blocks are present), the op is
1458 // redundant and it is erased.
1459 if (op.getBody().getBlocks().size() <= 2) {
1460 rewriter.eraseOp(op);
1461 return success();
1462 }
1463
1464 Location loc = op.getLoc();
1465
1466 // Split the current block after `spirv.mlir.selection`. The remaining ops
1467 // will be used in `continueBlock`.
1468 auto *currentBlock = rewriter.getInsertionBlock();
1469 rewriter.setInsertionPointAfter(op);
1470 auto position = rewriter.getInsertionPoint();
1471 auto *continueBlock = rewriter.splitBlock(currentBlock, position);
1472
1473 // Extract conditional branch information from the header block. By SPIR-V
1474 // dialect spec, it should contain `spirv.BranchConditional` or
1475 // `spirv.Switch` op. Note that `spirv.Switch op` is not supported at the
1476 // moment in the SPIR-V dialect. Remove this block when finished.
1477 auto *headerBlock = op.getHeaderBlock();
1478 assert(headerBlock->getOperations().size() == 1);
1479 auto condBrOp = dyn_cast<spirv::BranchConditionalOp>(
1480 headerBlock->getOperations().front());
1481 if (!condBrOp)
1482 return failure();
1483
1484 // Branch from merge block to continue block.
1485 auto *mergeBlock = op.getMergeBlock();
1486 Operation *terminator = mergeBlock->getTerminator();
1487 ValueRange terminatorOperands = terminator->getOperands();
1488 rewriter.setInsertionPointToEnd(mergeBlock);
1489 LLVM::BrOp::create(rewriter, loc, terminatorOperands, continueBlock);
1490
1491 // Link current block to `true` and `false` blocks within the selection.
1492 Block *trueBlock = condBrOp.getTrueBlock();
1493 Block *falseBlock = condBrOp.getFalseBlock();
1494 rewriter.setInsertionPointToEnd(currentBlock);
1495 LLVM::CondBrOp::create(rewriter, loc, condBrOp.getCondition(), trueBlock,
1496 condBrOp.getTrueTargetOperands(), falseBlock,
1497 condBrOp.getFalseTargetOperands());
1498
1499 rewriter.eraseBlock(headerBlock);
1500 rewriter.inlineRegionBefore(op.getBody(), continueBlock);
1501 rewriter.replaceOp(op, continueBlock->getArguments());
1502 return success();
1503 }
1504};
1505
1506/// Converts SPIR-V shift ops to LLVM shift ops. Since LLVM dialect
1507/// puts a restriction on `Shift` and `Base` to have the same bit width,
1508/// `Shift` is zero or sign extended to match this specification. Cases when
1509/// `Shift` bit width > `Base` bit width are considered to be illegal.
1510template <typename SPIRVOp, typename LLVMOp>
1511class ShiftPattern : public SPIRVToLLVMConversion<SPIRVOp> {
1512public:
1513 using SPIRVToLLVMConversion<SPIRVOp>::SPIRVToLLVMConversion;
1514
1515 LogicalResult
1516 matchAndRewrite(SPIRVOp op, typename SPIRVOp::Adaptor adaptor,
1517 ConversionPatternRewriter &rewriter) const override {
1518
1519 auto dstType = this->getTypeConverter()->convertType(op.getType());
1520 if (!dstType)
1521 return rewriter.notifyMatchFailure(op, "type conversion failed");
1522
1523 Type op1Type = op.getOperand1().getType();
1524 Type op2Type = op.getOperand2().getType();
1525
1526 if (op1Type == op2Type) {
1527 rewriter.template replaceOpWithNewOp<LLVMOp>(op, dstType,
1528 adaptor.getOperands());
1529 return success();
1530 }
1531
1532 std::optional<uint64_t> dstTypeWidth =
1534 std::optional<uint64_t> op2TypeWidth =
1536
1537 if (!dstTypeWidth || !op2TypeWidth)
1538 return failure();
1539
1540 Location loc = op.getLoc();
1541 Value extended;
1542 if (op2TypeWidth < dstTypeWidth) {
1543 if (isUnsignedIntegerOrVector(op2Type)) {
1544 extended =
1545 LLVM::ZExtOp::create(rewriter, loc, dstType, adaptor.getOperand2());
1546 } else {
1547 extended =
1548 LLVM::SExtOp::create(rewriter, loc, dstType, adaptor.getOperand2());
1549 }
1550 } else if (op2TypeWidth == dstTypeWidth) {
1551 extended = adaptor.getOperand2();
1552 } else {
1553 return failure();
1554 }
1555
1556 Value result =
1557 LLVMOp::create(rewriter, loc, dstType, adaptor.getOperand1(), extended);
1558 rewriter.replaceOp(op, result);
1559 return success();
1560 }
1561};
1562
1563class TanPattern : public SPIRVToLLVMConversion<spirv::GLTanOp> {
1564public:
1565 using SPIRVToLLVMConversion<spirv::GLTanOp>::SPIRVToLLVMConversion;
1566
1567 LogicalResult
1568 matchAndRewrite(spirv::GLTanOp tanOp, OpAdaptor adaptor,
1569 ConversionPatternRewriter &rewriter) const override {
1570 auto dstType = getTypeConverter()->convertType(tanOp.getType());
1571 if (!dstType)
1572 return rewriter.notifyMatchFailure(tanOp, "type conversion failed");
1573
1574 rewriter.replaceOpWithNewOp<LLVM::TanOp>(tanOp, dstType,
1575 adaptor.getOperands());
1576 return success();
1577 }
1578};
1579
1580class TanhPattern : public SPIRVToLLVMConversion<spirv::GLTanhOp> {
1581public:
1582 using SPIRVToLLVMConversion<spirv::GLTanhOp>::SPIRVToLLVMConversion;
1583
1584 LogicalResult
1585 matchAndRewrite(spirv::GLTanhOp tanhOp, OpAdaptor adaptor,
1586 ConversionPatternRewriter &rewriter) const override {
1587 auto srcType = tanhOp.getType();
1588 auto dstType = getTypeConverter()->convertType(srcType);
1589 if (!dstType)
1590 return rewriter.notifyMatchFailure(tanhOp, "type conversion failed");
1591
1592 rewriter.replaceOpWithNewOp<LLVM::TanhOp>(tanhOp, dstType,
1593 adaptor.getOperands());
1594 return success();
1595 }
1596};
1597
1598// `llvm.intr.abs` requires an `is_int_min_poison` immarg that `spirv.GL.SAbs`
1599// does not carry; default to `false` to preserve SPIR-V's well-defined
1600// behavior on INT_MIN.
1601class SAbsPattern : public SPIRVToLLVMConversion<spirv::GLSAbsOp> {
1602public:
1603 using SPIRVToLLVMConversion<spirv::GLSAbsOp>::SPIRVToLLVMConversion;
1604
1605 LogicalResult
1606 matchAndRewrite(spirv::GLSAbsOp op, OpAdaptor adaptor,
1607 ConversionPatternRewriter &rewriter) const override {
1608 Type dstType = getTypeConverter()->convertType(op.getType());
1609 if (!dstType)
1610 return rewriter.notifyMatchFailure(op, "type conversion failed");
1611
1612 rewriter.replaceOpWithNewOp<LLVM::AbsOp>(op, dstType, adaptor.getOperand(),
1613 /*is_int_min_poison=*/false);
1614 return success();
1615 }
1616};
1617
1618/// Converts `spirv.GL.Fract` to `x - floor(x)`.
1619class FractPattern : public SPIRVToLLVMConversion<spirv::GLFractOp> {
1620public:
1621 using SPIRVToLLVMConversion<spirv::GLFractOp>::SPIRVToLLVMConversion;
1622
1623 LogicalResult
1624 matchAndRewrite(spirv::GLFractOp op, OpAdaptor adaptor,
1625 ConversionPatternRewriter &rewriter) const override {
1626 Type dstType = getTypeConverter()->convertType(op.getType());
1627 if (!dstType)
1628 return rewriter.notifyMatchFailure(op, "type conversion failed");
1629
1630 Location loc = op.getLoc();
1631 Value operand = adaptor.getOperand();
1632 Value floored = LLVM::FFloorOp::create(rewriter, loc, dstType, operand);
1633 rewriter.replaceOpWithNewOp<LLVM::FSubOp>(op, dstType, operand, floored);
1634 return success();
1635 }
1636};
1637
1638/// Converts `spirv.GL.FMix` to `x * (1 - a) + y * a` as specified by
1639/// GL.std.450.
1640class GLFMixPattern : public SPIRVToLLVMConversion<spirv::GLFMixOp> {
1641public:
1642 using SPIRVToLLVMConversion<spirv::GLFMixOp>::SPIRVToLLVMConversion;
1643
1644 LogicalResult
1645 matchAndRewrite(spirv::GLFMixOp op, OpAdaptor adaptor,
1646 ConversionPatternRewriter &rewriter) const override {
1647 Type dstType = getTypeConverter()->convertType(op.getType());
1648 if (!dstType)
1649 return rewriter.notifyMatchFailure(op, "type conversion failed");
1650
1651 Location loc = op.getLoc();
1652 Value x = adaptor.getX();
1653 Value y = adaptor.getY();
1654 Value a = adaptor.getA();
1655 Value one = createFPConstant(loc, op.getType(), dstType, rewriter, 1.0);
1656 Value oneMinusA = LLVM::FSubOp::create(rewriter, loc, dstType, one, a);
1657 Value lhs = LLVM::FMulOp::create(rewriter, loc, dstType, x, oneMinusA);
1658 Value rhs = LLVM::FMulOp::create(rewriter, loc, dstType, y, a);
1659 rewriter.replaceOpWithNewOp<LLVM::FAddOp>(op, dstType, lhs, rhs);
1660 return success();
1661 }
1662};
1663
1664/// Converts `spirv.CL.mix` to `fma(a, y - x, x)`. The OpenCL spec defines
1665/// mix as `x + (y - x) * a` and explicitly permits FMA contractions.
1666class CLMixPattern : public SPIRVToLLVMConversion<spirv::CLMixOp> {
1667public:
1668 using SPIRVToLLVMConversion<spirv::CLMixOp>::SPIRVToLLVMConversion;
1669
1670 LogicalResult
1671 matchAndRewrite(spirv::CLMixOp op, OpAdaptor adaptor,
1672 ConversionPatternRewriter &rewriter) const override {
1673 Type dstType = getTypeConverter()->convertType(op.getType());
1674 if (!dstType)
1675 return rewriter.notifyMatchFailure(op, "type conversion failed");
1676
1677 Location loc = op.getLoc();
1678 Value x = adaptor.getX();
1679 Value y = adaptor.getY();
1680 Value a = adaptor.getZ();
1681 Value diff = LLVM::FSubOp::create(rewriter, loc, dstType, y, x);
1682 rewriter.replaceOpWithNewOp<LLVM::FMAOp>(op, dstType, a, diff, x);
1683 return success();
1684 }
1685};
1686
1687// Converts spirv.GL.Radians (scale = pi/180) and spirv.GL.Degrees
1688// (scale = 180/pi) by multiplying the operand by a compile-time constant.
1689template <typename SPIRVOp>
1690class ScalePattern : public SPIRVToLLVMConversion<SPIRVOp> {
1691public:
1692 template <typename... Args>
1693 ScalePattern(double scale, Args &&...args)
1694 : SPIRVToLLVMConversion<SPIRVOp>(std::forward<Args>(args)...),
1695 scale(scale) {}
1696
1697 LogicalResult
1698 matchAndRewrite(SPIRVOp op, typename SPIRVOp::Adaptor adaptor,
1699 ConversionPatternRewriter &rewriter) const override {
1700 Type srcType = op.getType();
1701 Type dstType = this->getTypeConverter()->convertType(srcType);
1702 if (!dstType)
1703 return rewriter.notifyMatchFailure(op, "type conversion failed");
1704
1705 Location loc = op.getLoc();
1706 Value factor = createFPConstant(loc, srcType, dstType, rewriter, scale);
1707 rewriter.replaceOpWithNewOp<LLVM::FMulOp>(op, dstType, adaptor.getOperand(),
1708 factor);
1709 return success();
1710 }
1711
1712private:
1713 double scale;
1714};
1715
1716class VariablePattern : public SPIRVToLLVMConversion<spirv::VariableOp> {
1717public:
1718 using SPIRVToLLVMConversion<spirv::VariableOp>::SPIRVToLLVMConversion;
1719
1720 LogicalResult
1721 matchAndRewrite(spirv::VariableOp varOp, OpAdaptor adaptor,
1722 ConversionPatternRewriter &rewriter) const override {
1723 auto srcType = varOp.getType();
1724 // Initialization is supported for scalars and vectors only.
1725 auto pointerTo = cast<spirv::PointerType>(srcType).getPointeeType();
1726 auto init = varOp.getInitializer();
1727 if (init && !pointerTo.isIntOrFloat() && !isa<VectorType>(pointerTo))
1728 return failure();
1729
1730 auto dstType = getTypeConverter()->convertType(srcType);
1731 if (!dstType)
1732 return rewriter.notifyMatchFailure(varOp, "type conversion failed");
1733
1734 Location loc = varOp.getLoc();
1735 Value size = createI32ConstantOf(loc, rewriter, 1);
1736 if (!init) {
1737 auto elementType = getTypeConverter()->convertType(pointerTo);
1738 if (!elementType)
1739 return rewriter.notifyMatchFailure(varOp, "type conversion failed");
1740 rewriter.replaceOpWithNewOp<LLVM::AllocaOp>(varOp, dstType, elementType,
1741 size);
1742 return success();
1743 }
1744 auto elementType = getTypeConverter()->convertType(pointerTo);
1745 if (!elementType)
1746 return rewriter.notifyMatchFailure(varOp, "type conversion failed");
1747 Value allocated =
1748 LLVM::AllocaOp::create(rewriter, loc, dstType, elementType, size);
1749 LLVM::StoreOp::create(rewriter, loc, adaptor.getInitializer(), allocated);
1750 rewriter.replaceOp(varOp, allocated);
1751 return success();
1752 }
1753};
1754
1755//===----------------------------------------------------------------------===//
1756// BitcastOp conversion
1757//===----------------------------------------------------------------------===//
1758
1759class BitcastConversionPattern
1760 : public SPIRVToLLVMConversion<spirv::BitcastOp> {
1761public:
1762 using SPIRVToLLVMConversion<spirv::BitcastOp>::SPIRVToLLVMConversion;
1763
1764 LogicalResult
1765 matchAndRewrite(spirv::BitcastOp bitcastOp, OpAdaptor adaptor,
1766 ConversionPatternRewriter &rewriter) const override {
1767 auto dstType = getTypeConverter()->convertType(bitcastOp.getType());
1768 if (!dstType)
1769 return rewriter.notifyMatchFailure(bitcastOp, "type conversion failed");
1770
1771 // LLVM's opaque pointers do not require bitcasts.
1772 if (isa<LLVM::LLVMPointerType>(dstType)) {
1773 rewriter.replaceOp(bitcastOp, adaptor.getOperand());
1774 return success();
1775 }
1776
1777 rewriter.replaceOpWithNewOp<LLVM::BitcastOp>(
1778 bitcastOp, dstType, adaptor.getOperands(), bitcastOp->getAttrs());
1779 return success();
1780 }
1781};
1782
1783//===----------------------------------------------------------------------===//
1784// FuncOp conversion
1785//===----------------------------------------------------------------------===//
1786
1787class FuncConversionPattern : public SPIRVToLLVMConversion<spirv::FuncOp> {
1788public:
1789 using SPIRVToLLVMConversion<spirv::FuncOp>::SPIRVToLLVMConversion;
1790
1791 LogicalResult
1792 matchAndRewrite(spirv::FuncOp funcOp, OpAdaptor adaptor,
1793 ConversionPatternRewriter &rewriter) const override {
1794
1795 // Convert function signature. At the moment LLVMType converter is enough
1796 // for currently supported types.
1797 auto funcType = funcOp.getFunctionType();
1798 TypeConverter::SignatureConversion signatureConverter(
1799 funcType.getNumInputs());
1800 auto llvmType = static_cast<const LLVMTypeConverter *>(getTypeConverter())
1801 ->convertFunctionSignature(
1802 funcType, /*isVariadic=*/false,
1803 /*useBarePtrCallConv=*/false, signatureConverter);
1804 if (!llvmType)
1805 return failure();
1806
1807 // Create a new `LLVMFuncOp`
1808 Location loc = funcOp.getLoc();
1809 StringRef name = funcOp.getName();
1810 auto newFuncOp = LLVM::LLVMFuncOp::create(rewriter, loc, name, llvmType);
1811
1812 // Convert SPIR-V Function Control to equivalent LLVM function attribute
1813 MLIRContext *context = funcOp.getContext();
1814 switch (funcOp.getFunctionControl()) {
1815 case spirv::FunctionControl::Inline:
1816 newFuncOp.setAlwaysInline(true);
1817 break;
1818 case spirv::FunctionControl::DontInline:
1819 newFuncOp.setNoInline(true);
1820 break;
1821
1822#define DISPATCH(functionControl, llvmAttr) \
1823 case functionControl: \
1824 newFuncOp->setAttr("passthrough", ArrayAttr::get(context, {llvmAttr})); \
1825 break;
1826
1827 DISPATCH(spirv::FunctionControl::Pure,
1828 StringAttr::get(context, "readonly"));
1829 DISPATCH(spirv::FunctionControl::Const,
1830 StringAttr::get(context, "readnone"));
1831
1832#undef DISPATCH
1833
1834 // Default: if `spirv::FunctionControl::None`, then no attributes are
1835 // needed.
1836 default:
1837 break;
1838 }
1839
1840 rewriter.inlineRegionBefore(funcOp.getBody(), newFuncOp.getBody(),
1841 newFuncOp.end());
1842 if (failed(rewriter.convertRegionTypes(
1843 &newFuncOp.getBody(), *getTypeConverter(), &signatureConverter))) {
1844 return failure();
1845 }
1846 rewriter.eraseOp(funcOp);
1847 return success();
1848 }
1849};
1850
1851//===----------------------------------------------------------------------===//
1852// ModuleOp conversion
1853//===----------------------------------------------------------------------===//
1854
1855class ModuleConversionPattern : public SPIRVToLLVMConversion<spirv::ModuleOp> {
1856public:
1857 using SPIRVToLLVMConversion<spirv::ModuleOp>::SPIRVToLLVMConversion;
1858
1859 LogicalResult
1860 matchAndRewrite(spirv::ModuleOp spvModuleOp, OpAdaptor adaptor,
1861 ConversionPatternRewriter &rewriter) const override {
1862
1863 auto newModuleOp =
1864 ModuleOp::create(rewriter, spvModuleOp.getLoc(), spvModuleOp.getName());
1865 rewriter.inlineRegionBefore(spvModuleOp.getRegion(), newModuleOp.getBody());
1866
1867 // Remove the terminator block that was automatically added by builder
1868 rewriter.eraseBlock(&newModuleOp.getBodyRegion().back());
1869 rewriter.eraseOp(spvModuleOp);
1870 return success();
1871 }
1872};
1873
1874//===----------------------------------------------------------------------===//
1875// VectorShuffleOp conversion
1876//===----------------------------------------------------------------------===//
1877
1878class VectorShufflePattern
1879 : public SPIRVToLLVMConversion<spirv::VectorShuffleOp> {
1880public:
1881 using SPIRVToLLVMConversion<spirv::VectorShuffleOp>::SPIRVToLLVMConversion;
1882 LogicalResult
1883 matchAndRewrite(spirv::VectorShuffleOp op, OpAdaptor adaptor,
1884 ConversionPatternRewriter &rewriter) const override {
1885 Location loc = op.getLoc();
1886 auto components = adaptor.getComponents();
1887 auto vector1 = adaptor.getVector1();
1888 auto vector2 = adaptor.getVector2();
1889 int vector1Size = cast<VectorType>(vector1.getType()).getNumElements();
1890 int vector2Size = cast<VectorType>(vector2.getType()).getNumElements();
1891 if (vector1Size == vector2Size) {
1892 rewriter.replaceOpWithNewOp<LLVM::ShuffleVectorOp>(
1893 op, vector1, vector2,
1894 LLVM::convertArrayToIndices<int32_t>(components));
1895 return success();
1896 }
1897
1898 auto dstType = getTypeConverter()->convertType(op.getType());
1899 if (!dstType)
1900 return rewriter.notifyMatchFailure(op, "type conversion failed");
1901 auto scalarType = cast<VectorType>(dstType).getElementType();
1902 auto componentsArray = components.getValue();
1903 auto *context = rewriter.getContext();
1904 auto llvmI32Type = IntegerType::get(context, 32);
1905 Value targetOp = LLVM::PoisonOp::create(rewriter, loc, dstType);
1906 for (unsigned i = 0; i < componentsArray.size(); i++) {
1907 if (!isa<IntegerAttr>(componentsArray[i]))
1908 return op.emitError("unable to support non-constant component");
1909
1910 int indexVal = cast<IntegerAttr>(componentsArray[i]).getInt();
1911 if (indexVal == -1)
1912 continue;
1913
1914 int offsetVal = 0;
1915 Value baseVector = vector1;
1916 if (indexVal >= vector1Size) {
1917 offsetVal = vector1Size;
1918 baseVector = vector2;
1919 }
1920
1921 Value dstIndex = LLVM::ConstantOp::create(
1922 rewriter, loc, llvmI32Type,
1923 rewriter.getIntegerAttr(rewriter.getI32Type(), i));
1924 Value index = LLVM::ConstantOp::create(
1925 rewriter, loc, llvmI32Type,
1926 rewriter.getIntegerAttr(rewriter.getI32Type(), indexVal - offsetVal));
1927
1928 auto extractOp = LLVM::ExtractElementOp::create(rewriter, loc, scalarType,
1929 baseVector, index);
1930 targetOp = LLVM::InsertElementOp::create(rewriter, loc, dstType, targetOp,
1931 extractOp, dstIndex);
1932 }
1933 rewriter.replaceOp(op, targetOp);
1934 return success();
1935 }
1936};
1937} // namespace
1938
1939//===----------------------------------------------------------------------===//
1940// Pattern population
1941//===----------------------------------------------------------------------===//
1942
1944 spirv::ClientAPI clientAPI) {
1945 typeConverter.addConversion([&](spirv::ArrayType type) {
1946 return convertArrayType(type, typeConverter);
1947 });
1948 typeConverter.addConversion([&, clientAPI](spirv::PointerType type) {
1949 return convertPointerType(type, typeConverter, clientAPI);
1950 });
1951 typeConverter.addConversion([&](spirv::RuntimeArrayType type) {
1952 return convertRuntimeArrayType(type, typeConverter);
1953 });
1954 typeConverter.addConversion([&](spirv::StructType type) {
1955 return convertStructType(type, typeConverter);
1956 });
1957}
1958
1960 const LLVMTypeConverter &typeConverter, RewritePatternSet &patterns,
1961 spirv::ClientAPI clientAPI) {
1962 patterns.add<
1963 // Arithmetic ops
1964 DirectConversionPattern<spirv::IAddOp, LLVM::AddOp>,
1965 DirectConversionPattern<spirv::IMulOp, LLVM::MulOp>,
1966 DirectConversionPattern<spirv::ISubOp, LLVM::SubOp>,
1967 DirectConversionPattern<spirv::FAddOp, LLVM::FAddOp>,
1968 DirectConversionPattern<spirv::FDivOp, LLVM::FDivOp>,
1969 DirectConversionPattern<spirv::FMulOp, LLVM::FMulOp>,
1970 DirectConversionPattern<spirv::FNegateOp, LLVM::FNegOp>,
1971 DirectConversionPattern<spirv::FRemOp, LLVM::FRemOp>,
1972 DirectConversionPattern<spirv::FSubOp, LLVM::FSubOp>,
1973 DirectConversionPattern<spirv::SDivOp, LLVM::SDivOp>,
1974 DirectConversionPattern<spirv::SRemOp, LLVM::SRemOp>,
1975 DirectConversionPattern<spirv::UDivOp, LLVM::UDivOp>,
1976 DirectConversionPattern<spirv::UModOp, LLVM::URemOp>, SNegatePattern,
1977
1978 // Bitwise ops
1979 BitFieldInsertPattern, BitFieldUExtractPattern, BitFieldSExtractPattern,
1980 DirectConversionPattern<spirv::BitCountOp, LLVM::CtPopOp>,
1981 DirectConversionPattern<spirv::BitReverseOp, LLVM::BitReverseOp>,
1982 DirectConversionPattern<spirv::BitwiseAndOp, LLVM::AndOp>,
1983 DirectConversionPattern<spirv::BitwiseOrOp, LLVM::OrOp>,
1984 DirectConversionPattern<spirv::BitwiseXorOp, LLVM::XOrOp>,
1985 NotPattern<spirv::NotOp>,
1986
1987 // Cast ops
1988 BitcastConversionPattern,
1989 DirectConversionPattern<spirv::ConvertFToSOp, LLVM::FPToSIOp>,
1990 DirectConversionPattern<spirv::ConvertFToUOp, LLVM::FPToUIOp>,
1991 DirectConversionPattern<spirv::ConvertSToFOp, LLVM::SIToFPOp>,
1992 DirectConversionPattern<spirv::ConvertUToFOp, LLVM::UIToFPOp>,
1993 IndirectCastPattern<spirv::FConvertOp, LLVM::FPExtOp, LLVM::FPTruncOp>,
1994 IndirectCastPattern<spirv::SConvertOp, LLVM::SExtOp, LLVM::TruncOp>,
1995 IndirectCastPattern<spirv::UConvertOp, LLVM::ZExtOp, LLVM::TruncOp>,
1996 DirectConversionPattern<spirv::ConvertPtrToUOp, LLVM::PtrToIntOp>,
1997 DirectConversionPattern<spirv::ConvertUToPtrOp, LLVM::IntToPtrOp>,
1998 DirectConversionPattern<spirv::PtrCastToGenericOp, LLVM::AddrSpaceCastOp>,
1999 DirectConversionPattern<spirv::GenericCastToPtrOp, LLVM::AddrSpaceCastOp>,
2000 DirectConversionPattern<spirv::GenericCastToPtrExplicitOp,
2001 LLVM::AddrSpaceCastOp>,
2002
2003 // Comparison ops
2004 IComparePattern<spirv::IEqualOp, LLVM::ICmpPredicate::eq>,
2005 IComparePattern<spirv::INotEqualOp, LLVM::ICmpPredicate::ne>,
2006 FComparePattern<spirv::FOrdEqualOp, LLVM::FCmpPredicate::oeq>,
2007 FComparePattern<spirv::FOrdGreaterThanOp, LLVM::FCmpPredicate::ogt>,
2008 FComparePattern<spirv::FOrdGreaterThanEqualOp, LLVM::FCmpPredicate::oge>,
2009 FComparePattern<spirv::FOrdLessThanEqualOp, LLVM::FCmpPredicate::ole>,
2010 FComparePattern<spirv::FOrdLessThanOp, LLVM::FCmpPredicate::olt>,
2011 FComparePattern<spirv::FOrdNotEqualOp, LLVM::FCmpPredicate::one>,
2012 FComparePattern<spirv::FUnordEqualOp, LLVM::FCmpPredicate::ueq>,
2013 FComparePattern<spirv::FUnordGreaterThanOp, LLVM::FCmpPredicate::ugt>,
2014 FComparePattern<spirv::FUnordGreaterThanEqualOp,
2015 LLVM::FCmpPredicate::uge>,
2016 FComparePattern<spirv::FUnordLessThanEqualOp, LLVM::FCmpPredicate::ule>,
2017 FComparePattern<spirv::FUnordLessThanOp, LLVM::FCmpPredicate::ult>,
2018 FComparePattern<spirv::FUnordNotEqualOp, LLVM::FCmpPredicate::une>,
2019 FComparePattern<spirv::OrderedOp, LLVM::FCmpPredicate::ord>,
2020 FComparePattern<spirv::UnorderedOp, LLVM::FCmpPredicate::uno>,
2021 IComparePattern<spirv::SGreaterThanOp, LLVM::ICmpPredicate::sgt>,
2022 IComparePattern<spirv::SGreaterThanEqualOp, LLVM::ICmpPredicate::sge>,
2023 IComparePattern<spirv::SLessThanEqualOp, LLVM::ICmpPredicate::sle>,
2024 IComparePattern<spirv::SLessThanOp, LLVM::ICmpPredicate::slt>,
2025 IComparePattern<spirv::UGreaterThanOp, LLVM::ICmpPredicate::ugt>,
2026 IComparePattern<spirv::UGreaterThanEqualOp, LLVM::ICmpPredicate::uge>,
2027 IComparePattern<spirv::ULessThanEqualOp, LLVM::ICmpPredicate::ule>,
2028 IComparePattern<spirv::ULessThanOp, LLVM::ICmpPredicate::ult>,
2029
2030 // Constant op
2031 ConstantScalarAndVectorPattern,
2032
2033 // Control Flow ops
2034 BranchConversionPattern, BranchConditionalConversionPattern,
2035 FunctionCallPattern, LoopPattern, SelectionPattern,
2036 ErasePattern<spirv::MergeOp>,
2037
2038 // Entry points and execution mode are handled separately.
2039 ErasePattern<spirv::EntryPointOp>, ExecutionModePattern,
2040
2041 // GLSL extended instruction set ops
2042 DirectConversionPattern<spirv::GLCeilOp, LLVM::FCeilOp>,
2043 DirectConversionPattern<spirv::GLCosOp, LLVM::CosOp>,
2044 DirectConversionPattern<spirv::GLExpOp, LLVM::ExpOp>,
2045 DirectConversionPattern<spirv::GLExp2Op, LLVM::Exp2Op>,
2046 DirectConversionPattern<spirv::GLFAbsOp, LLVM::FAbsOp>,
2047 DirectConversionPattern<spirv::GLFloorOp, LLVM::FFloorOp>,
2048 DirectConversionPattern<spirv::GLFmaOp, LLVM::FMAOp>,
2049 ClampPattern<spirv::GLFClampOp, LLVM::MinNumOp, LLVM::MaxNumOp>,
2050 ClampPattern<spirv::GLSClampOp, LLVM::SMinOp, LLVM::SMaxOp>,
2051 ClampPattern<spirv::GLUClampOp, LLVM::UMinOp, LLVM::UMaxOp>,
2052 DirectConversionPattern<spirv::GLFMaxOp, LLVM::MaxNumOp>,
2053 DirectConversionPattern<spirv::GLFMinOp, LLVM::MinNumOp>,
2054 DirectConversionPattern<spirv::GLNMaxOp, LLVM::MaxNumOp>,
2055 DirectConversionPattern<spirv::GLNMinOp, LLVM::MinNumOp>,
2056 DirectConversionPattern<spirv::GLLogOp, LLVM::LogOp>,
2057 DirectConversionPattern<spirv::GLLog2Op, LLVM::Log2Op>,
2058 DirectConversionPattern<spirv::GLPowOp, LLVM::PowOp>,
2059 DirectConversionPattern<spirv::GLRoundOp, LLVM::RoundOp>,
2060 DirectConversionPattern<spirv::GLRoundEvenOp, LLVM::RoundEvenOp>,
2061 DirectConversionPattern<spirv::GLSinOp, LLVM::SinOp>,
2062 DirectConversionPattern<spirv::GLSinhOp, LLVM::SinhOp>,
2063 DirectConversionPattern<spirv::GLCoshOp, LLVM::CoshOp>,
2064 DirectConversionPattern<spirv::GLSMaxOp, LLVM::SMaxOp>,
2065 DirectConversionPattern<spirv::GLSMinOp, LLVM::SMinOp>,
2066 DirectConversionPattern<spirv::GLSqrtOp, LLVM::SqrtOp>,
2067 DirectConversionPattern<spirv::GLUMaxOp, LLVM::UMaxOp>,
2068 DirectConversionPattern<spirv::GLUMinOp, LLVM::UMinOp>,
2069 DirectConversionPattern<spirv::GLTruncOp, LLVM::FTruncOp>,
2070 DirectConversionPattern<spirv::GLAsinOp, LLVM::ASinOp>,
2071 DirectConversionPattern<spirv::GLAcosOp, LLVM::ACosOp>,
2072 DirectConversionPattern<spirv::GLAtanOp, LLVM::ATanOp>,
2073 InverseSqrtPattern, SAbsPattern, TanPattern, TanhPattern, FractPattern,
2074 GLFMixPattern,
2075
2076 // OpenCL extended instruction set ops
2077 DirectConversionPattern<spirv::CLCeilOp, LLVM::FCeilOp>,
2078 DirectConversionPattern<spirv::CLCosOp, LLVM::CosOp>,
2079 DirectConversionPattern<spirv::CLExpOp, LLVM::ExpOp>,
2080 DirectConversionPattern<spirv::CLExp2Op, LLVM::Exp2Op>,
2081 DirectConversionPattern<spirv::CLExp10Op, LLVM::Exp10Op>,
2082 DirectConversionPattern<spirv::CLFAbsOp, LLVM::FAbsOp>,
2083 DirectConversionPattern<spirv::CLFloorOp, LLVM::FFloorOp>,
2084 DirectConversionPattern<spirv::CLFmaOp, LLVM::FMAOp>,
2085 DirectConversionPattern<spirv::CLFMaxOp, LLVM::MaxNumOp>,
2086 DirectConversionPattern<spirv::CLFMinOp, LLVM::MinNumOp>,
2087 DirectConversionPattern<spirv::CLLogOp, LLVM::LogOp>,
2088 DirectConversionPattern<spirv::CLLog2Op, LLVM::Log2Op>,
2089 DirectConversionPattern<spirv::CLLog10Op, LLVM::Log10Op>,
2090 DirectConversionPattern<spirv::CLPowOp, LLVM::PowOp>,
2091 DirectConversionPattern<spirv::CLRintOp, LLVM::RintOp>,
2092 DirectConversionPattern<spirv::CLRoundOp, LLVM::RoundOp>,
2093 DirectConversionPattern<spirv::CLSinOp, LLVM::SinOp>,
2094 DirectConversionPattern<spirv::CLSinhOp, LLVM::SinhOp>,
2095 DirectConversionPattern<spirv::CLCoshOp, LLVM::CoshOp>,
2096 DirectConversionPattern<spirv::CLTanOp, LLVM::TanOp>,
2097 DirectConversionPattern<spirv::CLTanhOp, LLVM::TanhOp>,
2098 DirectConversionPattern<spirv::CLAsinOp, LLVM::ASinOp>,
2099 DirectConversionPattern<spirv::CLAcosOp, LLVM::ACosOp>,
2100 DirectConversionPattern<spirv::CLAtanOp, LLVM::ATanOp>,
2101 DirectConversionPattern<spirv::CLAtan2Op, LLVM::ATan2Op>,
2102 DirectConversionPattern<spirv::CLSqrtOp, LLVM::SqrtOp>,
2103 DirectConversionPattern<spirv::CLTruncOp, LLVM::FTruncOp>,
2104 DirectConversionPattern<spirv::CLCopysignOp, LLVM::CopySignOp>,
2105 DirectConversionPattern<spirv::CLFmodOp, LLVM::FRemOp>,
2106 DirectConversionPattern<spirv::CLSMaxOp, LLVM::SMaxOp>,
2107 DirectConversionPattern<spirv::CLSMinOp, LLVM::SMinOp>,
2108 DirectConversionPattern<spirv::CLUMaxOp, LLVM::UMaxOp>,
2109 DirectConversionPattern<spirv::CLUMinOp, LLVM::UMinOp>, CLMixPattern,
2110
2111 // Logical ops
2112 DirectConversionPattern<spirv::LogicalAndOp, LLVM::AndOp>,
2113 DirectConversionPattern<spirv::LogicalOrOp, LLVM::OrOp>,
2114 IComparePattern<spirv::LogicalEqualOp, LLVM::ICmpPredicate::eq>,
2115 IComparePattern<spirv::LogicalNotEqualOp, LLVM::ICmpPredicate::ne>,
2116 NotPattern<spirv::LogicalNotOp>,
2117
2118 // Memory ops
2119 AccessChainPattern, AddressOfPattern, LoadStorePattern<spirv::LoadOp>,
2120 LoadStorePattern<spirv::StoreOp>, VariablePattern,
2121
2122 // Miscellaneous ops
2123 CompositeExtractPattern, CompositeInsertPattern,
2124 DirectConversionPattern<spirv::SelectOp, LLVM::SelectOp>,
2125 DirectConversionPattern<spirv::UndefOp, LLVM::UndefOp>,
2126 VectorShufflePattern,
2127
2128 // Shift ops
2129 ShiftPattern<spirv::ShiftRightArithmeticOp, LLVM::AShrOp>,
2130 ShiftPattern<spirv::ShiftRightLogicalOp, LLVM::LShrOp>,
2131 ShiftPattern<spirv::ShiftLeftLogicalOp, LLVM::ShlOp>,
2132
2133 // Return ops
2134 ReturnPattern, ReturnValuePattern,
2135
2136 // Unreachable op
2137 UnreachablePattern,
2138
2139 // Barrier ops
2140 ControlBarrierPattern<spirv::ControlBarrierOp>,
2141 ControlBarrierPattern<spirv::INTELControlBarrierArriveOp>,
2142 ControlBarrierPattern<spirv::INTELControlBarrierWaitOp>,
2143
2144 // Group reduction operations
2145 GroupReducePattern<spirv::GroupIAddOp>,
2146 GroupReducePattern<spirv::GroupFAddOp>,
2147 GroupReducePattern<spirv::GroupFMinOp>,
2148 GroupReducePattern<spirv::GroupUMinOp>,
2149 GroupReducePattern<spirv::GroupSMinOp, /*Signed=*/true>,
2150 GroupReducePattern<spirv::GroupFMaxOp>,
2151 GroupReducePattern<spirv::GroupUMaxOp>,
2152 GroupReducePattern<spirv::GroupSMaxOp, /*Signed=*/true>,
2153 GroupReducePattern<spirv::GroupNonUniformIAddOp, /*Signed=*/false,
2154 /*NonUniform=*/true>,
2155 GroupReducePattern<spirv::GroupNonUniformFAddOp, /*Signed=*/false,
2156 /*NonUniform=*/true>,
2157 GroupReducePattern<spirv::GroupNonUniformIMulOp, /*Signed=*/false,
2158 /*NonUniform=*/true>,
2159 GroupReducePattern<spirv::GroupNonUniformFMulOp, /*Signed=*/false,
2160 /*NonUniform=*/true>,
2161 GroupReducePattern<spirv::GroupNonUniformSMinOp, /*Signed=*/true,
2162 /*NonUniform=*/true>,
2163 GroupReducePattern<spirv::GroupNonUniformUMinOp, /*Signed=*/false,
2164 /*NonUniform=*/true>,
2165 GroupReducePattern<spirv::GroupNonUniformFMinOp, /*Signed=*/false,
2166 /*NonUniform=*/true>,
2167 GroupReducePattern<spirv::GroupNonUniformSMaxOp, /*Signed=*/true,
2168 /*NonUniform=*/true>,
2169 GroupReducePattern<spirv::GroupNonUniformUMaxOp, /*Signed=*/false,
2170 /*NonUniform=*/true>,
2171 GroupReducePattern<spirv::GroupNonUniformFMaxOp, /*Signed=*/false,
2172 /*NonUniform=*/true>,
2173 GroupReducePattern<spirv::GroupNonUniformBitwiseAndOp, /*Signed=*/false,
2174 /*NonUniform=*/true>,
2175 GroupReducePattern<spirv::GroupNonUniformBitwiseOrOp, /*Signed=*/false,
2176 /*NonUniform=*/true>,
2177 GroupReducePattern<spirv::GroupNonUniformBitwiseXorOp, /*Signed=*/false,
2178 /*NonUniform=*/true>,
2179 GroupReducePattern<spirv::GroupNonUniformLogicalAndOp, /*Signed=*/false,
2180 /*NonUniform=*/true>,
2181 GroupReducePattern<spirv::GroupNonUniformLogicalOrOp, /*Signed=*/false,
2182 /*NonUniform=*/true>,
2183 GroupReducePattern<spirv::GroupNonUniformLogicalXorOp, /*Signed=*/false,
2184 /*NonUniform=*/true>>(patterns.getContext(),
2185 typeConverter);
2186
2187 patterns.add<GlobalVariablePattern>(clientAPI, patterns.getContext(),
2188 typeConverter);
2189 // pi / 180
2190 patterns.add<ScalePattern<spirv::GLRadiansOp>>(
2191 0.017453292519943295, patterns.getContext(), typeConverter);
2192 // 180 / pi
2193 patterns.add<ScalePattern<spirv::GLDegreesOp>>(
2194 57.29577951308232, patterns.getContext(), typeConverter);
2195}
2196
2198 const LLVMTypeConverter &typeConverter, RewritePatternSet &patterns) {
2199 patterns.add<FuncConversionPattern>(patterns.getContext(), typeConverter);
2200}
2201
2203 const LLVMTypeConverter &typeConverter, RewritePatternSet &patterns) {
2204 patterns.add<ModuleConversionPattern>(patterns.getContext(), typeConverter);
2205}
2206
2207//===----------------------------------------------------------------------===//
2208// Pre-conversion hooks
2209//===----------------------------------------------------------------------===//
2210
2211/// Hook for descriptor set and binding number encoding.
2212static constexpr StringRef kBinding = "binding";
2213static constexpr StringRef kDescriptorSet = "descriptor_set";
2214void mlir::encodeBindAttribute(ModuleOp module) {
2215 auto spvModules = module.getOps<spirv::ModuleOp>();
2216 for (auto spvModule : spvModules) {
2217 spvModule.walk([&](spirv::GlobalVariableOp op) {
2218 IntegerAttr descriptorSet =
2219 op->getAttrOfType<IntegerAttr>(kDescriptorSet);
2220 IntegerAttr binding = op->getAttrOfType<IntegerAttr>(kBinding);
2221 // For every global variable in the module, get the ones with descriptor
2222 // set and binding numbers.
2223 if (descriptorSet && binding) {
2224 // Encode these numbers into the variable's symbolic name. If the
2225 // SPIR-V module has a name, add it at the beginning.
2226 auto moduleAndName =
2227 spvModule.getName().has_value()
2228 ? spvModule.getName()->str() + "_" + op.getSymName().str()
2229 : op.getSymName().str();
2230 std::string name =
2231 llvm::formatv("{0}_descriptor_set{1}_binding{2}", moduleAndName,
2232 std::to_string(descriptorSet.getInt()),
2233 std::to_string(binding.getInt()));
2234 auto nameAttr = StringAttr::get(op->getContext(), name);
2235
2236 // Replace all symbol uses and set the new symbol name. Finally, remove
2237 // descriptor set and binding attributes.
2238 if (failed(SymbolTable::replaceAllSymbolUses(op, nameAttr, spvModule)))
2239 op.emitError("unable to replace all symbol uses for ") << name;
2240 SymbolTable::setSymbolName(op, nameAttr);
2241 op->removeAttr(kDescriptorSet);
2242 op->removeAttr(kBinding);
2243 }
2244 });
2245 }
2246}
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.
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:233
FloatAttr getFloatAttr(Type type, double value)
Definition Builders.cpp:259
IntegerType getI32Type()
Definition Builders.cpp:67
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:209
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.