MLIR 24.0.0git
GPUToSPIRV.cpp
Go to the documentation of this file.
1//===- GPUToSPIRV.cpp - GPU to SPIR-V 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 GPU dialect to SPIR-V dialect.
10//
11//===----------------------------------------------------------------------===//
12
21#include "mlir/IR/Matchers.h"
23#include <optional>
24
25using namespace mlir;
26
27static constexpr const char kSPIRVModule[] = "__spv__";
28
29namespace {
30/// Pattern lowering GPU block/thread size/id to loading SPIR-V invocation
31/// builtin variables.
32template <typename SourceOp, spirv::BuiltIn builtin>
33class LaunchConfigConversion : public OpConversionPattern<SourceOp> {
34public:
35 using OpConversionPattern<SourceOp>::OpConversionPattern;
36
37 LogicalResult
38 matchAndRewrite(SourceOp op, typename SourceOp::Adaptor adaptor,
39 ConversionPatternRewriter &rewriter) const override;
40};
41
42/// Pattern lowering subgroup size/id to loading SPIR-V invocation
43/// builtin variables.
44template <typename SourceOp, spirv::BuiltIn builtin>
45class SingleDimLaunchConfigConversion : public OpConversionPattern<SourceOp> {
46public:
47 using OpConversionPattern<SourceOp>::OpConversionPattern;
48
49 LogicalResult
50 matchAndRewrite(SourceOp op, typename SourceOp::Adaptor adaptor,
51 ConversionPatternRewriter &rewriter) const override;
52};
53
54/// This is separate because in Vulkan workgroup size is exposed to shaders via
55/// a constant with WorkgroupSize decoration. So here we cannot generate a
56/// builtin variable; instead the information in the `spirv.entry_point_abi`
57/// attribute on the surrounding FuncOp is used to replace the gpu::BlockDimOp.
58class WorkGroupSizeConversion : public OpConversionPattern<gpu::BlockDimOp> {
59public:
60 WorkGroupSizeConversion(const TypeConverter &typeConverter,
61 MLIRContext *context)
62 : OpConversionPattern(typeConverter, context, /*benefit*/ 10) {}
63
64 LogicalResult
65 matchAndRewrite(gpu::BlockDimOp op, OpAdaptor adaptor,
66 ConversionPatternRewriter &rewriter) const override;
67};
68
69/// Pattern to convert a kernel function in GPU dialect within a spirv.module.
70class GPUFuncOpConversion final : public OpConversionPattern<gpu::GPUFuncOp> {
71public:
72 using Base::Base;
73
74 LogicalResult
75 matchAndRewrite(gpu::GPUFuncOp funcOp, OpAdaptor adaptor,
76 ConversionPatternRewriter &rewriter) const override;
77
78private:
79 SmallVector<int32_t, 3> workGroupSizeAsInt32;
80};
81
82/// Pattern to convert a gpu.module to a spirv.module.
83class GPUModuleConversion final : public OpConversionPattern<gpu::GPUModuleOp> {
84public:
85 using Base::Base;
86
87 LogicalResult
88 matchAndRewrite(gpu::GPUModuleOp moduleOp, OpAdaptor adaptor,
89 ConversionPatternRewriter &rewriter) const override;
90};
91
92/// Pattern to convert a gpu.return into a SPIR-V return.
93// TODO: This can go to DRR when GPU return has operands.
94class GPUReturnOpConversion final : public OpConversionPattern<gpu::ReturnOp> {
95public:
96 using Base::Base;
97
98 LogicalResult
99 matchAndRewrite(gpu::ReturnOp returnOp, OpAdaptor adaptor,
100 ConversionPatternRewriter &rewriter) const override;
101};
102
103/// Pattern to convert a gpu.barrier op into a spirv.ControlBarrier or
104/// spirv.MemoryNamedBarrier op.
105class GPUBarrierConversion final : public OpConversionPattern<gpu::BarrierOp> {
106public:
107 using Base::Base;
108
109 LogicalResult
110 matchAndRewrite(gpu::BarrierOp barrierOp, OpAdaptor adaptor,
111 ConversionPatternRewriter &rewriter) const override;
112};
113
114/// Pattern to convert a gpu.initialize_named_barrier into
115/// spirv.NamedBarrierInitialize.
116class GPUInitializeNamedBarrierConversion final
117 : public OpConversionPattern<gpu::InitializeNamedBarrierOp> {
118public:
119 using Base::Base;
120
121 LogicalResult
122 matchAndRewrite(gpu::InitializeNamedBarrierOp op, OpAdaptor adaptor,
123 ConversionPatternRewriter &rewriter) const override;
124};
125
126/// Pattern to convert a gpu.shuffle op into a spirv.GroupNonUniformShuffle op.
127class GPUShuffleConversion final : public OpConversionPattern<gpu::ShuffleOp> {
128public:
129 using Base::Base;
130
131 LogicalResult
132 matchAndRewrite(gpu::ShuffleOp shuffleOp, OpAdaptor adaptor,
133 ConversionPatternRewriter &rewriter) const override;
134};
135
136/// Pattern to convert a gpu.rotate op into a spirv.GroupNonUniformRotateKHROp.
137class GPURotateConversion final : public OpConversionPattern<gpu::RotateOp> {
138public:
139 using Base::Base;
140
141 LogicalResult
142 matchAndRewrite(gpu::RotateOp rotateOp, OpAdaptor adaptor,
143 ConversionPatternRewriter &rewriter) const override;
144};
145
146/// Pattern to convert a gpu.subgroup_broadcast op into a
147/// spirv.GroupNonUniformBroadcast op.
148class GPUSubgroupBroadcastConversion final
149 : public OpConversionPattern<gpu::SubgroupBroadcastOp> {
150public:
151 using Base::Base;
152
153 LogicalResult
154 matchAndRewrite(gpu::SubgroupBroadcastOp op, OpAdaptor adaptor,
155 ConversionPatternRewriter &rewriter) const override;
156};
157
158class GPUBallotConversion final : public OpConversionPattern<gpu::BallotOp> {
159public:
160 using Base::Base;
161
162 LogicalResult
163 matchAndRewrite(gpu::BallotOp ballotOp, OpAdaptor adaptor,
164 ConversionPatternRewriter &rewriter) const override;
165};
166
167class GPUPrintfConversion final : public OpConversionPattern<gpu::PrintfOp> {
168public:
169 using Base::Base;
170
171 LogicalResult
172 matchAndRewrite(gpu::PrintfOp gpuPrintfOp, OpAdaptor adaptor,
173 ConversionPatternRewriter &rewriter) const override;
174};
175
176} // namespace
177
178//===----------------------------------------------------------------------===//
179// Builtins.
180//===----------------------------------------------------------------------===//
181
182template <typename SourceOp, spirv::BuiltIn builtin>
183LogicalResult LaunchConfigConversion<SourceOp, builtin>::matchAndRewrite(
184 SourceOp op, typename SourceOp::Adaptor adaptor,
185 ConversionPatternRewriter &rewriter) const {
186 auto *typeConverter = this->template getTypeConverter<SPIRVTypeConverter>();
187 Type indexType = typeConverter->getIndexType();
188
189 // For Vulkan, these SPIR-V builtin variables are required to be a vector of
190 // type <3xi32> by the spec:
191 // https://registry.khronos.org/vulkan/specs/1.3-extensions/man/html/NumWorkgroups.html
192 // https://registry.khronos.org/vulkan/specs/1.3-extensions/man/html/WorkgroupId.html
193 // https://registry.khronos.org/vulkan/specs/1.3-extensions/man/html/WorkgroupSize.html
194 // https://registry.khronos.org/vulkan/specs/1.3-extensions/man/html/LocalInvocationId.html
195 // https://registry.khronos.org/vulkan/specs/1.3-extensions/man/html/LocalInvocationId.html
196 // https://registry.khronos.org/vulkan/specs/1.3-extensions/man/html/GlobalInvocationId.html
197 //
198 // For OpenCL, it depends on the Physical32/Physical64 addressing model:
199 // https://registry.khronos.org/OpenCL/specs/3.0-unified/html/OpenCL_Env.html#_built_in_variables
200 bool forShader =
201 typeConverter->getTargetEnv().allows(spirv::Capability::Shader);
202 Type builtinType = forShader ? rewriter.getIntegerType(32) : indexType;
203
204 Value vector =
205 spirv::getBuiltinVariableValue(op, builtin, builtinType, rewriter);
206 Value dim = spirv::CompositeExtractOp::create(
207 rewriter, op.getLoc(), builtinType, vector,
208 rewriter.getI32ArrayAttr({static_cast<int32_t>(op.getDimension())}));
209 if (forShader && builtinType != indexType)
210 dim = spirv::UConvertOp::create(rewriter, op.getLoc(), indexType, dim);
211 rewriter.replaceOp(op, dim);
212 return success();
213}
214
215template <typename SourceOp, spirv::BuiltIn builtin>
216LogicalResult
217SingleDimLaunchConfigConversion<SourceOp, builtin>::matchAndRewrite(
218 SourceOp op, typename SourceOp::Adaptor adaptor,
219 ConversionPatternRewriter &rewriter) const {
220 auto *typeConverter = this->template getTypeConverter<SPIRVTypeConverter>();
221 Type indexType = typeConverter->getIndexType();
222 Type i32Type = rewriter.getIntegerType(32);
223
224 // For Vulkan, these SPIR-V builtin variables are required to be a vector of
225 // type i32 by the spec:
226 // https://registry.khronos.org/vulkan/specs/1.3-extensions/man/html/NumSubgroups.html
227 // https://registry.khronos.org/vulkan/specs/1.3-extensions/man/html/SubgroupId.html
228 // https://registry.khronos.org/vulkan/specs/1.3-extensions/man/html/SubgroupSize.html
229 //
230 // For OpenCL, they are also required to be i32:
231 // https://registry.khronos.org/OpenCL/specs/3.0-unified/html/OpenCL_Env.html#_built_in_variables
232 Value builtinValue =
233 spirv::getBuiltinVariableValue(op, builtin, i32Type, rewriter);
234 if (i32Type != indexType)
235 builtinValue = spirv::UConvertOp::create(rewriter, op.getLoc(), indexType,
236 builtinValue);
237 rewriter.replaceOp(op, builtinValue);
238 return success();
239}
240
241LogicalResult WorkGroupSizeConversion::matchAndRewrite(
242 gpu::BlockDimOp op, OpAdaptor adaptor,
243 ConversionPatternRewriter &rewriter) const {
245 if (!workGroupSizeAttr)
246 return failure();
247
248 int val =
249 workGroupSizeAttr.asArrayRef()[static_cast<int32_t>(op.getDimension())];
250 auto convertedType =
251 getTypeConverter()->convertType(op.getResult().getType());
252 if (!convertedType)
253 return failure();
254 rewriter.replaceOpWithNewOp<spirv::ConstantOp>(
255 op, convertedType, IntegerAttr::get(convertedType, val));
256 return success();
257}
258
259//===----------------------------------------------------------------------===//
260// GPUFuncOp
261//===----------------------------------------------------------------------===//
262
263// Legalizes a GPU function as an entry SPIR-V function.
264static spirv::FuncOp
265lowerAsEntryFunction(gpu::GPUFuncOp funcOp, const TypeConverter &typeConverter,
266 ConversionPatternRewriter &rewriter,
267 spirv::EntryPointABIAttr entryPointInfo,
269 auto fnType = funcOp.getFunctionType();
270 if (fnType.getNumResults()) {
271 funcOp.emitError("SPIR-V lowering only supports entry functions"
272 "with no return values right now");
273 return nullptr;
274 }
275 if (!argABIInfo.empty() && fnType.getNumInputs() != argABIInfo.size()) {
276 funcOp.emitError(
277 "lowering as entry functions requires ABI info for all arguments "
278 "or none of them");
279 return nullptr;
280 }
281 // Update the signature to valid SPIR-V types and add the ABI
282 // attributes. These will be "materialized" by using the
283 // LowerABIAttributesPass.
284 TypeConverter::SignatureConversion signatureConverter(fnType.getNumInputs());
285 {
286 for (const auto &argType :
287 enumerate(funcOp.getFunctionType().getInputs())) {
288 auto convertedType = typeConverter.convertType(argType.value());
289 if (!convertedType)
290 return nullptr;
291 signatureConverter.addInputs(argType.index(), convertedType);
292 }
293 }
294 auto newFuncOp = spirv::FuncOp::create(
295 rewriter, funcOp.getLoc(), funcOp.getName(),
296 rewriter.getFunctionType(signatureConverter.getConvertedTypes(), {}));
297 newFuncOp.setArgAttrsAttr(funcOp.getArgAttrsAttr());
298 newFuncOp.setResAttrsAttr(funcOp.getResAttrsAttr());
299 cast<SymbolOpInterface>(newFuncOp.getOperation())
300 .setVisibility(
301 cast<SymbolOpInterface>(funcOp.getOperation()).getVisibility());
302
303 auto copyGPUProperty = [&](StringAttr name, Attribute value) {
304 if (value)
305 newFuncOp->setDiscardableAttr(name, value);
306 };
307 copyGPUProperty(funcOp.getWorkgroupAttribAttrsAttrName(),
308 funcOp.getWorkgroupAttribAttrsAttr());
309 copyGPUProperty(funcOp.getPrivateAttribAttrsAttrName(),
310 funcOp.getPrivateAttribAttrsAttr());
311 copyGPUProperty(funcOp.getKnownBlockSizeAttrName(),
312 funcOp.getKnownBlockSizeAttr());
313 copyGPUProperty(funcOp.getKnownGridSizeAttrName(),
314 funcOp.getKnownGridSizeAttr());
315 copyGPUProperty(funcOp.getKnownClusterSizeAttrName(),
316 funcOp.getKnownClusterSizeAttr());
317 copyGPUProperty(funcOp.getWorkgroupAttributionsAttrName(),
318 funcOp.getWorkgroupAttributionsAttr());
319 for (const auto &discardableAttr :
320 funcOp->getDiscardableAttrDictionary().getValue())
321 newFuncOp->setDiscardableAttr(discardableAttr.getName(),
322 discardableAttr.getValue());
323
324 rewriter.inlineRegionBefore(funcOp.getBody(), newFuncOp.getBody(),
325 newFuncOp.end());
326 if (failed(rewriter.convertRegionTypes(&newFuncOp.getBody(), typeConverter,
327 &signatureConverter)))
328 return nullptr;
329 rewriter.eraseOp(funcOp);
330
331 // Set the attributes for argument and the function.
332 StringRef argABIAttrName = spirv::getInterfaceVarABIAttrName();
333 for (auto argIndex : llvm::seq<unsigned>(0, argABIInfo.size())) {
334 newFuncOp.setArgAttr(argIndex, argABIAttrName, argABIInfo[argIndex]);
335 }
336 newFuncOp->setDiscardableAttr(spirv::getEntryPointABIAttrName(),
337 entryPointInfo);
338
339 return newFuncOp;
340}
341
342/// Populates `argABI` with spirv.interface_var_abi attributes for lowering
343/// gpu.func to spirv.func if no arguments have the attributes set
344/// already. Returns failure if any argument has the ABI attribute set already.
345static LogicalResult
346getDefaultABIAttrs(const spirv::TargetEnv &targetEnv, gpu::GPUFuncOp funcOp,
348 if (!spirv::needsInterfaceVarABIAttrs(targetEnv))
349 return success();
350
351 for (auto argIndex : llvm::seq<unsigned>(0, funcOp.getNumArguments())) {
352 if (funcOp.getArgAttrOfType<spirv::InterfaceVarABIAttr>(
354 return failure();
355 // Vulkan's interface variable requirements needs scalars to be wrapped in a
356 // struct. The struct held in storage buffer.
357 std::optional<spirv::StorageClass> sc;
358 if (funcOp.getArgument(argIndex).getType().isIntOrIndexOrFloat())
359 sc = spirv::StorageClass::StorageBuffer;
360 argABI.push_back(
361 spirv::getInterfaceVarABIAttr(0, argIndex, sc, funcOp.getContext()));
362 }
363 return success();
364}
365
366LogicalResult GPUFuncOpConversion::matchAndRewrite(
367 gpu::GPUFuncOp funcOp, OpAdaptor adaptor,
368 ConversionPatternRewriter &rewriter) const {
369 if (!gpu::GPUDialect::isKernel(funcOp))
370 return failure();
371
372 auto *typeConverter = getTypeConverter<SPIRVTypeConverter>();
373 SmallVector<spirv::InterfaceVarABIAttr, 4> argABI;
374 if (failed(
375 getDefaultABIAttrs(typeConverter->getTargetEnv(), funcOp, argABI))) {
376 argABI.clear();
377 for (auto argIndex : llvm::seq<unsigned>(0, funcOp.getNumArguments())) {
378 // If the ABI is already specified, use it.
379 auto abiAttr = funcOp.getArgAttrOfType<spirv::InterfaceVarABIAttr>(
381 if (!abiAttr) {
382 funcOp.emitRemark(
383 "match failure: missing 'spirv.interface_var_abi' attribute at "
384 "argument ")
385 << argIndex;
386 return failure();
387 }
388 argABI.push_back(abiAttr);
389 }
390 }
391
392 auto entryPointAttr = spirv::lookupEntryPointABI(funcOp);
393 if (!entryPointAttr) {
394 funcOp.emitRemark(
395 "match failure: missing 'spirv.entry_point_abi' attribute");
396 return failure();
397 }
398 spirv::FuncOp newFuncOp = lowerAsEntryFunction(
399 funcOp, *getTypeConverter(), rewriter, entryPointAttr, argABI);
400 if (!newFuncOp)
401 return failure();
402 newFuncOp->removeDiscardableAttr(
403 rewriter.getStringAttr(gpu::GPUDialect::getKernelFuncAttrName()));
404 return success();
405}
406
407//===----------------------------------------------------------------------===//
408// ModuleOp with gpu.module.
409//===----------------------------------------------------------------------===//
410
411LogicalResult GPUModuleConversion::matchAndRewrite(
412 gpu::GPUModuleOp moduleOp, OpAdaptor adaptor,
413 ConversionPatternRewriter &rewriter) const {
414 auto *typeConverter = getTypeConverter<SPIRVTypeConverter>();
415 const spirv::TargetEnv &targetEnv = typeConverter->getTargetEnv();
416 spirv::AddressingModel addressingModel = spirv::getAddressingModel(
417 targetEnv, typeConverter->getOptions().use64bitIndex);
418 FailureOr<spirv::MemoryModel> memoryModel = spirv::getMemoryModel(targetEnv);
419 if (failed(memoryModel))
420 return moduleOp.emitRemark(
421 "cannot deduce memory model from 'spirv.target_env'");
422
423 // Add a keyword to the module name to avoid symbolic conflict.
424 std::string spvModuleName = (kSPIRVModule + moduleOp.getName()).str();
425 auto spvModule = spirv::ModuleOp::create(
426 rewriter, moduleOp.getLoc(), addressingModel, *memoryModel, std::nullopt,
427 StringRef(spvModuleName));
428
429 // Move the region from the module op into the SPIR-V module.
430 Region &spvModuleRegion = spvModule.getRegion();
431 rewriter.inlineRegionBefore(moduleOp.getBodyRegion(), spvModuleRegion,
432 spvModuleRegion.begin());
433 // The spirv.module build method adds a block. Remove that.
434 rewriter.eraseBlock(&spvModuleRegion.back());
435
436 // Some of the patterns call `lookupTargetEnv` during conversion and they
437 // will fail if called after GPUModuleConversion and we don't preserve
438 // `TargetEnv` attribute.
439 // Copy TargetEnvAttr only if it is attached directly to the GPUModuleOp.
440 if (auto attr = moduleOp->getDiscardableAttrOfType<spirv::TargetEnvAttr>(
442 spvModule->setDiscardableAttr(spirv::getTargetEnvAttrName(), attr);
443 if (ArrayAttr targets = moduleOp.getTargetsAttr()) {
444 for (Attribute targetAttr : targets)
445 if (auto spirvTargetEnvAttr =
446 dyn_cast<spirv::TargetEnvAttr>(targetAttr)) {
447 spvModule->setDiscardableAttr(spirv::getTargetEnvAttrName(),
448 spirvTargetEnvAttr);
449 break;
450 }
451 }
452
453 rewriter.eraseOp(moduleOp);
454 return success();
455}
456
457//===----------------------------------------------------------------------===//
458// GPU return inside kernel functions to SPIR-V return.
459//===----------------------------------------------------------------------===//
460
461LogicalResult GPUReturnOpConversion::matchAndRewrite(
462 gpu::ReturnOp returnOp, OpAdaptor adaptor,
463 ConversionPatternRewriter &rewriter) const {
464 if (!adaptor.getOperands().empty())
465 return failure();
466
467 rewriter.replaceOpWithNewOp<spirv::ReturnOp>(returnOp);
468 return success();
469}
470
471//===----------------------------------------------------------------------===//
472// Barrier.
473//===----------------------------------------------------------------------===//
474
475/// Map gpu::BarrierScope to spirv::Scope.
476static FailureOr<spirv::Scope>
477mapGPUBarrierScopeToSPIRV(gpu::BarrierScope gpuScope) {
478 switch (gpuScope) {
479 case gpu::BarrierScope::Subgroup:
480 return spirv::Scope::Subgroup;
481 case gpu::BarrierScope::Workgroup:
482 return spirv::Scope::Workgroup;
483 case gpu::BarrierScope::Cluster:
484 return failure();
485 }
486 return failure();
487}
488
489LogicalResult GPUBarrierConversion::matchAndRewrite(
490 gpu::BarrierOp barrierOp, OpAdaptor adaptor,
491 ConversionPatternRewriter &rewriter) const {
492 MLIRContext *context = getContext();
493
494 // Map GPU scope to SPIR-V scope.
495 auto spirvScope = mapGPUBarrierScopeToSPIRV(barrierOp.getScope());
496 if (failed(spirvScope))
497 return rewriter.notifyMatchFailure(
498 barrierOp, "cluster scope is not supported in SPIR-V");
499
500 auto scopeAttr = spirv::ScopeAttr::get(context, *spirvScope);
501 auto memoryScopeAttr =
502 spirv::ScopeAttr::get(context, spirv::Scope::Workgroup);
503
504 // Require acquire and release memory semantics for workgroup memory.
505 auto memorySemantics = spirv::MemorySemanticsAttr::get(
506 context, spirv::MemorySemantics::WorkgroupMemory |
507 spirv::MemorySemantics::AcquireRelease);
508
509 if (adaptor.getNamedBarrier()) {
510 spirv::MemoryNamedBarrierOp::create(rewriter, barrierOp.getLoc(),
511 adaptor.getNamedBarrier(),
512 memoryScopeAttr, memorySemantics);
513 rewriter.eraseOp(barrierOp);
514 } else {
515 rewriter.replaceOpWithNewOp<spirv::ControlBarrierOp>(
516 barrierOp, scopeAttr, memoryScopeAttr, memorySemantics);
517 }
518 return success();
519}
520
521LogicalResult GPUInitializeNamedBarrierConversion::matchAndRewrite(
522 gpu::InitializeNamedBarrierOp op, OpAdaptor adaptor,
523 ConversionPatternRewriter &rewriter) const {
525 rewriter.replaceOpWithNewOp<spirv::NamedBarrierInitializeOp>(
526 op, nbType, adaptor.getMemberCount());
527 return success();
528}
529
530//===----------------------------------------------------------------------===//
531// Shuffle
532//===----------------------------------------------------------------------===//
533
534LogicalResult GPUShuffleConversion::matchAndRewrite(
535 gpu::ShuffleOp shuffleOp, OpAdaptor adaptor,
536 ConversionPatternRewriter &rewriter) const {
537 // Require the shuffle width to be the same as the target's subgroup size,
538 // given that for SPIR-V non-uniform subgroup ops, we cannot select
539 // participating invocations.
540 const spirv::TargetEnv &targetEnv =
541 getTypeConverter<SPIRVTypeConverter>()->getTargetEnv();
542 unsigned subgroupSize =
543 targetEnv.getAttr().getResourceLimits().getSubgroupSize();
544 IntegerAttr widthAttr;
545 if (!matchPattern(shuffleOp.getWidth(), m_Constant(&widthAttr)) ||
546 widthAttr.getValue().getZExtValue() != subgroupSize)
547 return rewriter.notifyMatchFailure(
548 shuffleOp, "shuffle width and target subgroup size mismatch");
549
550 assert(!adaptor.getOffset().getType().isSignedInteger() &&
551 "shuffle offset must be a signless/unsigned integer");
552
553 Location loc = shuffleOp.getLoc();
554 auto scope = rewriter.getAttr<spirv::ScopeAttr>(spirv::Scope::Subgroup);
555 Value result;
556 Value validVal;
557
558 switch (shuffleOp.getMode()) {
559 case gpu::ShuffleMode::XOR: {
560 result = spirv::GroupNonUniformShuffleXorOp::create(
561 rewriter, loc, scope, adaptor.getValue(), adaptor.getOffset());
562 validVal = spirv::ConstantOp::getOne(rewriter.getI1Type(),
563 shuffleOp.getLoc(), rewriter);
564 break;
565 }
566 case gpu::ShuffleMode::IDX: {
567 result = spirv::GroupNonUniformShuffleOp::create(
568 rewriter, loc, scope, adaptor.getValue(), adaptor.getOffset());
569 validVal = spirv::ConstantOp::getOne(rewriter.getI1Type(),
570 shuffleOp.getLoc(), rewriter);
571 break;
572 }
573 case gpu::ShuffleMode::DOWN: {
574 result = spirv::GroupNonUniformShuffleDownOp::create(
575 rewriter, loc, scope, adaptor.getValue(), adaptor.getOffset());
576
577 Value laneId = gpu::LaneIdOp::create(rewriter, loc, widthAttr);
578 Value resultLaneId =
579 arith::AddIOp::create(rewriter, loc, laneId, adaptor.getOffset());
580 validVal = arith::CmpIOp::create(rewriter, loc, arith::CmpIPredicate::ult,
581 resultLaneId, adaptor.getWidth());
582 break;
583 }
584 case gpu::ShuffleMode::UP: {
585 result = spirv::GroupNonUniformShuffleUpOp::create(
586 rewriter, loc, scope, adaptor.getValue(), adaptor.getOffset());
587
588 Value laneId = gpu::LaneIdOp::create(rewriter, loc, widthAttr);
589 Value resultLaneId =
590 arith::SubIOp::create(rewriter, loc, laneId, adaptor.getOffset());
591 auto i32Type = rewriter.getIntegerType(32);
592 validVal = arith::CmpIOp::create(
593 rewriter, loc, arith::CmpIPredicate::sge, resultLaneId,
594 arith::ConstantOp::create(rewriter, loc, i32Type,
595 rewriter.getIntegerAttr(i32Type, 0)));
596 break;
597 }
598 }
599
600 rewriter.replaceOp(shuffleOp, {result, validVal});
601 return success();
602}
603
604//===----------------------------------------------------------------------===//
605// Rotate
606//===----------------------------------------------------------------------===//
607
608LogicalResult GPURotateConversion::matchAndRewrite(
609 gpu::RotateOp rotateOp, OpAdaptor adaptor,
610 ConversionPatternRewriter &rewriter) const {
611 const spirv::TargetEnv &targetEnv =
612 getTypeConverter<SPIRVTypeConverter>()->getTargetEnv();
613 unsigned subgroupSize =
614 targetEnv.getAttr().getResourceLimits().getSubgroupSize();
615 unsigned width = rotateOp.getWidth();
616 if (width > subgroupSize)
617 return rewriter.notifyMatchFailure(
618 rotateOp, "rotate width is larger than target subgroup size");
619
620 Location loc = rotateOp.getLoc();
621 auto scope = rewriter.getAttr<spirv::ScopeAttr>(spirv::Scope::Subgroup);
622 Value offsetVal =
623 arith::ConstantOp::create(rewriter, loc, adaptor.getOffsetAttr());
624 Value widthVal =
625 arith::ConstantOp::create(rewriter, loc, adaptor.getWidthAttr());
626 Value rotateResult = spirv::GroupNonUniformRotateKHROp::create(
627 rewriter, loc, scope, adaptor.getValue(), offsetVal, widthVal);
628 Value validVal;
629 if (width == subgroupSize) {
630 validVal = spirv::ConstantOp::getOne(rewriter.getI1Type(), loc, rewriter);
631 } else {
632 IntegerAttr widthAttr = adaptor.getWidthAttr();
633 Value laneId = gpu::LaneIdOp::create(rewriter, loc, widthAttr);
634 validVal = arith::CmpIOp::create(rewriter, loc, arith::CmpIPredicate::ult,
635 laneId, widthVal);
636 }
637
638 rewriter.replaceOp(rotateOp, {rotateResult, validVal});
639 return success();
640}
641
642//===----------------------------------------------------------------------===//
643// Subgroup broadcast
644//===----------------------------------------------------------------------===//
645
646LogicalResult GPUSubgroupBroadcastConversion::matchAndRewrite(
647 gpu::SubgroupBroadcastOp op, OpAdaptor adaptor,
648 ConversionPatternRewriter &rewriter) const {
649 Location loc = op.getLoc();
650 auto scope = rewriter.getAttr<spirv::ScopeAttr>(spirv::Scope::Subgroup);
651 Value result;
652
653 switch (op.getBroadcastType()) {
654 case gpu::BroadcastType::specific_lane:
655 result = spirv::GroupNonUniformBroadcastOp::create(
656 rewriter, loc, scope, adaptor.getSrc(), adaptor.getLane());
657 break;
658 case gpu::BroadcastType::first_active_lane:
659 result = spirv::GroupNonUniformBroadcastFirstOp::create(
660 rewriter, loc, scope, adaptor.getSrc());
661 break;
662 }
663
664 rewriter.replaceOp(op, result);
665 return success();
666}
667
668LogicalResult GPUBallotConversion::matchAndRewrite(
669 gpu::BallotOp ballotOp, OpAdaptor adaptor,
670 ConversionPatternRewriter &rewriter) const {
671 Location loc = ballotOp.getLoc();
672 auto scope = rewriter.getAttr<spirv::ScopeAttr>(spirv::Scope::Subgroup);
673 auto int32Type = rewriter.getI32Type();
674 auto vec4i32Type = VectorType::get({4}, int32Type);
675
676 // SPIR-V ballot returns vector<4xi32> to support subgroups up to 128 lanes.
677 Value ballot = spirv::GroupNonUniformBallotOp::create(
678 rewriter, loc, vec4i32Type, scope, adaptor.getPredicate());
679
680 auto intType = cast<IntegerType>(ballotOp.getType());
681 unsigned width = intType.getWidth();
682
683 if (width == 32) {
684 Value result =
685 spirv::CompositeExtractOp::create(rewriter, loc, ballot, {0});
686 rewriter.replaceOp(ballotOp, result);
687 } else if (width == 64) {
688 // Combine first two vector elements: low 32 bits + (high 32 bits << 32).
689 Value low = spirv::CompositeExtractOp::create(rewriter, loc, ballot, {0});
690 Value high = spirv::CompositeExtractOp::create(rewriter, loc, ballot, {1});
691
692 auto int64Type = rewriter.getI64Type();
693 Value lowExt = spirv::UConvertOp::create(rewriter, loc, int64Type, low);
694 Value highExt = spirv::UConvertOp::create(rewriter, loc, int64Type, high);
695
696 Value shift32 = spirv::ConstantOp::create(
697 rewriter, loc, int64Type, rewriter.getIntegerAttr(int64Type, 32));
698 Value highShifted =
699 spirv::ShiftLeftLogicalOp::create(rewriter, loc, highExt, shift32);
700
701 Value result =
702 spirv::BitwiseOrOp::create(rewriter, loc, lowExt, highShifted);
703 rewriter.replaceOp(ballotOp, result);
704 } else {
705 return rewriter.notifyMatchFailure(
706 ballotOp, "only i32 and i64 result types are supported for SPIR-V");
707 }
708
709 return success();
710}
711
712//===----------------------------------------------------------------------===//
713// Group ops
714//===----------------------------------------------------------------------===//
715
716template <typename UniformOp, typename NonUniformOp>
718 Value arg, bool isGroup, bool isUniform,
719 std::optional<uint32_t> clusterSize) {
720 spirv::Scope scope =
721 isGroup ? spirv::Scope::Workgroup : spirv::Scope::Subgroup;
722 // GroupNonUniform* ops only support Subgroup scope.
723 if (!isUniform && scope != spirv::Scope::Subgroup)
724 return Value();
725
726 Type type = arg.getType();
727 auto scopeAttr = mlir::spirv::ScopeAttr::get(builder.getContext(), scope);
728 auto groupOp = spirv::GroupOperationAttr::get(
729 builder.getContext(), clusterSize.has_value()
730 ? spirv::GroupOperation::ClusteredReduce
731 : spirv::GroupOperation::Reduce);
732 if (isUniform) {
733 return UniformOp::create(builder, loc, type, scopeAttr, groupOp, arg)
734 .getResult();
735 }
736
737 Value clusterSizeValue;
738 if (clusterSize.has_value())
739 clusterSizeValue = spirv::ConstantOp::create(
740 builder, loc, builder.getI32Type(),
741 builder.getIntegerAttr(builder.getI32Type(), *clusterSize));
742
743 return NonUniformOp::create(builder, loc, type, scopeAttr, groupOp, arg,
744 clusterSizeValue)
745 .getResult();
746}
747
748template <typename NonUniformOp>
750 OpBuilder &builder, Location loc, Value arg, bool isGroup, bool isUniform,
751 std::optional<uint32_t> clusterSize) {
752 spirv::Scope scope =
753 isGroup ? spirv::Scope::Workgroup : spirv::Scope::Subgroup;
754 if (isUniform || scope != spirv::Scope::Subgroup)
755 return Value();
756
757 Type type = arg.getType();
758 auto scopeAttr = mlir::spirv::ScopeAttr::get(builder.getContext(), scope);
759 auto groupOp = spirv::GroupOperationAttr::get(
760 builder.getContext(), clusterSize.has_value()
761 ? spirv::GroupOperation::ClusteredReduce
762 : spirv::GroupOperation::Reduce);
763
764 Value clusterSizeValue;
765 if (clusterSize.has_value())
766 clusterSizeValue = spirv::ConstantOp::create(
767 builder, loc, builder.getI32Type(),
768 builder.getIntegerAttr(builder.getI32Type(), *clusterSize));
769
770 return NonUniformOp::create(builder, loc, type, scopeAttr, groupOp, arg,
771 clusterSizeValue)
772 .getResult();
773}
774
775static std::optional<Value>
777 gpu::AllReduceOperation opType, bool isGroup,
778 bool isUniform, std::optional<uint32_t> clusterSize) {
779 enum class ElemType { Float, Boolean, Integer };
780 using FuncT = Value (*)(OpBuilder &, Location, Value, bool, bool,
781 std::optional<uint32_t>);
782 struct OpHandler {
783 gpu::AllReduceOperation kind;
784 ElemType elemType;
785 FuncT func;
786 };
787
788 Type type = arg.getType();
789 ElemType elementType;
790 if (isa<FloatType>(type)) {
791 elementType = ElemType::Float;
792 } else if (auto intTy = dyn_cast<IntegerType>(type)) {
793 elementType = (intTy.getIntOrFloatBitWidth() == 1) ? ElemType::Boolean
794 : ElemType::Integer;
795 } else {
796 return std::nullopt;
797 }
798
799 // TODO(https://github.com/llvm/llvm-project/issues/73459): The SPIR-V spec
800 // does not specify how -0.0 / +0.0 and NaN values are handled in *FMin/*FMax
801 // reduction ops. We should account possible precision requirements in this
802 // conversion.
803
804 using ReduceType = gpu::AllReduceOperation;
805 const OpHandler handlers[] = {
806 {ReduceType::ADD, ElemType::Integer,
807 &createGroupReduceOpImpl<spirv::GroupIAddOp,
808 spirv::GroupNonUniformIAddOp>},
809 {ReduceType::ADD, ElemType::Float,
810 &createGroupReduceOpImpl<spirv::GroupFAddOp,
811 spirv::GroupNonUniformFAddOp>},
812 {ReduceType::MUL, ElemType::Integer,
813 &createGroupReduceOpImpl<spirv::GroupIMulKHROp,
814 spirv::GroupNonUniformIMulOp>},
815 {ReduceType::MUL, ElemType::Float,
816 &createGroupReduceOpImpl<spirv::GroupFMulKHROp,
817 spirv::GroupNonUniformFMulOp>},
818 {ReduceType::MINUI, ElemType::Integer,
819 &createGroupReduceOpImpl<spirv::GroupUMinOp,
820 spirv::GroupNonUniformUMinOp>},
821 {ReduceType::MINSI, ElemType::Integer,
822 &createGroupReduceOpImpl<spirv::GroupSMinOp,
823 spirv::GroupNonUniformSMinOp>},
824 {ReduceType::MINNUMF, ElemType::Float,
825 &createGroupReduceOpImpl<spirv::GroupFMinOp,
826 spirv::GroupNonUniformFMinOp>},
827 {ReduceType::MAXUI, ElemType::Integer,
828 &createGroupReduceOpImpl<spirv::GroupUMaxOp,
829 spirv::GroupNonUniformUMaxOp>},
830 {ReduceType::MAXSI, ElemType::Integer,
831 &createGroupReduceOpImpl<spirv::GroupSMaxOp,
832 spirv::GroupNonUniformSMaxOp>},
833 {ReduceType::MAXNUMF, ElemType::Float,
834 &createGroupReduceOpImpl<spirv::GroupFMaxOp,
835 spirv::GroupNonUniformFMaxOp>},
836 {ReduceType::MINIMUMF, ElemType::Float,
837 &createGroupReduceOpImpl<spirv::GroupFMinOp,
838 spirv::GroupNonUniformFMinOp>},
839 {ReduceType::MAXIMUMF, ElemType::Float,
840 &createGroupReduceOpImpl<spirv::GroupFMaxOp,
841 spirv::GroupNonUniformFMaxOp>},
842 {ReduceType::AND, ElemType::Integer,
844 spirv::GroupNonUniformBitwiseAndOp>},
845 {ReduceType::OR, ElemType::Integer,
847 spirv::GroupNonUniformBitwiseOrOp>},
848 {ReduceType::XOR, ElemType::Integer,
850 spirv::GroupNonUniformBitwiseXorOp>}};
851
852 for (const OpHandler &handler : handlers)
853 if (handler.kind == opType && elementType == handler.elemType)
854 if (Value result =
855 handler.func(builder, loc, arg, isGroup, isUniform, clusterSize))
856 return result;
857
858 return std::nullopt;
859}
860
861/// Pattern to convert a gpu.all_reduce op into a SPIR-V group op.
863 : public OpConversionPattern<gpu::AllReduceOp> {
864public:
865 using Base::Base;
866
867 LogicalResult
868 matchAndRewrite(gpu::AllReduceOp op, OpAdaptor adaptor,
869 ConversionPatternRewriter &rewriter) const override {
870 auto opType = op.getOp();
871
872 // gpu.all_reduce can have either reduction op attribute or reduction
873 // region. Only attribute version is supported.
874 if (!opType)
875 return failure();
876
877 auto result =
878 createGroupReduceOp(rewriter, op.getLoc(), adaptor.getValue(), *opType,
879 /*isGroup*/ true, op.getUniform(), std::nullopt);
880 if (!result)
881 return failure();
882
883 rewriter.replaceOp(op, *result);
884 return success();
885 }
886};
887
888/// Pattern to convert a gpu.subgroup_reduce op into a SPIR-V group op.
890 : public OpConversionPattern<gpu::SubgroupReduceOp> {
891public:
892 using Base::Base;
893
894 LogicalResult
895 matchAndRewrite(gpu::SubgroupReduceOp op, OpAdaptor adaptor,
896 ConversionPatternRewriter &rewriter) const override {
897 if (op.getClusterStride() > 1) {
898 return rewriter.notifyMatchFailure(
899 op, "lowering for cluster stride > 1 is not implemented");
900 }
901
902 if (!isa<spirv::ScalarType>(adaptor.getValue().getType()))
903 return rewriter.notifyMatchFailure(op, "reduction type is not a scalar");
904
906 rewriter, op.getLoc(), adaptor.getValue(), adaptor.getOp(),
907 /*isGroup=*/false, adaptor.getUniform(), op.getClusterSize());
908 if (!result)
909 return failure();
910
911 rewriter.replaceOp(op, *result);
912 return success();
913 }
914};
915
916// Formulate a unique variable/constant name after
917// searching in the module for existing variable/constant names.
918// This is to avoid name collision with existing variables.
919// Example: printfMsg0, printfMsg1, printfMsg2, ...
920static std::string makeVarName(spirv::ModuleOp moduleOp, llvm::Twine prefix) {
921 std::string name;
922 unsigned number = 0;
923
924 do {
925 name.clear();
926 name = (prefix + llvm::Twine(number++)).str();
927 } while (moduleOp.lookupSymbol(name));
928
929 return name;
930}
931
932/// Pattern to convert a gpu.printf op into a SPIR-V CLPrintf op.
933
934LogicalResult GPUPrintfConversion::matchAndRewrite(
935 gpu::PrintfOp gpuPrintfOp, OpAdaptor adaptor,
936 ConversionPatternRewriter &rewriter) const {
937
938 Location loc = gpuPrintfOp.getLoc();
939
940 auto moduleOp = gpuPrintfOp->getParentOfType<spirv::ModuleOp>();
941 if (!moduleOp)
942 return failure();
943
944 // SPIR-V global variable is used to initialize printf
945 // format string value, if there are multiple printf messages,
946 // each global var needs to be created with a unique name.
947 std::string globalVarName = makeVarName(moduleOp, llvm::Twine("printfMsg"));
948 spirv::GlobalVariableOp globalVar;
949
950 IntegerType i8Type = rewriter.getI8Type();
951 IntegerType i32Type = rewriter.getI32Type();
952
953 // Each character of printf format string is
954 // stored as a spec constant. We need to create
955 // unique name for this spec constant like
956 // @printfMsg0_sc0, @printfMsg0_sc1, ... by searching in the module
957 // for existing spec constant names.
958 auto createSpecConstant = [&](unsigned value) {
959 auto attr = rewriter.getI8IntegerAttr(value);
960 std::string specCstName =
961 makeVarName(moduleOp, llvm::Twine(globalVarName) + "_sc");
962
963 return spirv::SpecConstantOp::create(
964 rewriter, loc, rewriter.getStringAttr(specCstName), attr,
965 /*sym_visibility=*/nullptr);
966 };
967 {
968 Operation *parent =
969 SymbolTable::getNearestSymbolTable(gpuPrintfOp->getParentOp());
970
971 ConversionPatternRewriter::InsertionGuard guard(rewriter);
972
973 Block &entryBlock = *parent->getRegion(0).begin();
974 rewriter.setInsertionPointToStart(
975 &entryBlock); // insertion point at module level
976
977 // Create Constituents with SpecConstant by scanning format string
978 // Each character of format string is stored as a spec constant
979 // and then these spec constants are used to create a
980 // SpecConstantCompositeOp.
981 llvm::SmallString<20> formatString(adaptor.getFormat());
982 formatString.push_back('\0'); // Null terminate for C.
983 SmallVector<Attribute, 4> constituents;
984 for (char c : formatString) {
985 spirv::SpecConstantOp cSpecConstantOp = createSpecConstant(c);
986 constituents.push_back(SymbolRefAttr::get(cSpecConstantOp));
987 }
988
989 // Create SpecConstantCompositeOp to initialize the global variable
990 size_t contentSize = constituents.size();
991 auto globalType = spirv::ArrayType::get(i8Type, contentSize);
992 spirv::SpecConstantCompositeOp specCstComposite;
993 // There will be one SpecConstantCompositeOp per printf message/global var,
994 // so no need do lookup for existing ones.
995 std::string specCstCompositeName =
996 (llvm::Twine(globalVarName) + "_scc").str();
997
998 specCstComposite = spirv::SpecConstantCompositeOp::create(
999 rewriter, loc, TypeAttr::get(globalType),
1000 rewriter.getStringAttr(specCstCompositeName),
1001 rewriter.getArrayAttr(constituents), /*sym_visibility=*/nullptr);
1002
1003 auto ptrType = spirv::PointerType::get(
1004 globalType, spirv::StorageClass::UniformConstant);
1005
1006 // Define a GlobalVarOp initialized using specialized constants
1007 // that is used to specify the printf format string
1008 // to be passed to the SPIRV CLPrintfOp.
1009 globalVar = spirv::GlobalVariableOp::create(
1010 rewriter, loc, ptrType, globalVarName,
1011 FlatSymbolRefAttr::get(specCstComposite));
1012
1013 globalVar->setDiscardableAttr("Constant", rewriter.getUnitAttr());
1014 }
1015 // Get SSA value of Global variable and create pointer to i8 to point to
1016 // the format string.
1017 Value globalPtr = spirv::AddressOfOp::create(rewriter, loc, globalVar);
1018 Value fmtStr = spirv::BitcastOp::create(
1019 rewriter, loc,
1020 spirv::PointerType::get(i8Type, spirv::StorageClass::UniformConstant),
1021 globalPtr);
1022
1023 // Get printf arguments.
1024 auto printfArgs = llvm::to_vector_of<Value, 4>(adaptor.getArgs());
1025
1026 spirv::CLPrintfOp::create(rewriter, loc, i32Type, fmtStr, printfArgs);
1027
1028 // Need to erase the gpu.printf op as gpu.printf does not use result vs
1029 // spirv::CLPrintfOp has i32 resultType so cannot replace with new SPIR-V
1030 // printf op.
1031 rewriter.eraseOp(gpuPrintfOp);
1032
1033 return success();
1034}
1035
1036//===----------------------------------------------------------------------===//
1037// GPU To SPIRV Patterns.
1038//===----------------------------------------------------------------------===//
1039
1041 RewritePatternSet &patterns) {
1042 patterns.add<
1043 GPUBarrierConversion, GPUInitializeNamedBarrierConversion,
1044 GPUBallotConversion, GPUFuncOpConversion, GPUModuleConversion,
1045 GPUReturnOpConversion, GPUShuffleConversion, GPURotateConversion,
1046 GPUSubgroupBroadcastConversion,
1047 LaunchConfigConversion<gpu::BlockIdOp, spirv::BuiltIn::WorkgroupId>,
1048 LaunchConfigConversion<gpu::GridDimOp, spirv::BuiltIn::NumWorkgroups>,
1049 LaunchConfigConversion<gpu::BlockDimOp, spirv::BuiltIn::WorkgroupSize>,
1050 LaunchConfigConversion<gpu::ThreadIdOp,
1051 spirv::BuiltIn::LocalInvocationId>,
1052 LaunchConfigConversion<gpu::GlobalIdOp,
1053 spirv::BuiltIn::GlobalInvocationId>,
1054 SingleDimLaunchConfigConversion<gpu::SubgroupIdOp,
1055 spirv::BuiltIn::SubgroupId>,
1056 SingleDimLaunchConfigConversion<gpu::NumSubgroupsOp,
1057 spirv::BuiltIn::NumSubgroups>,
1058 SingleDimLaunchConfigConversion<gpu::SubgroupSizeOp,
1059 spirv::BuiltIn::SubgroupSize>,
1060 SingleDimLaunchConfigConversion<
1061 gpu::LaneIdOp, spirv::BuiltIn::SubgroupLocalInvocationId>,
1062 WorkGroupSizeConversion, GPUAllReduceConversion,
1063 GPUSubgroupReduceConversion, GPUPrintfConversion>(typeConverter,
1064 patterns.getContext());
1065}
1066
1068 SPIRVTypeConverter &typeConverter) {
1069 typeConverter.addConversion([](gpu::NamedBarrierType type) {
1070 return spirv::NamedBarrierType::get(type.getContext());
1071 });
1072}
return success()
static std::optional< Value > createGroupReduceOp(OpBuilder &builder, Location loc, Value arg, gpu::AllReduceOperation opType, bool isGroup, bool isUniform, std::optional< uint32_t > clusterSize)
static FailureOr< spirv::Scope > mapGPUBarrierScopeToSPIRV(gpu::BarrierScope gpuScope)
Map gpu::BarrierScope to spirv::Scope.
static LogicalResult getDefaultABIAttrs(const spirv::TargetEnv &targetEnv, gpu::GPUFuncOp funcOp, SmallVectorImpl< spirv::InterfaceVarABIAttr > &argABI)
Populates argABI with spirv.interface_var_abi attributes for lowering gpu.func to spirv....
static constexpr const char kSPIRVModule[]
static Value createGroupReduceOpImpl(OpBuilder &builder, Location loc, Value arg, bool isGroup, bool isUniform, std::optional< uint32_t > clusterSize)
static Value createGroupNonUniformBitwiseReduceOpImpl(OpBuilder &builder, Location loc, Value arg, bool isGroup, bool isUniform, std::optional< uint32_t > clusterSize)
static spirv::FuncOp lowerAsEntryFunction(gpu::GPUFuncOp funcOp, const TypeConverter &typeConverter, ConversionPatternRewriter &rewriter, spirv::EntryPointABIAttr entryPointInfo, ArrayRef< spirv::InterfaceVarABIAttr > argABIInfo)
static std::string makeVarName(spirv::ModuleOp moduleOp, llvm::Twine prefix)
ArrayAttr()
b getContext())
Pattern to convert a gpu.all_reduce op into a SPIR-V group op.
LogicalResult matchAndRewrite(gpu::AllReduceOp op, OpAdaptor adaptor, ConversionPatternRewriter &rewriter) const override
Pattern to convert a gpu.subgroup_reduce op into a SPIR-V group op.
LogicalResult matchAndRewrite(gpu::SubgroupReduceOp op, OpAdaptor adaptor, ConversionPatternRewriter &rewriter) const override
Attributes are known-constant values of operations.
Definition Attributes.h:25
IntegerAttr getIntegerAttr(Type type, int64_t value)
Definition Builders.cpp:237
IntegerType getI32Type()
Definition Builders.cpp:71
MLIRContext * getContext() const
Definition Builders.h:56
static FlatSymbolRefAttr get(StringAttr value)
Construct a symbol reference for the given value name.
This class defines the main interface for locations in MLIR and acts as a non-nullable wrapper around...
Definition Location.h:76
MLIRContext is the top-level object for a collection of MLIR operations.
Definition MLIRContext.h:63
This class helps build Operations.
Definition Builders.h:210
Region & getRegion(unsigned index)
Returns the region held by this operation at position 'index'.
Definition Operation.h:738
Block & back()
Definition Region.h:64
iterator begin()
Definition Region.h:55
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.
Type conversion from builtin types to SPIR-V types for shader interface.
static Operation * getNearestSymbolTable(Operation *from)
Returns the nearest symbol table from a given operation from.
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
Definition Types.h:74
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 ArrayType get(Type elementType, unsigned elementCount)
An attribute that specifies the information regarding the interface variable: descriptor set,...
static NamedBarrierType get(MLIRContext *context)
static PointerType get(Type pointeeType, StorageClass storageClass)
ResourceLimitsAttr getResourceLimits() const
Returns the target resource limits.
A wrapper class around a spirv::TargetEnvAttr to provide query methods for allowed version/capabiliti...
TargetEnvAttr getAttr() const
detail::InFlightRemark failed(Location loc, RemarkOpts opts)
Report an optimization remark that failed.
Definition Remarks.h:732
StringRef getInterfaceVarABIAttrName()
Returns the attribute name for specifying argument ABI information.
bool needsInterfaceVarABIAttrs(TargetEnvAttr targetAttr)
Returns whether the given SPIR-V target (described by TargetEnvAttr) needs ABI attributes for interfa...
InterfaceVarABIAttr getInterfaceVarABIAttr(unsigned descriptorSet, unsigned binding, std::optional< StorageClass > storageClass, MLIRContext *context)
Gets the InterfaceVarABIAttr given its fields.
Value getBuiltinVariableValue(Operation *op, BuiltIn builtin, Type integerType, OpBuilder &builder, StringRef prefix="__builtin__", StringRef suffix="__")
Returns the value for the given builtin variable.
EntryPointABIAttr lookupEntryPointABI(Operation *op)
Queries the entry point ABI on the nearest function-like op containing the given op.
StringRef getTargetEnvAttrName()
Returns the attribute name for specifying SPIR-V target environment.
DenseI32ArrayAttr lookupLocalWorkGroupSize(Operation *op)
Queries the local workgroup size from entry point ABI on the nearest function-like op containing the ...
AddressingModel getAddressingModel(TargetEnvAttr targetAttr, bool use64bitAddress)
Returns addressing model selected based on target environment.
FailureOr< MemoryModel > getMemoryModel(TargetEnvAttr targetAttr)
Returns memory model selected based on target environment.
StringRef getEntryPointABIAttrName()
Returns the attribute name for specifying entry point information.
Include the generated interface declarations.
bool matchPattern(Value value, const Pattern &pattern)
Entry point for matching a pattern over a Value.
Definition Matchers.h:490
void populateGPUNamedBarrierToSPIRVTypeConversion(SPIRVTypeConverter &typeConverter)
Adds gpu::NamedBarrierType to spirv::NamedBarrierType conversion.
detail::DenseArrayAttrImpl< int32_t > DenseI32ArrayAttr
void populateGPUToSPIRVPatterns(const SPIRVTypeConverter &typeConverter, RewritePatternSet &patterns)
Appends to a pattern list additional patterns for translating GPU Ops to SPIR-V ops.
detail::constant_op_matcher m_Constant()
Matches a constant foldable operation.
Definition Matchers.h:369