19#include "llvm/Support/FormatVariadic.h"
43 impl = std::make_unique<::ScfToSPIRVContextImpl>();
59template <
typename ScfOp,
typename OpTy>
60void replaceSCFOutputValue(ScfOp scfOp, OpTy newOp,
61 ConversionPatternRewriter &rewriter,
66 auto &allocas = scfToSPIRVContext->
outputVars[newOp];
71 for (
Type convertedType : returnTypes) {
74 rewriter.setInsertionPoint(newOp);
75 auto alloc = spirv::VariableOp::create(rewriter, loc, pointerType,
76 spirv::StorageClass::Function,
78 allocas.push_back(alloc);
79 rewriter.setInsertionPointAfter(newOp);
80 Value loadResult = spirv::LoadOp::create(rewriter, loc, alloc);
81 resultValue.push_back(loadResult);
83 rewriter.replaceOp(scfOp, resultValue);
95template <
typename OpTy>
96class SCFToSPIRVPattern :
public OpConversionPattern<OpTy> {
98 SCFToSPIRVPattern(MLIRContext *context,
const SPIRVTypeConverter &converter,
99 ScfToSPIRVContextImpl *scfToSPIRVContext)
100 : OpConversionPattern<OpTy>::OpConversionPattern(converter, context),
101 scfToSPIRVContext(scfToSPIRVContext), typeConverter(converter) {}
104 ScfToSPIRVContextImpl *scfToSPIRVContext;
119 const SPIRVTypeConverter &typeConverter;
127struct ForOpConversion final : SCFToSPIRVPattern<scf::ForOp> {
128 using SCFToSPIRVPattern::SCFToSPIRVPattern;
131 matchAndRewrite(scf::ForOp forOp, OpAdaptor adaptor,
132 ConversionPatternRewriter &rewriter)
const override {
138 auto loc = forOp.getLoc();
139 auto loopControl = spirv::LoopControl::None;
140 if (
auto attr = forOp->getAttrOfType<spirv::LoopControlAttr>(
142 loopControl = attr.getValue();
143 auto loopOp = spirv::LoopOp::create(rewriter, loc, loopControl);
144 loopOp.addEntryAndMergeBlock(rewriter);
146 OpBuilder::InsertionGuard guard(rewriter);
148 Block *header = rewriter.createBlock(&loopOp.getBody(),
149 getBlockIt(loopOp.getBody(), 1));
150 rewriter.setInsertionPointAfter(loopOp);
153 Value adapLowerBound = adaptor.getLowerBound();
154 BlockArgument newIndVar =
156 for (Value arg : adaptor.getInitArgs())
158 Block *body = forOp.getBody();
163 TypeConverter::SignatureConversion signatureConverter(
165 signatureConverter.remapInput(0, newIndVar);
167 signatureConverter.remapInput(i, header->
getArgument(i));
168 body = rewriter.applySignatureConversion(&forOp.getRegion().front(),
173 rewriter.inlineRegionBefore(forOp->getRegion(0), loopOp.getBody(),
174 getBlockIt(loopOp.getBody(), 2));
176 SmallVector<Value, 8> args(1, adaptor.getLowerBound());
177 args.append(adaptor.getInitArgs().begin(), adaptor.getInitArgs().end());
179 rewriter.setInsertionPointToEnd(&(loopOp.getBody().front()));
180 spirv::BranchOp::create(rewriter, loc, header, args);
183 rewriter.setInsertionPointToEnd(header);
184 auto *mergeBlock = loopOp.getMergeBlock();
186 if (forOp.getUnsignedCmp()) {
187 cmpOp = spirv::ULessThanOp::create(rewriter, loc, rewriter.getI1Type(),
188 newIndVar, adaptor.getUpperBound());
190 cmpOp = spirv::SLessThanOp::create(rewriter, loc, rewriter.getI1Type(),
191 newIndVar, adaptor.getUpperBound());
194 spirv::BranchConditionalOp::create(rewriter, loc, cmpOp, body,
195 ArrayRef<Value>(), mergeBlock,
200 Block *continueBlock = loopOp.getContinueBlock();
201 rewriter.setInsertionPointToEnd(continueBlock);
204 Value updatedIndVar = spirv::IAddOp::create(
205 rewriter, loc, newIndVar.
getType(), newIndVar, adaptor.getStep());
206 spirv::BranchOp::create(rewriter, loc, header, updatedIndVar);
212 SmallVector<Type, 8> initTypes;
213 for (
auto arg : adaptor.getInitArgs())
214 initTypes.push_back(arg.getType());
215 replaceSCFOutputValue(forOp, loopOp, rewriter, scfToSPIRVContext,
220 std::optional<APInt> tripCount = forOp.getStaticTripCount();
221 if (!tripCount || tripCount->isZero()) {
222 auto &allocas = scfToSPIRVContext->
outputVars[loopOp];
223 rewriter.setInsertionPoint(loopOp);
224 for (
auto [alloca, init] : llvm::zip(allocas, adaptor.getInitArgs()))
225 spirv::StoreOp::create(rewriter, loc, alloca, init);
237struct IfOpConversion : SCFToSPIRVPattern<scf::IfOp> {
238 using SCFToSPIRVPattern::SCFToSPIRVPattern;
241 matchAndRewrite(scf::IfOp ifOp, OpAdaptor adaptor,
242 ConversionPatternRewriter &rewriter)
const override {
246 auto loc = ifOp.getLoc();
249 SmallVector<Type, 8> returnTypes;
250 for (
auto result : ifOp.getResults()) {
251 auto convertedType = typeConverter.convertType(
result.getType());
253 return rewriter.notifyMatchFailure(
255 llvm::formatv(
"failed to convert type '{0}'",
result.getType()));
257 returnTypes.push_back(convertedType);
262 auto selectionControl = spirv::SelectionControl::None;
263 if (
auto attr = ifOp->getAttrOfType<spirv::SelectionControlAttr>(
265 selectionControl = attr.getValue();
267 spirv::SelectionOp::create(rewriter, loc, selectionControl);
268 auto *mergeBlock = rewriter.createBlock(&selectionOp.getBody(),
269 selectionOp.getBody().end());
270 spirv::MergeOp::create(rewriter, loc);
272 OpBuilder::InsertionGuard guard(rewriter);
273 auto *selectionHeaderBlock =
274 rewriter.createBlock(&selectionOp.getBody().front());
277 auto &thenRegion = ifOp.getThenRegion();
278 auto *thenBlock = &thenRegion.front();
279 rewriter.setInsertionPointToEnd(&thenRegion.back());
280 spirv::BranchOp::create(rewriter, loc, mergeBlock);
281 rewriter.inlineRegionBefore(thenRegion, mergeBlock);
283 auto *elseBlock = mergeBlock;
286 if (!ifOp.getElseRegion().empty()) {
287 auto &elseRegion = ifOp.getElseRegion();
288 elseBlock = &elseRegion.front();
289 rewriter.setInsertionPointToEnd(&elseRegion.back());
290 spirv::BranchOp::create(rewriter, loc, mergeBlock);
291 rewriter.inlineRegionBefore(elseRegion, mergeBlock);
295 rewriter.setInsertionPointToEnd(selectionHeaderBlock);
296 spirv::BranchConditionalOp::create(rewriter, loc, adaptor.getCondition(),
297 thenBlock, ArrayRef<Value>(), elseBlock,
300 replaceSCFOutputValue(ifOp, selectionOp, rewriter, scfToSPIRVContext,
312struct IndexSwitchOpConversion final : SCFToSPIRVPattern<scf::IndexSwitchOp> {
313 using SCFToSPIRVPattern::SCFToSPIRVPattern;
316 matchAndRewrite(scf::IndexSwitchOp switchOp, OpAdaptor adaptor,
317 ConversionPatternRewriter &rewriter)
const override {
318 Location loc = switchOp.getLoc();
321 SmallVector<Type, 8> returnTypes;
322 for (Value
result : switchOp.getResults()) {
323 Type convertedType = typeConverter.convertType(
result.getType());
325 return rewriter.notifyMatchFailure(
327 llvm::formatv(
"failed to convert type '{0}'",
result.getType()));
328 returnTypes.push_back(convertedType);
333 Value selector = adaptor.getArg();
334 auto selectorType = dyn_cast<IntegerType>(selector.
getType());
336 return rewriter.notifyMatchFailure(loc,
337 "selector type is not an integer");
338 unsigned selectorWidth = selectorType.getWidth();
341 auto selectionControl = spirv::SelectionControl::None;
342 if (
auto attr = switchOp->getAttrOfType<spirv::SelectionControlAttr>(
344 selectionControl = attr.getValue();
346 spirv::SelectionOp::create(rewriter, loc, selectionControl);
347 auto *mergeBlock = rewriter.createBlock(&selectionOp.getBody(),
348 selectionOp.getBody().end());
349 spirv::MergeOp::create(rewriter, loc);
351 OpBuilder::InsertionGuard guard(rewriter);
352 auto *headerBlock = rewriter.createBlock(&selectionOp.getBody().front());
355 SmallVector<APInt> caseLiterals;
356 SmallVector<Block *> caseBlocks;
357 ArrayRef<int64_t> cases = switchOp.getCases();
358 for (
auto [caseValue, caseRegion] :
359 llvm::zip_equal(cases, switchOp.getCaseRegions())) {
360 Block *caseBlock = &caseRegion.front();
361 rewriter.setInsertionPointToEnd(&caseRegion.back());
362 spirv::BranchOp::create(rewriter, loc, mergeBlock);
363 rewriter.inlineRegionBefore(caseRegion, mergeBlock);
364 caseLiterals.push_back(
365 APInt(selectorWidth, caseValue,
true));
366 caseBlocks.push_back(caseBlock);
370 Region &defaultRegion = switchOp.getDefaultRegion();
372 rewriter.setInsertionPointToEnd(&defaultRegion.
back());
373 spirv::BranchOp::create(rewriter, loc, mergeBlock);
374 rewriter.inlineRegionBefore(defaultRegion, mergeBlock);
379 SmallVector<ValueRange> caseOperands(caseBlocks.size(),
ValueRange());
380 rewriter.setInsertionPointToEnd(headerBlock);
381 spirv::SwitchOp::create(rewriter, loc, selector, defaultBlock,
ValueRange(),
382 caseLiterals, caseBlocks, caseOperands);
384 replaceSCFOutputValue(switchOp, selectionOp, rewriter, scfToSPIRVContext,
394struct TerminatorOpConversion final : SCFToSPIRVPattern<scf::YieldOp> {
396 using SCFToSPIRVPattern::SCFToSPIRVPattern;
399 matchAndRewrite(scf::YieldOp terminatorOp, OpAdaptor adaptor,
400 ConversionPatternRewriter &rewriter)
const override {
403 Operation *parent = terminatorOp->getParentOp();
407 scf::SCFDialect::getDialectNamespace() &&
408 !isa<scf::IfOp, scf::ForOp, scf::WhileOp, scf::IndexSwitchOp>(parent))
409 return rewriter.notifyMatchFailure(
411 llvm::formatv(
"conversion not supported for parent op: '{0}'",
416 if (!operands.empty()) {
417 auto &allocas = scfToSPIRVContext->
outputVars[parent];
418 if (allocas.size() != operands.size())
421 auto loc = terminatorOp.getLoc();
422 for (
unsigned i = 0, e = operands.size(); i < e; i++)
423 spirv::StoreOp::create(rewriter, loc, allocas[i], operands[i]);
424 if (isa<spirv::LoopOp>(parent)) {
427 auto br = cast<spirv::BranchOp>(
428 rewriter.getInsertionBlock()->getTerminator());
429 SmallVector<Value, 8> args(br.getBlockArguments());
430 args.append(operands.begin(), operands.end());
431 rewriter.setInsertionPoint(br);
432 spirv::BranchOp::create(rewriter, terminatorOp.getLoc(), br.getTarget(),
434 rewriter.eraseOp(br);
437 rewriter.eraseOp(terminatorOp);
446struct WhileOpConversion final : SCFToSPIRVPattern<scf::WhileOp> {
447 using SCFToSPIRVPattern::SCFToSPIRVPattern;
450 matchAndRewrite(scf::WhileOp whileOp, OpAdaptor adaptor,
451 ConversionPatternRewriter &rewriter)
const override {
452 auto loc = whileOp.getLoc();
453 auto loopControl = spirv::LoopControl::None;
454 if (
auto attr = whileOp->getAttrOfType<spirv::LoopControlAttr>(
456 loopControl = attr.getValue();
457 auto loopOp = spirv::LoopOp::create(rewriter, loc, loopControl);
458 loopOp.addEntryAndMergeBlock(rewriter);
460 Region &beforeRegion = whileOp.getBefore();
461 Region &afterRegion = whileOp.getAfter();
463 if (
failed(rewriter.convertRegionTypes(&beforeRegion, typeConverter)) ||
464 failed(rewriter.convertRegionTypes(&afterRegion, typeConverter)))
465 return rewriter.notifyMatchFailure(whileOp,
466 "Failed to convert region types");
468 OpBuilder::InsertionGuard guard(rewriter);
470 Block &entryBlock = *loopOp.getEntryBlock();
473 Block &mergeBlock = *loopOp.getMergeBlock();
475 auto cond = cast<scf::ConditionOp>(beforeBlock.
getTerminator());
476 SmallVector<Value> condArgs;
477 if (
failed(rewriter.getRemappedValues(cond.getArgs(), condArgs)))
480 Value conditionVal = rewriter.getRemappedValue(cond.getCondition());
485 SmallVector<Value> yieldArgs;
486 if (
failed(rewriter.getRemappedValues(yield.getResults(), yieldArgs)))
490 rewriter.inlineRegionBefore(beforeRegion, loopOp.getBody(),
491 getBlockIt(loopOp.getBody(), 1));
494 rewriter.inlineRegionBefore(afterRegion, loopOp.getBody(),
495 getBlockIt(loopOp.getBody(), 2));
498 rewriter.setInsertionPointToEnd(&entryBlock);
499 spirv::BranchOp::create(rewriter, loc, &beforeBlock, adaptor.getInits());
501 auto condLoc = cond.getLoc();
503 SmallVector<Value> resultValues(condArgs.size());
512 for (
const auto &it : llvm::enumerate(condArgs)) {
513 auto res = it.value();
519 rewriter.setInsertionPoint(loopOp);
520 auto alloc = spirv::VariableOp::create(rewriter, condLoc, pointerType,
521 spirv::StorageClass::Function,
525 rewriter.setInsertionPointAfter(loopOp);
526 auto loadResult = spirv::LoadOp::create(rewriter, condLoc, alloc);
527 resultValues[i] = loadResult;
530 rewriter.setInsertionPointToEnd(&beforeBlock);
531 spirv::StoreOp::create(rewriter, condLoc, alloc, res);
534 rewriter.setInsertionPointToEnd(&beforeBlock);
535 rewriter.replaceOpWithNewOp<spirv::BranchConditionalOp>(
536 cond, conditionVal, &afterBlock, condArgs, &mergeBlock,
ValueRange());
539 rewriter.setInsertionPointToEnd(&afterBlock);
540 rewriter.replaceOpWithNewOp<spirv::BranchOp>(yield, &beforeBlock,
543 rewriter.replaceOp(whileOp, resultValues);
556 patterns.
add<ForOpConversion, IfOpConversion, IndexSwitchOpConversion,
557 TerminatorOpConversion, WhileOpConversion>(
BlockArgument getArgument(unsigned i)
unsigned getNumArguments()
Operation * getTerminator()
Get the terminator operation of this block.
BlockArgument addArgument(Type type, Location loc)
Add one value to the argument list.
StringRef getNamespace() const
This class defines the main interface for locations in MLIR and acts as a non-nullable wrapper around...
Dialect * getDialect()
Return the dialect this operation is associated with, or nullptr if the associated dialect is not loa...
OperationName getName()
The name of an operation is the key identifier for it.
This class contains a list of basic blocks and a link to the parent operation it is attached to.
BlockListType::iterator iterator
MLIRContext * getContext() const
RewritePatternSet & add(ConstructorArg &&arg, ConstructorArgs &&...args)
Add an instance of each of the pattern types 'Ts' to the pattern list with the given arguments.
Type conversion from builtin types to SPIR-V types for shader interface.
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
Type getType() const
Return the type of this value.
Location getLoc() const
Return the location of this value.
static PointerType get(Type pointeeType, StorageClass storageClass)
StringRef getLoopControlAttrName()
Returns the attribute name for specifying loop control.
StringRef getSelectionControlAttrName()
Returns the attribute name for specifying selection control.
Include the generated interface declarations.
void populateSCFToSPIRVPatterns(const SPIRVTypeConverter &typeConverter, ScfToSPIRVContext &scfToSPIRVContext, RewritePatternSet &patterns)
Collects a set of patterns to lower from scf.for, scf.if, and loop.terminator to CFG operations withi...
llvm::DenseMap< KeyT, ValueT, KeyInfoT, BucketT > DenseMap
DenseMap< Operation *, SmallVector< spirv::VariableOp, 8 > > outputVars
ScfToSPIRVContext()
We use ScfToSPIRVContext to store information about the lowering of the scf region that need to be us...
ScfToSPIRVContextImpl * getImpl()