MLIR 24.0.0git
UpdateVCEPass.cpp
Go to the documentation of this file.
1//===- DeduceVersionExtensionCapabilityPass.cpp ---------------------------===//
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-exception
6//
7//===----------------------------------------------------------------------===//
8//
9// This file implements a pass to deduce minimal version/extension/capability
10// requirements for a spirv::ModuleOp.
11//
12//===----------------------------------------------------------------------===//
13
15
19#include "mlir/IR/Builders.h"
20#include "mlir/IR/Visitors.h"
21#include "llvm/ADT/StringExtras.h"
22#include <optional>
23
24namespace mlir {
25namespace spirv {
26#define GEN_PASS_DEF_SPIRVUPDATEVCEPASS
27#include "mlir/Dialect/SPIRV/Transforms/Passes.h.inc"
28} // namespace spirv
29} // namespace mlir
30
31using namespace mlir;
32
33namespace {
34/// Pass to deduce minimal version/extension/capability requirements for a
35/// spirv::ModuleOp.
36class UpdateVCEPass final
37 : public spirv::impl::SPIRVUpdateVCEPassBase<UpdateVCEPass> {
38 void runOnOperation() override;
39};
40} // namespace
41
42/// Checks that `candidates` extension requirements are possible to be satisfied
43/// with the given `targetEnv` and updates `deducedExtensions` if so. Emits
44/// errors attaching to the given `op` on failures.
45///
46/// `candidates` is a vector of vector for extension requirements following
47/// ((Extension::A OR Extension::B) AND (Extension::C OR Extension::D))
48/// convention.
50 Operation *op, const spirv::TargetEnv &targetEnv,
52 SetVector<spirv::Extension> &deducedExtensions) {
53 for (const auto &ors : candidates) {
54 if (std::optional<spirv::Extension> chosen = targetEnv.allows(ors)) {
55 deducedExtensions.insert(*chosen);
56 } else {
58 for (spirv::Extension ext : ors)
59 extStrings.push_back(spirv::stringifyExtension(ext));
60
61 return op->emitError("'")
62 << op->getName() << "' requires at least one extension in ["
63 << llvm::join(extStrings, ", ")
64 << "] but none allowed in target environment";
65 }
66 }
67 return success();
68}
69
70/// Checks that `candidates`capability requirements are possible to be satisfied
71/// with the given `targetEnv` and updates `deducedCapabilities` if so. Emits
72/// errors attaching to the given `op` on failures.
73///
74/// `candidates` is a vector of vector for capability requirements following
75/// ((Capability::A OR Capability::B) AND (Capability::C OR Capability::D))
76/// convention.
78 Operation *op, const spirv::TargetEnv &targetEnv,
80 SetVector<spirv::Capability> &deducedCapabilities) {
81 for (const auto &ors : candidates) {
82 if (std::optional<spirv::Capability> chosen = targetEnv.allows(ors)) {
83 deducedCapabilities.insert(*chosen);
84 } else {
86 for (spirv::Capability cap : ors)
87 capStrings.push_back(spirv::stringifyCapability(cap));
88
89 return op->emitError("'")
90 << op->getName() << "' requires at least one capability in ["
91 << llvm::join(capStrings, ", ")
92 << "] but none allowed in target environment";
93 }
94 }
95 return success();
96}
97
100 for (spirv::Capability cap : caps)
101 tmp.insert_range(getRecursiveImpliedCapabilities(cap));
102 caps.insert_range(std::move(tmp));
103}
104
105void UpdateVCEPass::runOnOperation() {
106 spirv::ModuleOp module = getOperation();
107
108 spirv::TargetEnvAttr targetAttr = spirv::lookupTargetEnv(module);
109 if (!targetAttr) {
110 module.emitError("missing 'spirv.target_env' attribute");
111 return signalPassFailure();
112 }
113
114 spirv::TargetEnv targetEnv(targetAttr);
115 spirv::Version allowedVersion = targetAttr.getVersion();
116
117 spirv::Version deducedVersion = spirv::Version::V_1_0;
118 SetVector<spirv::Extension> deducedExtensions;
119 SetVector<spirv::Capability> deducedCapabilities;
120
121 // Walk each SPIR-V op to deduce the minimal version/extension/capability
122 // requirements.
123 WalkResult walkResult = module.walk([&](Operation *op) -> WalkResult {
124 // Op min version requirements
125 if (auto minVersionIfx = dyn_cast<spirv::QueryMinVersionInterface>(op)) {
126 std::optional<spirv::Version> minVersion = minVersionIfx.getMinVersion();
127 if (minVersion) {
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);
135 }
136 }
137 }
138
139 // Op max version requirements
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);
148 }
149 }
150
151 // Op extension requirements
152 if (auto extensions = dyn_cast<spirv::QueryExtensionInterface>(op))
154 op, targetEnv, extensions.getExtensions(), deducedExtensions)))
155 return WalkResult::interrupt();
156
157 // Op capability requirements
158 if (auto capabilities = dyn_cast<spirv::QueryCapabilityInterface>(op))
160 op, targetEnv, capabilities.getCapabilities(),
161 deducedCapabilities)))
162 return WalkResult::interrupt();
163
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());
167
168 // Per the SPIR-V spec Decoration table, the `LinkageAttributes` decoration
169 // requires the `Linkage` capability, and specific linkage types pull in
170 // additional extensions (e.g., `LinkOnceODR` -> `SPV_KHR_linkonce_odr`).
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)))
176 return failure();
177 }
178 if (auto exts = spirv::getExtensions(linkageType)) {
179 SmallVector<ArrayRef<spirv::Extension>, 1> extCandidates = {*exts};
181 op, targetEnv, extCandidates, deducedExtensions)))
182 return failure();
183 }
184 return success();
185 };
186
187 // Special treatment for global variables, whose type requirements are
188 // conveyed by type attributes.
189 if (auto globalVar = dyn_cast<spirv::GlobalVariableOp>(op)) {
190 valueTypes.push_back(globalVar.getType());
191
192 // The `DescriptorSet` and `Binding` decorations (represented by the
193 // `binding` and `descriptor_set` attributes) require the `Shader`
194 // capability per the SPIR-V spec Decoration table.
195 if (globalVar.getBinding() || globalVar.getDescriptorSet()) {
196 spirv::Capability shader = spirv::Capability::Shader;
197 SmallVector<ArrayRef<spirv::Capability>, 1> caps = {shader};
198 if (failed(checkAndUpdateCapabilityRequirements(op, targetEnv, caps,
199 deducedCapabilities)))
200 return WalkResult::interrupt();
201 }
202
203 if (auto linkage = globalVar.getLinkageAttributes())
204 if (failed(requireLinkage(linkage->getLinkageType().getValue())))
205 return WalkResult::interrupt();
206 }
207
208 if (auto funcOp = dyn_cast<spirv::FuncOp>(op))
209 if (auto linkage = funcOp.getLinkageAttributes())
210 if (failed(requireLinkage(linkage->getLinkageType().getValue())))
211 return WalkResult::interrupt();
212
213 // If the op is FunctionLike make sure to process input and result types.
214 if (auto funcOpInterface = dyn_cast<FunctionOpInterface>(op)) {
215 llvm::append_range(valueTypes, funcOpInterface.getArgumentTypes());
216 llvm::append_range(valueTypes, funcOpInterface.getResultTypes());
217 }
218
219 // Requirements from values' types
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)))
227 return WalkResult::interrupt();
228
229 typeCapabilities.clear();
230 cast<spirv::SPIRVType>(valueType).getCapabilities(typeCapabilities);
232 op, targetEnv, typeCapabilities, deducedCapabilities)))
233 return WalkResult::interrupt();
234 }
235
236 return WalkResult::advance();
237 });
238
239 if (walkResult.wasInterrupted())
240 return signalPassFailure();
241
242 addAllImpliedCapabilities(deducedCapabilities);
243
244 // Update min version requirement for capabilities after deducing them.
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();
255 }
256 }
257 }
258
259 auto triple = spirv::VerCapExtAttr::get(
260 deducedVersion, deducedCapabilities.getArrayRef(),
261 deducedExtensions.getArrayRef(), &getContext());
262 module->setAttr(spirv::ModuleOp::getVCETripleAttrName(), triple);
263}
return success()
b getContext())
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.
Definition Operation.h:87
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.
Definition Operation.h:115
static WalkResult advance()
Definition WalkResult.h:47
bool wasInterrupted() const
Returns true if the walk was interrupted.
Definition WalkResult.h:51
static WalkResult interrupt()
Definition WalkResult.h:46
SmallVectorImpl< ArrayRef< Capability > > CapabilityArrayRefVector
The capability requirements for each type are following the ((Capability::A OR Extension::B) AND (Cap...
Definition SPIRVTypes.h:66
SmallVectorImpl< ArrayRef< Extension > > ExtensionArrayRefVector
The extension requirements for each type are following the ((Extension::A OR Extension::B) AND (Exten...
Definition SPIRVTypes.h:55
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.
detail::InFlightRemark failed(Location loc, RemarkOpts opts)
Report an optimization remark that failed.
Definition Remarks.h:717
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
Definition LLVM.h:125