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->getDiscardableAttrOfType<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->getDiscardableAttrOfType<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;
343 switchOp->getDiscardableAttrOfType<spirv::SelectionControlAttr>(
345 selectionControl = attr.getValue();
347 spirv::SelectionOp::create(rewriter, loc, selectionControl);
348 auto *mergeBlock = rewriter.createBlock(&selectionOp.getBody(),
349 selectionOp.getBody().end());
350 spirv::MergeOp::create(rewriter, loc);
352 OpBuilder::InsertionGuard guard(rewriter);
353 auto *headerBlock = rewriter.createBlock(&selectionOp.getBody().front());
356 SmallVector<APInt> caseLiterals;
357 SmallVector<Block *> caseBlocks;
358 ArrayRef<int64_t> cases = switchOp.getCases();
359 for (
auto [caseValue, caseRegion] :
360 llvm::zip_equal(cases, switchOp.getCaseRegions())) {
361 Block *caseBlock = &caseRegion.front();
362 rewriter.setInsertionPointToEnd(&caseRegion.back());
363 spirv::BranchOp::create(rewriter, loc, mergeBlock);
364 rewriter.inlineRegionBefore(caseRegion, mergeBlock);
365 caseLiterals.push_back(
366 APInt(selectorWidth, caseValue,
true));
367 caseBlocks.push_back(caseBlock);
371 Region &defaultRegion = switchOp.getDefaultRegion();
373 rewriter.setInsertionPointToEnd(&defaultRegion.
back());
374 spirv::BranchOp::create(rewriter, loc, mergeBlock);
375 rewriter.inlineRegionBefore(defaultRegion, mergeBlock);
380 SmallVector<ValueRange> caseOperands(caseBlocks.size(),
ValueRange());
381 rewriter.setInsertionPointToEnd(headerBlock);
382 spirv::SwitchOp::create(rewriter, loc, selector, defaultBlock,
ValueRange(),
383 caseLiterals, caseBlocks, caseOperands);
385 replaceSCFOutputValue(switchOp, selectionOp, rewriter, scfToSPIRVContext,
395struct TerminatorOpConversion final : SCFToSPIRVPattern<scf::YieldOp> {
397 using SCFToSPIRVPattern::SCFToSPIRVPattern;
400 matchAndRewrite(scf::YieldOp terminatorOp, OpAdaptor adaptor,
401 ConversionPatternRewriter &rewriter)
const override {
404 Operation *parent = terminatorOp->getParentOp();
408 scf::SCFDialect::getDialectNamespace() &&
409 !isa<scf::IfOp, scf::ForOp, scf::WhileOp, scf::IndexSwitchOp>(parent))
410 return rewriter.notifyMatchFailure(
412 llvm::formatv(
"conversion not supported for parent op: '{0}'",
417 if (!operands.empty()) {
418 auto &allocas = scfToSPIRVContext->
outputVars[parent];
419 if (allocas.size() != operands.size())
422 auto loc = terminatorOp.getLoc();
423 for (
unsigned i = 0, e = operands.size(); i < e; i++)
424 spirv::StoreOp::create(rewriter, loc, allocas[i], operands[i]);
425 if (isa<spirv::LoopOp>(parent)) {
428 auto br = cast<spirv::BranchOp>(
429 rewriter.getInsertionBlock()->getTerminator());
430 SmallVector<Value, 8> args(br.getBlockArguments());
431 args.append(operands.begin(), operands.end());
432 rewriter.setInsertionPoint(br);
433 spirv::BranchOp::create(rewriter, terminatorOp.getLoc(), br.getTarget(),
435 rewriter.eraseOp(br);
438 rewriter.eraseOp(terminatorOp);
447struct WhileOpConversion final : SCFToSPIRVPattern<scf::WhileOp> {
448 using SCFToSPIRVPattern::SCFToSPIRVPattern;
451 matchAndRewrite(scf::WhileOp whileOp, OpAdaptor adaptor,
452 ConversionPatternRewriter &rewriter)
const override {
453 auto loc = whileOp.getLoc();
454 auto loopControl = spirv::LoopControl::None;
455 if (
auto attr = whileOp->getDiscardableAttrOfType<spirv::LoopControlAttr>(
457 loopControl = attr.getValue();
458 auto loopOp = spirv::LoopOp::create(rewriter, loc, loopControl);
459 loopOp.addEntryAndMergeBlock(rewriter);
461 Region &beforeRegion = whileOp.getBefore();
462 Region &afterRegion = whileOp.getAfter();
464 if (
failed(rewriter.convertRegionTypes(&beforeRegion, typeConverter)) ||
465 failed(rewriter.convertRegionTypes(&afterRegion, typeConverter)))
466 return rewriter.notifyMatchFailure(whileOp,
467 "Failed to convert region types");
469 OpBuilder::InsertionGuard guard(rewriter);
471 Block &entryBlock = *loopOp.getEntryBlock();
474 Block &mergeBlock = *loopOp.getMergeBlock();
476 auto cond = cast<scf::ConditionOp>(beforeBlock.
getTerminator());
477 SmallVector<Value> condArgs;
478 if (
failed(rewriter.getRemappedValues(cond.getArgs(), condArgs)))
481 Value conditionVal = rewriter.getRemappedValue(cond.getCondition());
486 SmallVector<Value> yieldArgs;
487 if (
failed(rewriter.getRemappedValues(yield.getResults(), yieldArgs)))
491 rewriter.inlineRegionBefore(beforeRegion, loopOp.getBody(),
492 getBlockIt(loopOp.getBody(), 1));
495 rewriter.inlineRegionBefore(afterRegion, loopOp.getBody(),
496 getBlockIt(loopOp.getBody(), 2));
499 rewriter.setInsertionPointToEnd(&entryBlock);
500 spirv::BranchOp::create(rewriter, loc, &beforeBlock, adaptor.getInits());
502 auto condLoc = cond.getLoc();
504 SmallVector<Value> resultValues(condArgs.size());
513 for (
const auto &it : llvm::enumerate(condArgs)) {
514 auto res = it.value();
520 rewriter.setInsertionPoint(loopOp);
521 auto alloc = spirv::VariableOp::create(rewriter, condLoc, pointerType,
522 spirv::StorageClass::Function,
526 rewriter.setInsertionPointAfter(loopOp);
527 auto loadResult = spirv::LoadOp::create(rewriter, condLoc, alloc);
528 resultValues[i] = loadResult;
531 rewriter.setInsertionPointToEnd(&beforeBlock);
532 spirv::StoreOp::create(rewriter, condLoc, alloc, res);
535 rewriter.setInsertionPointToEnd(&beforeBlock);
536 rewriter.replaceOpWithNewOp<spirv::BranchConditionalOp>(
537 cond, conditionVal, &afterBlock, condArgs, &mergeBlock,
ValueRange());
540 rewriter.setInsertionPointToEnd(&afterBlock);
541 rewriter.replaceOpWithNewOp<spirv::BranchOp>(yield, &beforeBlock,
544 rewriter.replaceOp(whileOp, resultValues);
557 patterns.
add<ForOpConversion, IfOpConversion, IndexSwitchOpConversion,
558 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()