262 lines · cpp
1//===- TestAvailability.cpp - Pass to test SPIR-V op availability ---------===//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-exception6//7//===----------------------------------------------------------------------===//8 9#include "mlir/Dialect/Func/IR/FuncOps.h"10#include "mlir/Dialect/SPIRV/IR/SPIRVAttributes.h"11#include "mlir/Dialect/SPIRV/IR/SPIRVOps.h"12#include "mlir/Dialect/SPIRV/Transforms/SPIRVConversion.h"13#include "mlir/Pass/Pass.h"14 15using namespace mlir;16 17//===----------------------------------------------------------------------===//18// Printing op availability pass19//===----------------------------------------------------------------------===//20 21namespace {22/// A pass for testing SPIR-V op availability.23struct PrintOpAvailability24 : public PassWrapper<PrintOpAvailability, OperationPass<mlir::ModuleOp>> {25 MLIR_DEFINE_EXPLICIT_INTERNAL_INLINE_TYPE_ID(PrintOpAvailability)26 27 void runOnOperation() override;28 StringRef getArgument() const final { return "test-spirv-op-availability"; }29 StringRef getDescription() const final {30 return "Test SPIR-V op availability";31 }32};33} // namespace34 35void PrintOpAvailability::runOnOperation() {36 mlir::ModuleOp moduleOp = getOperation();37 Dialect *spirvDialect = getContext().getLoadedDialect("spirv");38 39 auto opCallback = [&](Operation *op) {40 if (op->getDialect() != spirvDialect)41 return WalkResult::advance();42 43 auto opName = op->getName();44 auto &os = llvm::outs();45 46 if (auto minVersionIfx = dyn_cast<spirv::QueryMinVersionInterface>(op)) {47 std::optional<spirv::Version> minVersion = minVersionIfx.getMinVersion();48 os << opName << " min version: ";49 if (minVersion)50 os << spirv::stringifyVersion(*minVersion) << "\n";51 else52 os << "None\n";53 }54 55 if (auto maxVersionIfx = dyn_cast<spirv::QueryMaxVersionInterface>(op)) {56 std::optional<spirv::Version> maxVersion = maxVersionIfx.getMaxVersion();57 os << opName << " max version: ";58 if (maxVersion)59 os << spirv::stringifyVersion(*maxVersion) << "\n";60 else61 os << "None\n";62 }63 64 if (auto extension = dyn_cast<spirv::QueryExtensionInterface>(op)) {65 os << opName << " extensions: [";66 for (const auto &exts : extension.getExtensions()) {67 os << " [";68 llvm::interleaveComma(exts, os, [&](spirv::Extension ext) {69 os << spirv::stringifyExtension(ext);70 });71 os << "]";72 }73 os << " ]\n";74 }75 76 if (auto capability = dyn_cast<spirv::QueryCapabilityInterface>(op)) {77 os << opName << " capabilities: [";78 for (const auto &caps : capability.getCapabilities()) {79 os << " [";80 llvm::interleaveComma(caps, os, [&](spirv::Capability cap) {81 os << spirv::stringifyCapability(cap);82 });83 os << "]";84 }85 os << " ]\n";86 }87 os.flush();88 89 return WalkResult::advance();90 };91 92 moduleOp.walk([&](func::FuncOp f) {93 llvm::outs() << f.getName() << "\n";94 f->walk(opCallback);95 });96 97 moduleOp.walk([&](spirv::GraphARMOp g) {98 llvm::outs() << g.getName() << "\n";99 g->walk(opCallback);100 });101}102 103namespace mlir {104void registerPrintSpirvAvailabilityPass() {105 PassRegistration<PrintOpAvailability>();106}107} // namespace mlir108 109//===----------------------------------------------------------------------===//110// Converting target environment pass111//===----------------------------------------------------------------------===//112 113namespace {114/// A pass for testing SPIR-V op availability.115struct ConvertToTargetEnv116 : public PassWrapper<ConvertToTargetEnv, OperationPass<func::FuncOp>> {117 MLIR_DEFINE_EXPLICIT_INTERNAL_INLINE_TYPE_ID(ConvertToTargetEnv)118 119 StringRef getArgument() const override { return "test-spirv-target-env"; }120 StringRef getDescription() const override {121 return "Test SPIR-V target environment";122 }123 void runOnOperation() override;124};125 126struct ConvertToAtomCmpExchangeWeak : RewritePattern {127 ConvertToAtomCmpExchangeWeak(MLIRContext *context)128 : RewritePattern("test.convert_to_atomic_compare_exchange_weak_op", 1,129 context, {"spirv.AtomicCompareExchangeWeak"}) {}130 131 LogicalResult matchAndRewrite(Operation *op,132 PatternRewriter &rewriter) const override {133 Value ptr = op->getOperand(0);134 Value value = op->getOperand(1);135 Value comparator = op->getOperand(2);136 137 // Create a spirv.AtomicCompareExchangeWeak op with AtomicCounterMemory bits138 // in memory semantics to additionally require AtomicStorage capability.139 rewriter.replaceOpWithNewOp<spirv::AtomicCompareExchangeWeakOp>(140 op, value.getType(), ptr, spirv::Scope::Workgroup,141 spirv::MemorySemantics::AcquireRelease |142 spirv::MemorySemantics::AtomicCounterMemory,143 spirv::MemorySemantics::Acquire, value, comparator);144 return success();145 }146};147 148struct ConvertToBitReverse : RewritePattern {149 ConvertToBitReverse(MLIRContext *context)150 : RewritePattern("test.convert_to_bit_reverse_op", 1, context,151 {"spirv.BitReverse"}) {}152 153 LogicalResult matchAndRewrite(Operation *op,154 PatternRewriter &rewriter) const override {155 Value predicate = op->getOperand(0);156 rewriter.replaceOpWithNewOp<spirv::BitReverseOp>(157 op, op->getResult(0).getType(), predicate);158 return success();159 }160};161 162struct ConvertToGroupNonUniformBallot : RewritePattern {163 ConvertToGroupNonUniformBallot(MLIRContext *context)164 : RewritePattern("test.convert_to_group_non_uniform_ballot_op", 1,165 context, {"spirv.GroupNonUniformBallot"}) {}166 167 LogicalResult matchAndRewrite(Operation *op,168 PatternRewriter &rewriter) const override {169 Value predicate = op->getOperand(0);170 rewriter.replaceOpWithNewOp<spirv::GroupNonUniformBallotOp>(171 op, op->getResult(0).getType(), spirv::Scope::Workgroup, predicate);172 return success();173 }174};175 176struct ConvertToModule : RewritePattern {177 ConvertToModule(MLIRContext *context)178 : RewritePattern("test.convert_to_module_op", 1, context,179 {"spirv.module"}) {}180 181 LogicalResult matchAndRewrite(Operation *op,182 PatternRewriter &rewriter) const override {183 rewriter.replaceOpWithNewOp<spirv::ModuleOp>(184 op, spirv::AddressingModel::PhysicalStorageBuffer64,185 spirv::MemoryModel::Vulkan);186 return success();187 }188};189 190struct ConvertToSubgroupBallot : RewritePattern {191 ConvertToSubgroupBallot(MLIRContext *context)192 : RewritePattern("test.convert_to_subgroup_ballot_op", 1, context,193 {"spirv.KHR.SubgroupBallot"}) {}194 195 LogicalResult matchAndRewrite(Operation *op,196 PatternRewriter &rewriter) const override {197 Value predicate = op->getOperand(0);198 rewriter.replaceOpWithNewOp<spirv::KHRSubgroupBallotOp>(199 op, op->getResult(0).getType(), predicate);200 return success();201 }202};203 204template <const char *TestOpName, typename SPIRVOp>205struct ConvertToIntegerDotProd : RewritePattern {206 ConvertToIntegerDotProd(MLIRContext *context)207 : RewritePattern(TestOpName, 1, context, {SPIRVOp::getOperationName()}) {}208 209 LogicalResult matchAndRewrite(Operation *op,210 PatternRewriter &rewriter) const override {211 rewriter.replaceOpWithNewOp<SPIRVOp>(op, op->getResultTypes(),212 op->getOperands(), op->getAttrs());213 return success();214 }215};216} // namespace217 218void ConvertToTargetEnv::runOnOperation() {219 MLIRContext *context = &getContext();220 func::FuncOp fn = getOperation();221 222 auto targetEnv = dyn_cast_or_null<spirv::TargetEnvAttr>(223 fn.getOperation()->getDiscardableAttr(spirv::getTargetEnvAttrName()));224 if (!targetEnv) {225 fn.emitError("missing 'spirv.target_env' attribute");226 return signalPassFailure();227 }228 229 auto target = SPIRVConversionTarget::get(targetEnv);230 231 static constexpr char sDotTestOpName[] = "test.convert_to_sdot_op";232 static constexpr char suDotTestOpName[] = "test.convert_to_sudot_op";233 static constexpr char uDotTestOpName[] = "test.convert_to_udot_op";234 static constexpr char sDotAccSatTestOpName[] =235 "test.convert_to_sdot_acc_sat_op";236 static constexpr char suDotAccSatTestOpName[] =237 "test.convert_to_sudot_acc_sat_op";238 static constexpr char uDotAccSatTestOpName[] =239 "test.convert_to_udot_acc_sat_op";240 241 RewritePatternSet patterns(context);242 patterns.add<243 ConvertToAtomCmpExchangeWeak, ConvertToBitReverse,244 ConvertToGroupNonUniformBallot, ConvertToModule, ConvertToSubgroupBallot,245 ConvertToIntegerDotProd<sDotTestOpName, spirv::SDotOp>,246 ConvertToIntegerDotProd<suDotTestOpName, spirv::SUDotOp>,247 ConvertToIntegerDotProd<uDotTestOpName, spirv::UDotOp>,248 ConvertToIntegerDotProd<sDotAccSatTestOpName, spirv::SDotAccSatOp>,249 ConvertToIntegerDotProd<suDotAccSatTestOpName, spirv::SUDotAccSatOp>,250 ConvertToIntegerDotProd<uDotAccSatTestOpName, spirv::UDotAccSatOp>>(251 context);252 253 if (failed(applyPartialConversion(fn, *target, std::move(patterns))))254 return signalPassFailure();255}256 257namespace mlir {258void registerConvertToTargetEnvPass() {259 PassRegistration<ConvertToTargetEnv>();260}261} // namespace mlir262