21#include "llvm/ADT/StringExtras.h"
26#define GEN_PASS_DEF_SPIRVUPDATEVCEPASS
27#include "mlir/Dialect/SPIRV/Transforms/Passes.h.inc"
36class UpdateVCEPass final
37 :
public spirv::impl::SPIRVUpdateVCEPassBase<UpdateVCEPass> {
38 void runOnOperation()
override;
53 for (
const auto &ors : candidates) {
54 if (std::optional<spirv::Extension> chosen = targetEnv.
allows(ors)) {
55 deducedExtensions.insert(*chosen);
58 for (spirv::Extension ext : ors)
59 extStrings.push_back(spirv::stringifyExtension(ext));
62 << op->
getName() <<
"' requires at least one extension in ["
63 << llvm::join(extStrings,
", ")
64 <<
"] but none allowed in target environment";
81 for (
const auto &ors : candidates) {
82 if (std::optional<spirv::Capability> chosen = targetEnv.
allows(ors)) {
83 deducedCapabilities.insert(*chosen);
86 for (spirv::Capability cap : ors)
87 capStrings.push_back(spirv::stringifyCapability(cap));
90 << op->
getName() <<
"' requires at least one capability in ["
91 << llvm::join(capStrings,
", ")
92 <<
"] but none allowed in target environment";
100 for (spirv::Capability cap : caps)
101 tmp.insert_range(getRecursiveImpliedCapabilities(cap));
102 caps.insert_range(std::move(tmp));
105void UpdateVCEPass::runOnOperation() {
106 spirv::ModuleOp module = getOperation();
110 module.emitError("missing 'spirv.target_env' attribute");
111 return signalPassFailure();
114 spirv::TargetEnv targetEnv(targetAttr);
115 spirv::Version allowedVersion = targetAttr.
getVersion();
117 spirv::Version deducedVersion = spirv::Version::V_1_0;
123 WalkResult walkResult =
module.walk([&](Operation *op) -> WalkResult {
125 if (auto minVersionIfx = dyn_cast<spirv::QueryMinVersionInterface>(op)) {
126 std::optional<spirv::Version> minVersion = minVersionIfx.getMinVersion();
128 deducedVersion = std::max(deducedVersion, *minVersion);
129 if (deducedVersion > allowedVersion) {
130 return op->emitError(
"'")
131 << op->getName() <<
"' requires min version "
132 << spirv::stringifyVersion(deducedVersion)
133 <<
" but target environment allows up to "
134 << spirv::stringifyVersion(allowedVersion);
140 if (
auto maxVersionIfx = dyn_cast<spirv::QueryMaxVersionInterface>(op)) {
141 std::optional<spirv::Version> maxVersion = maxVersionIfx.getMaxVersion();
142 if (maxVersion && *maxVersion < allowedVersion) {
143 return op->emitError(
"'")
144 << op->getName() <<
"' is missing after version "
145 << spirv::stringifyVersion(*maxVersion)
146 <<
" but target environment is "
147 << spirv::stringifyVersion(allowedVersion);
152 if (
auto extensions = dyn_cast<spirv::QueryExtensionInterface>(op))
154 op, targetEnv, extensions.getExtensions(), deducedExtensions)))
158 if (
auto capabilities = dyn_cast<spirv::QueryCapabilityInterface>(op))
160 op, targetEnv, capabilities.getCapabilities(),
161 deducedCapabilities)))
164 SmallVector<Type, 4> valueTypes;
165 valueTypes.append(op->operand_type_begin(), op->operand_type_end());
166 valueTypes.append(op->result_type_begin(), op->result_type_end());
171 auto requireLinkage = [&](spirv::LinkageType linkageType) -> LogicalResult {
172 if (
auto caps = spirv::getCapabilities(linkageType)) {
173 SmallVector<ArrayRef<spirv::Capability>, 1> capCandidates = {*caps};
175 op, targetEnv, capCandidates, deducedCapabilities)))
178 if (
auto exts = spirv::getExtensions(linkageType)) {
179 SmallVector<ArrayRef<spirv::Extension>, 1> extCandidates = {*exts};
181 op, targetEnv, extCandidates, deducedExtensions)))
189 if (
auto globalVar = dyn_cast<spirv::GlobalVariableOp>(op)) {
190 valueTypes.push_back(globalVar.getType());
195 if (globalVar.getBinding() || globalVar.getDescriptorSet()) {
196 spirv::Capability shader = spirv::Capability::Shader;
197 SmallVector<ArrayRef<spirv::Capability>, 1> caps = {shader};
199 deducedCapabilities)))
203 if (
auto linkage = globalVar.getLinkageAttributes())
204 if (
failed(requireLinkage(linkage->getLinkageType().getValue())))
208 if (
auto funcOp = dyn_cast<spirv::FuncOp>(op))
209 if (
auto linkage = funcOp.getLinkageAttributes())
210 if (
failed(requireLinkage(linkage->getLinkageType().getValue())))
214 if (
auto funcOpInterface = dyn_cast<FunctionOpInterface>(op)) {
215 llvm::append_range(valueTypes, funcOpInterface.getArgumentTypes());
216 llvm::append_range(valueTypes, funcOpInterface.getResultTypes());
220 SmallVector<ArrayRef<spirv::Extension>, 4> typeExtensions;
221 SmallVector<ArrayRef<spirv::Capability>, 8> typeCapabilities;
222 for (Type valueType : valueTypes) {
223 typeExtensions.clear();
224 cast<spirv::SPIRVType>(valueType).getExtensions(typeExtensions);
226 op, targetEnv, typeExtensions, deducedExtensions)))
229 typeCapabilities.clear();
230 cast<spirv::SPIRVType>(valueType).getCapabilities(typeCapabilities);
232 op, targetEnv, typeCapabilities, deducedCapabilities)))
240 return signalPassFailure();
245 for (spirv::Capability cap : deducedCapabilities) {
246 if (std::optional<spirv::Version> minVersion = spirv::getMinVersion(cap)) {
247 deducedVersion = std::max(deducedVersion, *minVersion);
248 if (deducedVersion > allowedVersion) {
249 module.emitError("Capability '")
250 << spirv::stringifyCapability(cap) << "' requires min version "
251 << spirv::stringifyVersion(deducedVersion)
252 << " but target environment allows up to "
253 << spirv::stringifyVersion(allowedVersion);
254 return signalPassFailure();
260 deducedVersion, deducedCapabilities.getArrayRef(),
261 deducedExtensions.getArrayRef(), &
getContext());
262 module->setAttr(spirv::ModuleOp::getVCETripleAttrName(), triple);
static LogicalResult checkAndUpdateExtensionRequirements(Operation *op, const spirv::TargetEnv &targetEnv, const spirv::SPIRVType::ExtensionArrayRefVector &candidates, SetVector< spirv::Extension > &deducedExtensions)
Checks that candidates extension requirements are possible to be satisfied with the given targetEnv a...
static void addAllImpliedCapabilities(SetVector< spirv::Capability > &caps)
static LogicalResult checkAndUpdateCapabilityRequirements(Operation *op, const spirv::TargetEnv &targetEnv, const spirv::SPIRVType::CapabilityArrayRefVector &candidates, SetVector< spirv::Capability > &deducedCapabilities)
Checks that candidatescapability requirements are possible to be satisfied with the given targetEnv a...
Operation is the basic unit of execution within MLIR.
InFlightDiagnostic emitError(const Twine &message={})
Emit an error about fatal conditions with this operation, reporting up to any diagnostic handlers tha...
OperationName getName()
The name of an operation is the key identifier for it.
static WalkResult advance()
bool wasInterrupted() const
Returns true if the walk was interrupted.
static WalkResult interrupt()
SmallVectorImpl< ArrayRef< Capability > > CapabilityArrayRefVector
The capability requirements for each type are following the ((Capability::A OR Extension::B) AND (Cap...
SmallVectorImpl< ArrayRef< Extension > > ExtensionArrayRefVector
The extension requirements for each type are following the ((Extension::A OR Extension::B) AND (Exten...
Version getVersion() const
Returns the target version.
A wrapper class around a spirv::TargetEnvAttr to provide query methods for allowed version/capabiliti...
bool allows(Capability) const
Returns true if the given capability is allowed.
static VerCapExtAttr get(Version version, ArrayRef< Capability > capabilities, ArrayRef< Extension > extensions, MLIRContext *context)
Gets a VerCapExtAttr instance.
TargetEnvAttr lookupTargetEnv(Operation *op)
Queries the target environment recursively from enclosing symbol table ops containing the given op.
Include the generated interface declarations.
llvm::SetVector< T, Vector, Set, N > SetVector