brintos

brintos / llvm-project-archived public Read only

0
0
Text · 9.3 KiB · 5643a0f Raw
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