llvm-project/mlir/test/Dialect/SPIRV/TestAvailability.cpp
Lei Zhang 267483ac70 [mlir][spirv] Support implied extensions and capabilities
In SPIR-V, when a new version is introduced, it is possible some
existing extensions will be incorporated into it so that it becomes
implicitly declared if targeting the new version. This affects
conversion target specification because we need to take this into
account when allowing what extensions to use.

For a capability, it may also implies some other capabilities,
for example, the `Shader` capability implies `Matrix` the capability.
This should also be taken into consideration when preparing the
conversion target: when we specify an capability is allowed, all
its recursively implied capabilities are also allowed.

This commit adds utility functions to query implied extensions for
a given version and implied capabilities for a given capability
and updated SPIRVConversionTarget to use them.

This commit also fixes a bug in availability spec. When a symbol
(op or enum case) can be enabled by an extension, we should drop
it's minimal version requirement. Being enabled by an extension
naturally means the symbol can be used by *any* SPIR-V version
as long as the extension is supported. The grammar still encodes
the 'version' field for such cases, but it should be interpreted
as a different way: rather than meaning a minimal version
requirement, it says the symbol becomes core at that specific
version.

Differential Revision: https://reviews.llvm.org/D72765
2020-01-17 08:01:57 -05:00

219 lines
7.9 KiB
C++

//===- TestAvailability.cpp - Pass to test SPIR-V op availability ---------===//
//
// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.
// See https://llvm.org/LICENSE.txt for license information.
// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
//
//===----------------------------------------------------------------------===//
#include "mlir/Dialect/SPIRV/SPIRVLowering.h"
#include "mlir/Dialect/SPIRV/SPIRVOps.h"
#include "mlir/Dialect/SPIRV/SPIRVTypes.h"
#include "mlir/IR/Function.h"
#include "mlir/Pass/Pass.h"
using namespace mlir;
//===----------------------------------------------------------------------===//
// Printing op availability pass
//===----------------------------------------------------------------------===//
namespace {
/// A pass for testing SPIR-V op availability.
struct PrintOpAvailability : public FunctionPass<PrintOpAvailability> {
void runOnFunction() override;
};
} // end anonymous namespace
void PrintOpAvailability::runOnFunction() {
auto f = getFunction();
llvm::outs() << f.getName() << "\n";
Dialect *spvDialect = getContext().getRegisteredDialect("spv");
f.getOperation()->walk([&](Operation *op) {
if (op->getDialect() != spvDialect)
return WalkResult::advance();
auto opName = op->getName();
auto &os = llvm::outs();
if (auto minVersion = dyn_cast<spirv::QueryMinVersionInterface>(op))
os << opName << " min version: "
<< spirv::stringifyVersion(minVersion.getMinVersion()) << "\n";
if (auto maxVersion = dyn_cast<spirv::QueryMaxVersionInterface>(op))
os << opName << " max version: "
<< spirv::stringifyVersion(maxVersion.getMaxVersion()) << "\n";
if (auto extension = dyn_cast<spirv::QueryExtensionInterface>(op)) {
os << opName << " extensions: [";
for (const auto &exts : extension.getExtensions()) {
os << " [";
interleaveComma(exts, os, [&](spirv::Extension ext) {
os << spirv::stringifyExtension(ext);
});
os << "]";
}
os << " ]\n";
}
if (auto capability = dyn_cast<spirv::QueryCapabilityInterface>(op)) {
os << opName << " capabilities: [";
for (const auto &caps : capability.getCapabilities()) {
os << " [";
interleaveComma(caps, os, [&](spirv::Capability cap) {
os << spirv::stringifyCapability(cap);
});
os << "]";
}
os << " ]\n";
}
os.flush();
return WalkResult::advance();
});
}
static PassRegistration<PrintOpAvailability>
printOpAvailabilityPass("test-spirv-op-availability",
"Test SPIR-V op availability");
//===----------------------------------------------------------------------===//
// Converting target environment pass
//===----------------------------------------------------------------------===//
namespace {
/// A pass for testing SPIR-V op availability.
struct ConvertToTargetEnv : public FunctionPass<ConvertToTargetEnv> {
void runOnFunction() override;
};
struct ConvertToAtomCmpExchangeWeak : public RewritePattern {
ConvertToAtomCmpExchangeWeak(MLIRContext *context);
PatternMatchResult matchAndRewrite(Operation *op,
PatternRewriter &rewriter) const override;
};
struct ConvertToBitReverse : public RewritePattern {
ConvertToBitReverse(MLIRContext *context);
PatternMatchResult matchAndRewrite(Operation *op,
PatternRewriter &rewriter) const override;
};
struct ConvertToGroupNonUniformBallot : public RewritePattern {
ConvertToGroupNonUniformBallot(MLIRContext *context);
PatternMatchResult matchAndRewrite(Operation *op,
PatternRewriter &rewriter) const override;
};
struct ConvertToModule : public RewritePattern {
ConvertToModule(MLIRContext *context);
PatternMatchResult matchAndRewrite(Operation *op,
PatternRewriter &rewriter) const override;
};
struct ConvertToSubgroupBallot : public RewritePattern {
ConvertToSubgroupBallot(MLIRContext *context);
PatternMatchResult matchAndRewrite(Operation *op,
PatternRewriter &rewriter) const override;
};
} // end anonymous namespace
void ConvertToTargetEnv::runOnFunction() {
MLIRContext *context = &getContext();
FuncOp fn = getFunction();
auto targetEnv = fn.getOperation()
->getAttr(spirv::getTargetEnvAttrName())
.cast<spirv::TargetEnvAttr>();
auto target = spirv::SPIRVConversionTarget::get(targetEnv, context);
OwningRewritePatternList patterns;
patterns.insert<ConvertToAtomCmpExchangeWeak, ConvertToBitReverse,
ConvertToGroupNonUniformBallot, ConvertToModule,
ConvertToSubgroupBallot>(context);
if (failed(applyPartialConversion(fn, *target, patterns)))
return signalPassFailure();
}
ConvertToAtomCmpExchangeWeak::ConvertToAtomCmpExchangeWeak(MLIRContext *context)
: RewritePattern("test.convert_to_atomic_compare_exchange_weak_op",
{"spv.AtomicCompareExchangeWeak"}, 1, context) {}
PatternMatchResult
ConvertToAtomCmpExchangeWeak::matchAndRewrite(Operation *op,
PatternRewriter &rewriter) const {
Value ptr = op->getOperand(0);
Value value = op->getOperand(1);
Value comparator = op->getOperand(2);
// Create a spv.AtomicCompareExchangeWeak op with AtomicCounterMemory bits in
// memory semantics to additionally require AtomicStorage capability.
rewriter.replaceOpWithNewOp<spirv::AtomicCompareExchangeWeakOp>(
op, value.getType(), ptr, spirv::Scope::Workgroup,
spirv::MemorySemantics::AcquireRelease |
spirv::MemorySemantics::AtomicCounterMemory,
spirv::MemorySemantics::Acquire, value, comparator);
return matchSuccess();
}
ConvertToBitReverse::ConvertToBitReverse(MLIRContext *context)
: RewritePattern("test.convert_to_bit_reverse_op", {"spv.BitReverse"}, 1,
context) {}
PatternMatchResult
ConvertToBitReverse::matchAndRewrite(Operation *op,
PatternRewriter &rewriter) const {
Value predicate = op->getOperand(0);
rewriter.replaceOpWithNewOp<spirv::BitReverseOp>(
op, op->getResult(0).getType(), predicate);
return matchSuccess();
}
ConvertToGroupNonUniformBallot::ConvertToGroupNonUniformBallot(
MLIRContext *context)
: RewritePattern("test.convert_to_group_non_uniform_ballot_op",
{"spv.GroupNonUniformBallot"}, 1, context) {}
PatternMatchResult ConvertToGroupNonUniformBallot::matchAndRewrite(
Operation *op, PatternRewriter &rewriter) const {
Value predicate = op->getOperand(0);
rewriter.replaceOpWithNewOp<spirv::GroupNonUniformBallotOp>(
op, op->getResult(0).getType(), spirv::Scope::Workgroup, predicate);
return matchSuccess();
}
ConvertToModule::ConvertToModule(MLIRContext *context)
: RewritePattern("test.convert_to_module_op", {"spv.module"}, 1, context) {}
PatternMatchResult
ConvertToModule::matchAndRewrite(Operation *op,
PatternRewriter &rewriter) const {
rewriter.replaceOpWithNewOp<spirv::ModuleOp>(
op, spirv::AddressingModel::PhysicalStorageBuffer64,
spirv::MemoryModel::Vulkan);
return matchSuccess();
}
ConvertToSubgroupBallot::ConvertToSubgroupBallot(MLIRContext *context)
: RewritePattern("test.convert_to_subgroup_ballot_op",
{"spv.SubgroupBallotKHR"}, 1, context) {}
PatternMatchResult
ConvertToSubgroupBallot::matchAndRewrite(Operation *op,
PatternRewriter &rewriter) const {
Value predicate = op->getOperand(0);
rewriter.replaceOpWithNewOp<spirv::SubgroupBallotKHROp>(
op, op->getResult(0).getType(), predicate);
return matchSuccess();
}
static PassRegistration<ConvertToTargetEnv>
convertToTargetEnvPass("test-spirv-target-env",
"Test SPIR-V target environment");