32template <
typename SourceOp, spirv::BuiltIn builtin>
33class LaunchConfigConversion :
public OpConversionPattern<SourceOp> {
35 using OpConversionPattern<SourceOp>::OpConversionPattern;
38 matchAndRewrite(SourceOp op,
typename SourceOp::Adaptor adaptor,
39 ConversionPatternRewriter &rewriter)
const override;
44template <
typename SourceOp, spirv::BuiltIn builtin>
45class SingleDimLaunchConfigConversion :
public OpConversionPattern<SourceOp> {
47 using OpConversionPattern<SourceOp>::OpConversionPattern;
50 matchAndRewrite(SourceOp op,
typename SourceOp::Adaptor adaptor,
51 ConversionPatternRewriter &rewriter)
const override;
58class WorkGroupSizeConversion :
public OpConversionPattern<gpu::BlockDimOp> {
62 : OpConversionPattern(typeConverter, context, 10) {}
65 matchAndRewrite(gpu::BlockDimOp op, OpAdaptor adaptor,
66 ConversionPatternRewriter &rewriter)
const override;
70class GPUFuncOpConversion final :
public OpConversionPattern<gpu::GPUFuncOp> {
75 matchAndRewrite(gpu::GPUFuncOp funcOp, OpAdaptor adaptor,
76 ConversionPatternRewriter &rewriter)
const override;
83class GPUModuleConversion final :
public OpConversionPattern<gpu::GPUModuleOp> {
88 matchAndRewrite(gpu::GPUModuleOp moduleOp, OpAdaptor adaptor,
89 ConversionPatternRewriter &rewriter)
const override;
94class GPUReturnOpConversion final :
public OpConversionPattern<gpu::ReturnOp> {
99 matchAndRewrite(gpu::ReturnOp returnOp, OpAdaptor adaptor,
100 ConversionPatternRewriter &rewriter)
const override;
105class GPUBarrierConversion final :
public OpConversionPattern<gpu::BarrierOp> {
110 matchAndRewrite(gpu::BarrierOp barrierOp, OpAdaptor adaptor,
111 ConversionPatternRewriter &rewriter)
const override;
116class GPUInitializeNamedBarrierConversion final
117 :
public OpConversionPattern<gpu::InitializeNamedBarrierOp> {
122 matchAndRewrite(gpu::InitializeNamedBarrierOp op, OpAdaptor adaptor,
123 ConversionPatternRewriter &rewriter)
const override;
127class GPUShuffleConversion final :
public OpConversionPattern<gpu::ShuffleOp> {
132 matchAndRewrite(gpu::ShuffleOp shuffleOp, OpAdaptor adaptor,
133 ConversionPatternRewriter &rewriter)
const override;
137class GPURotateConversion final :
public OpConversionPattern<gpu::RotateOp> {
142 matchAndRewrite(gpu::RotateOp rotateOp, OpAdaptor adaptor,
143 ConversionPatternRewriter &rewriter)
const override;
148class GPUSubgroupBroadcastConversion final
149 :
public OpConversionPattern<gpu::SubgroupBroadcastOp> {
154 matchAndRewrite(gpu::SubgroupBroadcastOp op, OpAdaptor adaptor,
155 ConversionPatternRewriter &rewriter)
const override;
158class GPUBallotConversion final :
public OpConversionPattern<gpu::BallotOp> {
163 matchAndRewrite(gpu::BallotOp ballotOp, OpAdaptor adaptor,
164 ConversionPatternRewriter &rewriter)
const override;
167class GPUPrintfConversion final :
public OpConversionPattern<gpu::PrintfOp> {
172 matchAndRewrite(gpu::PrintfOp gpuPrintfOp, OpAdaptor adaptor,
173 ConversionPatternRewriter &rewriter)
const override;
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();
201 typeConverter->getTargetEnv().allows(spirv::Capability::Shader);
202 Type builtinType = forShader ? rewriter.getIntegerType(32) : indexType;
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);
215template <
typename SourceOp, spirv::BuiltIn builtin>
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);
234 if (i32Type != indexType)
235 builtinValue = spirv::UConvertOp::create(rewriter, op.getLoc(), indexType,
237 rewriter.replaceOp(op, builtinValue);
241LogicalResult WorkGroupSizeConversion::matchAndRewrite(
242 gpu::BlockDimOp op, OpAdaptor adaptor,
243 ConversionPatternRewriter &rewriter)
const {
245 if (!workGroupSizeAttr)
249 workGroupSizeAttr.
asArrayRef()[
static_cast<int32_t
>(op.getDimension())];
251 getTypeConverter()->convertType(op.getResult().getType());
254 rewriter.replaceOpWithNewOp<spirv::ConstantOp>(
255 op, convertedType, IntegerAttr::get(convertedType, val));
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");
275 if (!argABIInfo.empty() && fnType.getNumInputs() != argABIInfo.size()) {
277 "lowering as entry functions requires ABI info for all arguments "
284 TypeConverter::SignatureConversion signatureConverter(fnType.getNumInputs());
286 for (
const auto &argType :
287 enumerate(funcOp.getFunctionType().getInputs())) {
288 auto convertedType = typeConverter.convertType(argType.value());
291 signatureConverter.addInputs(argType.index(), convertedType);
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())
301 cast<SymbolOpInterface>(funcOp.getOperation()).getVisibility());
303 auto copyGPUProperty = [&](StringAttr name,
Attribute value) {
305 newFuncOp->setDiscardableAttr(name, value);
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());
324 rewriter.inlineRegionBefore(funcOp.getBody(), newFuncOp.getBody(),
326 if (failed(rewriter.convertRegionTypes(&newFuncOp.getBody(), typeConverter,
327 &signatureConverter)))
329 rewriter.eraseOp(funcOp);
333 for (
auto argIndex : llvm::seq<unsigned>(0, argABIInfo.size())) {
334 newFuncOp.setArgAttr(argIndex, argABIAttrName, argABIInfo[argIndex]);
351 for (
auto argIndex : llvm::seq<unsigned>(0, funcOp.getNumArguments())) {
357 std::optional<spirv::StorageClass> sc;
358 if (funcOp.getArgument(argIndex).getType().isIntOrIndexOrFloat())
359 sc = spirv::StorageClass::StorageBuffer;
366LogicalResult GPUFuncOpConversion::matchAndRewrite(
367 gpu::GPUFuncOp funcOp, OpAdaptor adaptor,
368 ConversionPatternRewriter &rewriter)
const {
369 if (!gpu::GPUDialect::isKernel(funcOp))
372 auto *typeConverter = getTypeConverter<SPIRVTypeConverter>();
373 SmallVector<spirv::InterfaceVarABIAttr, 4> argABI;
377 for (
auto argIndex : llvm::seq<unsigned>(0, funcOp.getNumArguments())) {
379 auto abiAttr = funcOp.getArgAttrOfType<spirv::InterfaceVarABIAttr>(
383 "match failure: missing 'spirv.interface_var_abi' attribute at "
388 argABI.push_back(abiAttr);
393 if (!entryPointAttr) {
395 "match failure: missing 'spirv.entry_point_abi' attribute");
399 funcOp, *getTypeConverter(), rewriter, entryPointAttr, argABI);
402 newFuncOp->removeDiscardableAttr(
403 rewriter.getStringAttr(gpu::GPUDialect::getKernelFuncAttrName()));
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();
417 targetEnv, typeConverter->getOptions().use64bitIndex);
420 return moduleOp.emitRemark(
421 "cannot deduce memory model from 'spirv.target_env'");
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));
430 Region &spvModuleRegion = spvModule.getRegion();
431 rewriter.inlineRegionBefore(moduleOp.getBodyRegion(), spvModuleRegion,
432 spvModuleRegion.
begin());
434 rewriter.eraseBlock(&spvModuleRegion.
back());
440 if (
auto attr = moduleOp->getDiscardableAttrOfType<spirv::TargetEnvAttr>(
443 if (
ArrayAttr targets = moduleOp.getTargetsAttr()) {
444 for (Attribute targetAttr : targets)
445 if (
auto spirvTargetEnvAttr =
446 dyn_cast<spirv::TargetEnvAttr>(targetAttr)) {
453 rewriter.eraseOp(moduleOp);
461LogicalResult GPUReturnOpConversion::matchAndRewrite(
462 gpu::ReturnOp returnOp, OpAdaptor adaptor,
463 ConversionPatternRewriter &rewriter)
const {
464 if (!adaptor.getOperands().empty())
467 rewriter.replaceOpWithNewOp<spirv::ReturnOp>(returnOp);
476static FailureOr<spirv::Scope>
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:
489LogicalResult GPUBarrierConversion::matchAndRewrite(
490 gpu::BarrierOp barrierOp, OpAdaptor adaptor,
491 ConversionPatternRewriter &rewriter)
const {
497 return rewriter.notifyMatchFailure(
498 barrierOp,
"cluster scope is not supported in SPIR-V");
500 auto scopeAttr = spirv::ScopeAttr::get(context, *spirvScope);
501 auto memoryScopeAttr =
502 spirv::ScopeAttr::get(context, spirv::Scope::Workgroup);
505 auto memorySemantics = spirv::MemorySemanticsAttr::get(
506 context, spirv::MemorySemantics::WorkgroupMemory |
507 spirv::MemorySemantics::AcquireRelease);
509 if (adaptor.getNamedBarrier()) {
510 spirv::MemoryNamedBarrierOp::create(rewriter, barrierOp.getLoc(),
511 adaptor.getNamedBarrier(),
512 memoryScopeAttr, memorySemantics);
513 rewriter.eraseOp(barrierOp);
515 rewriter.replaceOpWithNewOp<spirv::ControlBarrierOp>(
516 barrierOp, scopeAttr, memoryScopeAttr, memorySemantics);
521LogicalResult GPUInitializeNamedBarrierConversion::matchAndRewrite(
522 gpu::InitializeNamedBarrierOp op, OpAdaptor adaptor,
523 ConversionPatternRewriter &rewriter)
const {
525 rewriter.replaceOpWithNewOp<spirv::NamedBarrierInitializeOp>(
526 op, nbType, adaptor.getMemberCount());
534LogicalResult GPUShuffleConversion::matchAndRewrite(
535 gpu::ShuffleOp shuffleOp, OpAdaptor adaptor,
536 ConversionPatternRewriter &rewriter)
const {
540 const spirv::TargetEnv &targetEnv =
541 getTypeConverter<SPIRVTypeConverter>()->getTargetEnv();
542 unsigned subgroupSize =
544 IntegerAttr widthAttr;
546 widthAttr.getValue().getZExtValue() != subgroupSize)
547 return rewriter.notifyMatchFailure(
548 shuffleOp,
"shuffle width and target subgroup size mismatch");
550 assert(!adaptor.getOffset().getType().isSignedInteger() &&
551 "shuffle offset must be a signless/unsigned integer");
553 Location loc = shuffleOp.getLoc();
554 auto scope = rewriter.getAttr<spirv::ScopeAttr>(spirv::Scope::Subgroup);
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);
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);
573 case gpu::ShuffleMode::DOWN: {
574 result = spirv::GroupNonUniformShuffleDownOp::create(
575 rewriter, loc, scope, adaptor.getValue(), adaptor.getOffset());
577 Value laneId = gpu::LaneIdOp::create(rewriter, loc, widthAttr);
579 arith::AddIOp::create(rewriter, loc, laneId, adaptor.getOffset());
580 validVal = arith::CmpIOp::create(rewriter, loc, arith::CmpIPredicate::ult,
581 resultLaneId, adaptor.getWidth());
584 case gpu::ShuffleMode::UP: {
585 result = spirv::GroupNonUniformShuffleUpOp::create(
586 rewriter, loc, scope, adaptor.getValue(), adaptor.getOffset());
588 Value laneId = gpu::LaneIdOp::create(rewriter, loc, widthAttr);
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)));
600 rewriter.replaceOp(shuffleOp, {
result, validVal});
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 =
615 unsigned width = rotateOp.getWidth();
616 if (width > subgroupSize)
617 return rewriter.notifyMatchFailure(
618 rotateOp,
"rotate width is larger than target subgroup size");
620 Location loc = rotateOp.getLoc();
621 auto scope = rewriter.getAttr<spirv::ScopeAttr>(spirv::Scope::Subgroup);
623 arith::ConstantOp::create(rewriter, loc, adaptor.getOffsetAttr());
625 arith::ConstantOp::create(rewriter, loc, adaptor.getWidthAttr());
626 Value rotateResult = spirv::GroupNonUniformRotateKHROp::create(
627 rewriter, loc, scope, adaptor.getValue(), offsetVal, widthVal);
629 if (width == subgroupSize) {
630 validVal = spirv::ConstantOp::getOne(rewriter.getI1Type(), loc, rewriter);
632 IntegerAttr widthAttr = adaptor.getWidthAttr();
633 Value laneId = gpu::LaneIdOp::create(rewriter, loc, widthAttr);
634 validVal = arith::CmpIOp::create(rewriter, loc, arith::CmpIPredicate::ult,
638 rewriter.replaceOp(rotateOp, {rotateResult, validVal});
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);
653 switch (op.getBroadcastType()) {
654 case gpu::BroadcastType::specific_lane:
655 result = spirv::GroupNonUniformBroadcastOp::create(
656 rewriter, loc, scope, adaptor.getSrc(), adaptor.getLane());
658 case gpu::BroadcastType::first_active_lane:
659 result = spirv::GroupNonUniformBroadcastFirstOp::create(
660 rewriter, loc, scope, adaptor.getSrc());
664 rewriter.replaceOp(op,
result);
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);
677 Value ballot = spirv::GroupNonUniformBallotOp::create(
678 rewriter, loc, vec4i32Type, scope, adaptor.getPredicate());
680 auto intType = cast<IntegerType>(ballotOp.getType());
681 unsigned width = intType.getWidth();
685 spirv::CompositeExtractOp::create(rewriter, loc, ballot, {0});
686 rewriter.replaceOp(ballotOp,
result);
687 }
else if (width == 64) {
689 Value low = spirv::CompositeExtractOp::create(rewriter, loc, ballot, {0});
690 Value high = spirv::CompositeExtractOp::create(rewriter, loc, ballot, {1});
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);
696 Value shift32 = spirv::ConstantOp::create(
697 rewriter, loc, int64Type, rewriter.getIntegerAttr(int64Type, 32));
699 spirv::ShiftLeftLogicalOp::create(rewriter, loc, highExt, shift32);
702 spirv::BitwiseOrOp::create(rewriter, loc, lowExt, highShifted);
703 rewriter.replaceOp(ballotOp,
result);
705 return rewriter.notifyMatchFailure(
706 ballotOp,
"only i32 and i64 result types are supported for SPIR-V");
716template <
typename UniformOp,
typename NonUniformOp>
718 Value arg,
bool isGroup,
bool isUniform,
719 std::optional<uint32_t> clusterSize) {
721 isGroup ? spirv::Scope::Workgroup : spirv::Scope::Subgroup;
723 if (!isUniform && scope != spirv::Scope::Subgroup)
727 auto scopeAttr = mlir::spirv::ScopeAttr::get(builder.
getContext(), scope);
728 auto groupOp = spirv::GroupOperationAttr::get(
730 ? spirv::GroupOperation::ClusteredReduce
731 : spirv::GroupOperation::Reduce);
733 return UniformOp::create(builder, loc, type, scopeAttr, groupOp, arg)
737 Value clusterSizeValue;
738 if (clusterSize.has_value())
739 clusterSizeValue = spirv::ConstantOp::create(
743 return NonUniformOp::create(builder, loc, type, scopeAttr, groupOp, arg,
748template <
typename NonUniformOp>
751 std::optional<uint32_t> clusterSize) {
753 isGroup ? spirv::Scope::Workgroup : spirv::Scope::Subgroup;
754 if (isUniform || scope != spirv::Scope::Subgroup)
758 auto scopeAttr = mlir::spirv::ScopeAttr::get(builder.
getContext(), scope);
759 auto groupOp = spirv::GroupOperationAttr::get(
761 ? spirv::GroupOperation::ClusteredReduce
762 : spirv::GroupOperation::Reduce);
764 Value clusterSizeValue;
765 if (clusterSize.has_value())
766 clusterSizeValue = spirv::ConstantOp::create(
770 return NonUniformOp::create(builder, loc, type, scopeAttr, groupOp, arg,
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 };
781 std::optional<uint32_t>);
783 gpu::AllReduceOperation kind;
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
804 using ReduceType = gpu::AllReduceOperation;
805 const OpHandler handlers[] = {
806 {ReduceType::ADD, ElemType::Integer,
808 spirv::GroupNonUniformIAddOp>},
809 {ReduceType::ADD, ElemType::Float,
811 spirv::GroupNonUniformFAddOp>},
812 {ReduceType::MUL, ElemType::Integer,
814 spirv::GroupNonUniformIMulOp>},
815 {ReduceType::MUL, ElemType::Float,
817 spirv::GroupNonUniformFMulOp>},
818 {ReduceType::MINUI, ElemType::Integer,
820 spirv::GroupNonUniformUMinOp>},
821 {ReduceType::MINSI, ElemType::Integer,
823 spirv::GroupNonUniformSMinOp>},
824 {ReduceType::MINNUMF, ElemType::Float,
826 spirv::GroupNonUniformFMinOp>},
827 {ReduceType::MAXUI, ElemType::Integer,
829 spirv::GroupNonUniformUMaxOp>},
830 {ReduceType::MAXSI, ElemType::Integer,
832 spirv::GroupNonUniformSMaxOp>},
833 {ReduceType::MAXNUMF, ElemType::Float,
835 spirv::GroupNonUniformFMaxOp>},
836 {ReduceType::MINIMUMF, ElemType::Float,
838 spirv::GroupNonUniformFMinOp>},
839 {ReduceType::MAXIMUMF, ElemType::Float,
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>}};
852 for (
const OpHandler &handler : handlers)
853 if (handler.kind == opType && elementType == handler.elemType)
855 handler.func(builder, loc, arg, isGroup, isUniform, clusterSize))
863 :
public OpConversionPattern<gpu::AllReduceOp> {
869 ConversionPatternRewriter &rewriter)
const override {
870 auto opType = op.getOp();
879 true, op.getUniform(), std::nullopt);
883 rewriter.replaceOp(op, *
result);
890 :
public OpConversionPattern<gpu::SubgroupReduceOp> {
896 ConversionPatternRewriter &rewriter)
const override {
897 if (op.getClusterStride() > 1) {
898 return rewriter.notifyMatchFailure(
899 op,
"lowering for cluster stride > 1 is not implemented");
902 if (!isa<spirv::ScalarType>(adaptor.getValue().getType()))
903 return rewriter.notifyMatchFailure(op,
"reduction type is not a scalar");
906 rewriter, op.getLoc(), adaptor.getValue(), adaptor.getOp(),
907 false, adaptor.getUniform(), op.getClusterSize());
911 rewriter.replaceOp(op, *
result);
920static std::string
makeVarName(spirv::ModuleOp moduleOp, llvm::Twine prefix) {
926 name = (prefix + llvm::Twine(number++)).str();
927 }
while (moduleOp.lookupSymbol(name));
934LogicalResult GPUPrintfConversion::matchAndRewrite(
935 gpu::PrintfOp gpuPrintfOp, OpAdaptor adaptor,
936 ConversionPatternRewriter &rewriter)
const {
938 Location loc = gpuPrintfOp.getLoc();
940 auto moduleOp = gpuPrintfOp->getParentOfType<spirv::ModuleOp>();
947 std::string globalVarName =
makeVarName(moduleOp, llvm::Twine(
"printfMsg"));
948 spirv::GlobalVariableOp globalVar;
950 IntegerType i8Type = rewriter.getI8Type();
951 IntegerType i32Type = rewriter.getI32Type();
958 auto createSpecConstant = [&](
unsigned value) {
959 auto attr = rewriter.getI8IntegerAttr(value);
960 std::string specCstName =
961 makeVarName(moduleOp, llvm::Twine(globalVarName) +
"_sc");
963 return spirv::SpecConstantOp::create(
964 rewriter, loc, rewriter.getStringAttr(specCstName), attr,
971 ConversionPatternRewriter::InsertionGuard guard(rewriter);
974 rewriter.setInsertionPointToStart(
981 llvm::SmallString<20> formatString(adaptor.getFormat());
982 formatString.push_back(
'\0');
983 SmallVector<Attribute, 4> constituents;
984 for (
char c : formatString) {
985 spirv::SpecConstantOp cSpecConstantOp = createSpecConstant(c);
986 constituents.push_back(SymbolRefAttr::get(cSpecConstantOp));
990 size_t contentSize = constituents.size();
992 spirv::SpecConstantCompositeOp specCstComposite;
995 std::string specCstCompositeName =
996 (llvm::Twine(globalVarName) +
"_scc").str();
998 specCstComposite = spirv::SpecConstantCompositeOp::create(
999 rewriter, loc, TypeAttr::get(globalType),
1000 rewriter.getStringAttr(specCstCompositeName),
1001 rewriter.getArrayAttr(constituents),
nullptr);
1004 globalType, spirv::StorageClass::UniformConstant);
1009 globalVar = spirv::GlobalVariableOp::create(
1010 rewriter, loc, ptrType, globalVarName,
1013 globalVar->setDiscardableAttr(
"Constant", rewriter.getUnitAttr());
1017 Value globalPtr = spirv::AddressOfOp::create(rewriter, loc, globalVar);
1018 Value fmtStr = spirv::BitcastOp::create(
1024 auto printfArgs = llvm::to_vector_of<Value, 4>(adaptor.getArgs());
1026 spirv::CLPrintfOp::create(rewriter, loc, i32Type, fmtStr, printfArgs);
1031 rewriter.eraseOp(gpuPrintfOp);
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>,
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)
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.
IntegerAttr getIntegerAttr(Type type, int64_t value)
MLIRContext * getContext() const
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...
MLIRContext is the top-level object for a collection of MLIR operations.
This class helps build Operations.
Region & getRegion(unsigned index)
Returns the region held by this operation at position 'index'.
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...
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
Type getType() const
Return the type of this value.
ArrayRef< T > asArrayRef() const
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
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.
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.