29#include "llvm/Support/LogicalResult.h"
32#define DEBUG_TYPE "wasm-convert"
35#define GEN_PASS_DEF_RAISEWASMMLIR
36#include "mlir/Conversion/Passes.h.inc"
43template <
typename SourceOp,
typename TargetIntOp,
typename TargetFPOp>
44struct IntFPDispatchMappingConversion : OpConversionPattern<SourceOp> {
45 using OpConversionPattern<SourceOp>::OpConversionPattern;
48 matchAndRewrite(SourceOp srcOp,
typename SourceOp::Adaptor adaptor,
49 ConversionPatternRewriter &rewriter)
const override {
50 Type type = srcOp.getRhs().getType();
52 rewriter.replaceOpWithNewOp<TargetIntOp>(srcOp, srcOp->getResultTypes(),
53 adaptor.getOperands());
58 rewriter.replaceOpWithNewOp<TargetFPOp>(srcOp, srcOp->getResultTypes(),
59 adaptor.getOperands());
64using WasmAddOpConversion =
65 IntFPDispatchMappingConversion<AddOp, arith::AddIOp, arith::AddFOp>;
66using WasmMulOpConversion =
67 IntFPDispatchMappingConversion<MulOp, arith::MulIOp, arith::MulFOp>;
68using WasmSubOpConversion =
69 IntFPDispatchMappingConversion<SubOp, arith::SubIOp, arith::SubFOp>;
73template <
typename SourceOp,
typename TargetOp>
74struct OpMappingConversion : OpConversionPattern<SourceOp> {
75 using OpConversionPattern<SourceOp>::OpConversionPattern;
78 matchAndRewrite(SourceOp srcOp,
typename SourceOp::Adaptor adaptor,
79 ConversionPatternRewriter &rewriter)
const override {
80 rewriter.replaceOpWithNewOp<TargetOp>(srcOp, srcOp->getResultTypes(),
81 adaptor.getOperands());
86using WasmAndOpConversion = OpMappingConversion<AndOp, arith::AndIOp>;
87using WasmCeilOpConversion = OpMappingConversion<CeilOp, math::CeilOp>;
90using WasmConvertSOpConversion =
91 OpMappingConversion<ConvertSOp, arith::SIToFPOp>;
92using WasmConvertUOpConversion =
93 OpMappingConversion<ConvertUOp, arith::UIToFPOp>;
94using WasmDemoteOpConversion = OpMappingConversion<DemoteOp, arith::TruncFOp>;
95using WasmDivFPOpConversion = OpMappingConversion<DivOp, arith::DivFOp>;
96using WasmDivSIOpConversion = OpMappingConversion<DivSIOp, arith::DivSIOp>;
97using WasmDivUIOpConversion = OpMappingConversion<DivUIOp, arith::DivUIOp>;
98using WasmExtendSOpConversion =
99 OpMappingConversion<ExtendSI32Op, arith::ExtSIOp>;
100using WasmExtendUOpConversion =
101 OpMappingConversion<ExtendUI32Op, arith::ExtUIOp>;
102using WasmFloorOpConversion = OpMappingConversion<FloorOp, math::FloorOp>;
103using WasmMaxOpConversion = OpMappingConversion<MaxOp, arith::MaximumFOp>;
104using WasmMinOpConversion = OpMappingConversion<MinOp, arith::MinimumFOp>;
105using WasmOrOpConversion = OpMappingConversion<OrOp, arith::OrIOp>;
106using WasmPromoteOpConversion = OpMappingConversion<PromoteOp, arith::ExtFOp>;
107using WasmRemSIOpConversion = OpMappingConversion<RemSIOp, arith::RemSIOp>;
108using WasmRemUIOpConversion = OpMappingConversion<RemUIOp, arith::RemUIOp>;
109using WasmReinterpretOpConversion =
110 OpMappingConversion<ReinterpretOp, arith::BitcastOp>;
111using WasmShLOpConversion = OpMappingConversion<ShLOp, arith::ShLIOp>;
112using WasmShRSOpConversion = OpMappingConversion<ShRSOp, arith::ShRSIOp>;
113using WasmShRUOpConversion = OpMappingConversion<ShRUOp, arith::ShRUIOp>;
114using WasmXOrOpConversion = OpMappingConversion<XOrOp, arith::XOrIOp>;
115using WasmNegOpConversion = OpMappingConversion<NegOp, arith::NegFOp>;
116using WasmCopySignOpConversion =
117 OpMappingConversion<CopySignOp, math::CopySignOp>;
118using WasmClzOpConversion =
119 OpMappingConversion<ClzOp, math::CountLeadingZerosOp>;
120using WasmCtzOpConversion =
121 OpMappingConversion<CtzOp, math::CountTrailingZerosOp>;
122using WasmPopCntOpConversion = OpMappingConversion<PopCntOp, math::CtPopOp>;
123using WasmAbsOpConversion = OpMappingConversion<AbsOp, math::AbsFOp>;
124using WasmTruncOpConversion = OpMappingConversion<TruncOp, math::TruncOp>;
125using WasmSqrtOpConversion = OpMappingConversion<SqrtOp, math::SqrtOp>;
126using WasmWrapOpConversion = OpMappingConversion<WrapOp, arith::TruncIOp>;
148template <
typename SourceOp,
typename LHSShiftOp,
typename RHSShiftOp>
149struct RotateOpConversion : OpConversionPattern<SourceOp> {
150 using OpConversionPattern<SourceOp>::OpConversionPattern;
153 matchAndRewrite(SourceOp srcOp,
typename SourceOp::Adaptor adaptor,
154 ConversionPatternRewriter &rewriter)
const override {
155 const Type ty = srcOp->getResultTypes()[0];
156 const Location loc = srcOp->getLoc();
157 const Value val = adaptor.getVal();
158 const Value bits = adaptor.getBits();
162 auto cstWidthMinusOne =
163 ConstOp::create(rewriter, loc, IntegerAttr::get(ty, width - 1));
167 auto orLHS = LHSShiftOp::create(
169 AndOp::create(rewriter, loc, bits, cstWidthMinusOne));
173 auto orRHS = RHSShiftOp::create(
176 AndOp::create(rewriter, loc,
178 SubOp::create(rewriter, loc,
179 ConstOp::create(rewriter, loc,
180 IntegerAttr::get(ty, 0)),
186 rewriter.replaceOpWithNewOp<OrOp>(srcOp, orLHS, orRHS);
191using WasmRotrOpConversion = RotateOpConversion<RotrOp, ShRUOp, ShLOp>;
192using WasmRotlOpConversion = RotateOpConversion<RotlOp, ShLOp, ShRUOp>;
194template <
typename SourceOp,
typename TargetOp,
typename AttrType,
195 typename ValType, ValType flag>
196struct ComparisonOpConversion : OpConversionPattern<SourceOp> {
197 using OpConversionPattern<SourceOp>::OpConversionPattern;
200 matchAndRewrite(SourceOp srcOp,
typename SourceOp::Adaptor adaptor,
201 ConversionPatternRewriter &rewriter)
const override {
203 TargetOp::create(rewriter, srcOp.getLoc(), rewriter.getI1Type(),
204 AttrType::get(rewriter.getContext(), flag),
205 adaptor.getLhs(), adaptor.getRhs())
207 rewriter.replaceOpWithNewOp<arith::ExtUIOp>(srcOp, rewriter.getI32Type(),
214template <
typename SourceOp, arith::CmpFPredicate compFlag>
215using FPComparisonConversion =
216 ComparisonOpConversion<SourceOp, arith::CmpFOp, arith::CmpFPredicateAttr,
217 arith::CmpFPredicate, compFlag>;
219template <
typename SourceOp, arith::CmpIPredicate compFlag>
220using IntComparisonConversion =
221 ComparisonOpConversion<SourceOp, arith::CmpIOp, arith::CmpIPredicateAttr,
222 arith::CmpIPredicate, compFlag>;
224using WasmLtSIOpConversion =
225 IntComparisonConversion<LtSIOp, arith::CmpIPredicate::slt>;
226using WasmLeSIOpConversion =
227 IntComparisonConversion<LeSIOp, arith::CmpIPredicate::sle>;
228using WasmGtSIOpConversion =
229 IntComparisonConversion<GtSIOp, arith::CmpIPredicate::sgt>;
230using WasmGeSIOpConversion =
231 IntComparisonConversion<GeSIOp, arith::CmpIPredicate::sge>;
232using WasmLtUIOpConversion =
233 IntComparisonConversion<LtUIOp, arith::CmpIPredicate::ult>;
234using WasmLeUIOpConversion =
235 IntComparisonConversion<LeUIOp, arith::CmpIPredicate::ule>;
236using WasmGtUIOpConversion =
237 IntComparisonConversion<GtUIOp, arith::CmpIPredicate::ugt>;
238using WasmGeUIOpConversion =
239 IntComparisonConversion<GeUIOp, arith::CmpIPredicate::uge>;
240using WasmLtOpConversion =
241 FPComparisonConversion<LtOp, arith::CmpFPredicate::OLT>;
242using WasmLeOpConversion =
243 FPComparisonConversion<LeOp, arith::CmpFPredicate::OLE>;
244using WasmGtOpConversion =
245 FPComparisonConversion<GtOp, arith::CmpFPredicate::OGT>;
246using WasmGeOpConversion =
247 FPComparisonConversion<GeOp, arith::CmpFPredicate::OGE>;
249template <
typename SourceOp, arith::CmpIPredicate IntFlag,
250 arith::CmpFPredicate FloatFlag>
251struct IntFpComparisonOpConversion : OpConversionPattern<SourceOp> {
252 using OpConversionPattern<SourceOp>::OpConversionPattern;
255 matchAndRewrite(SourceOp srcOp,
typename SourceOp::Adaptor adaptor,
256 ConversionPatternRewriter &rewriter)
const override {
257 Value comparisonResult;
258 if (srcOp.getLhs().getType().isInteger())
260 arith::CmpIOp::create(
261 rewriter, srcOp.getLoc(), rewriter.getI1Type(),
262 arith::CmpIPredicateAttr::get(rewriter.getContext(), IntFlag),
263 adaptor.getLhs(), adaptor.getRhs())
265 else if (srcOp.getLhs().getType().isFloat())
267 arith::CmpFOp::create(
268 rewriter, srcOp.getLoc(), rewriter.getI1Type(),
269 arith::CmpFPredicateAttr::get(rewriter.getContext(), FloatFlag),
270 adaptor.getLhs(), adaptor.getRhs())
273 return rewriter.notifyMatchFailure(
274 srcOp.getLoc(),
"Unsupported datatype for comparison OP.");
276 rewriter.replaceOpWithNewOp<arith::ExtUIOp>(srcOp, rewriter.getI32Type(),
282using WasmEqOpConversion =
283 IntFpComparisonOpConversion<EqOp, arith::CmpIPredicate::eq,
284 arith::CmpFPredicate::OEQ>;
285using WasmNeOpConversion =
286 IntFpComparisonOpConversion<NeOp, arith::CmpIPredicate::ne,
287 arith::CmpFPredicate::ONE>;
289struct WasmCallOpConversion : OpConversionPattern<FuncCallOp> {
290 using OpConversionPattern::OpConversionPattern;
293 matchAndRewrite(FuncCallOp funcCallOp, FuncCallOp::Adaptor adaptor,
294 ConversionPatternRewriter &rewriter)
const override {
295 rewriter.replaceOpWithNewOp<func::CallOp>(
296 funcCallOp, funcCallOp.getCallee(), funcCallOp.getResults().getTypes(),
297 funcCallOp.getOperands());
302struct WasmConstOpConversion : OpConversionPattern<ConstOp> {
303 using OpConversionPattern::OpConversionPattern;
306 matchAndRewrite(ConstOp constOp, ConstOp::Adaptor adaptor,
307 ConversionPatternRewriter &rewriter)
const override {
308 rewriter.replaceOpWithNewOp<arith::ConstantOp>(constOp, constOp.getValue());
313struct WasmEqzOpConversion : OpConversionPattern<EqzOp> {
314 using OpConversionPattern::OpConversionPattern;
317 matchAndRewrite(EqzOp eqzOp, EqzOp::Adaptor adaptor,
318 ConversionPatternRewriter &rewriter)
const override {
319 auto loc = eqzOp->getLoc();
320 auto zero = arith::ConstantOp::create(
322 rewriter.getIntegerAttr(adaptor.getInput().getType(), 0))
324 auto cmpRes = arith::CmpIOp::create(
325 rewriter, loc, rewriter.getI1Type(),
326 arith::CmpIPredicateAttr::get(rewriter.getContext(),
327 arith::CmpIPredicate::eq),
328 adaptor.getInput(), zero)
330 rewriter.replaceOpWithNewOp<arith::ExtUIOp>(eqzOp, rewriter.getI32Type(),
337struct WasmExtendLowBitsOpConversion : OpConversionPattern<ExtendLowBitsSOp> {
338 using OpConversionPattern::OpConversionPattern;
341 matchAndRewrite(ExtendLowBitsSOp extendLowBytesSOp,
342 ExtendLowBitsSOp::Adaptor adaptor,
343 ConversionPatternRewriter &rewriter)
const override {
344 auto truncWidth = extendLowBytesSOp.getBitsToTake().getInt();
345 auto truncation = arith::TruncIOp::create(
346 rewriter, extendLowBytesSOp->getLoc(),
347 rewriter.getIntegerType(truncWidth), adaptor.getInput());
348 rewriter.replaceOpWithNewOp<arith::ExtSIOp>(
349 extendLowBytesSOp, extendLowBytesSOp.getResult().
getType(),
350 truncation.getResult());
355struct WasmFuncImportOpConversion : OpConversionPattern<FuncImportOp> {
356 using OpConversionPattern::OpConversionPattern;
359 matchAndRewrite(FuncImportOp funcImportOp, FuncImportOp::Adaptor,
360 ConversionPatternRewriter &rewriter)
const override {
361 auto nFunc = rewriter.replaceOpWithNewOp<func::FuncOp>(
362 funcImportOp, funcImportOp.getSymName(), funcImportOp.getType());
363 nFunc.setVisibility(SymbolTable::Visibility::Private);
368struct WasmFuncOpConversion : OpConversionPattern<FuncOp> {
369 using OpConversionPattern::OpConversionPattern;
376 class CFRewriterVisitor {
378 using branch_to_dest_t = llvm::DenseMap<LabelBranchingOpInterface, Block *>;
379 Value getCompResultAsI1(Value compResult,
380 ConversionPatternRewriter &rewriter) {
381 auto testValue = arith::ConstantOp::create(rewriter, compResult.
getLoc(),
382 rewriter.getI32IntegerAttr(0));
383 auto flag = arith::CmpIOp::create(
384 rewriter, compResult.
getLoc(), rewriter.getIntegerType(1),
385 arith::CmpIPredicate::ne, compResult, testValue)
390 void replaceNestLevelWithBranch(BlockOp blockOp,
391 llvm::ArrayRef<Block *> regionsToEntry,
392 ConversionPatternRewriter &rewriter) {
393 rewriter.replaceOpWithNewOp<cf::BranchOp>(blockOp, regionsToEntry[0],
394 blockOp->getOperands());
397 void replaceNestLevelWithBranch(LoopOp loopOp,
398 llvm::ArrayRef<Block *> regionsToEntry,
399 ConversionPatternRewriter &rewriter) {
400 rewriter.replaceOpWithNewOp<cf::BranchOp>(loopOp, regionsToEntry[0],
401 loopOp->getOperands());
404 void replaceNestLevelWithBranch(IfOp ifOp,
405 llvm::ArrayRef<Block *> regionsToEntry,
406 ConversionPatternRewriter &rewriter) {
408 regionsToEntry.size() == 2 ? regionsToEntry[1] : ifOp.getTarget();
409 auto flag = getCompResultAsI1(ifOp.getCondition(), rewriter);
410 rewriter.replaceOpWithNewOp<cf::CondBranchOp>(
411 ifOp, flag, regionsToEntry[0], ifOp.getInputs(), falseDest,
415 template <
typename LevelType>
417 replaceNestLevelWithBranchWrapper(LabelLevelOpInterface nestingOp,
418 llvm::ArrayRef<Block *> regionsToEntry,
419 ConversionPatternRewriter &rewriter) {
420 auto cast = dyn_cast<LevelType>(nestingOp.getOperation());
423 replaceNestLevelWithBranch(cast, regionsToEntry, rewriter);
427 template <
typename... LevelTypes>
428 LogicalResult inlineNestDispatcher(LabelLevelOpInterface nestingOp,
429 ConversionPatternRewriter &rewriter) {
430 auto sip = rewriter.saveInsertionPoint();
431 Block *blockSuccessor = nestingOp->getSuccessor(0);
432 llvm::SmallVector<Block *, 2> regionEntries;
433 LLVM_DEBUG(llvm::dbgs()
434 <<
"Starting inlining blocks for " << nestingOp <<
"\n";);
435 for (
auto ®ion : nestingOp->getRegions()) {
438 regionEntries.push_back(®ion.front());
440 llvm::SmallVector<LabelLevelOpInterface> nestedOps{
441 region.getOps<LabelLevelOpInterface>()};
442 for (
auto nestedOp : nestedOps) {
443 LLVM_DEBUG(llvm::dbgs() <<
" Found nested op: " << nestedOp);
444 if (
failed(inlineBlocks(nestedOp, rewriter)))
447 rewriter.inlineRegionBefore(region, blockSuccessor);
449 LLVM_DEBUG(llvm::dbgs() <<
"End of region inlining\n");
450 LLVM_DEBUG(llvm::dbgs() <<
"Replacing initial op with branching\n");
451 rewriter.setInsertionPoint(nestingOp);
453 (... || succeeded(replaceNestLevelWithBranchWrapper<LevelTypes>(
454 nestingOp, regionEntries, rewriter))));
455 rewriter.restoreInsertionPoint(sip);
458 "Unable to inline the operation regions.");
463 LogicalResult inlineBlocks(LabelLevelOpInterface nestingOp,
464 ConversionPatternRewriter &rewriter) {
465 return inlineNestDispatcher<BlockOp, IfOp, LoopOp>(nestingOp, rewriter);
468 llvm::FailureOr<Block *> getBlockFor(LabelBranchingOpInterface branchOp) {
469 auto destIter = branchToDest.find(branchOp);
470 if (destIter == branchToDest.end())
471 return branchOp->emitError(
"No indexed label op for this operation: ")
473 return destIter->second;
476 inline void convertBranch(BranchIfOp brOp,
Block *dest,
477 ConversionPatternRewriter &rewriter) {
478 auto flag = getCompResultAsI1(brOp.getCondition(), rewriter);
479 rewriter.replaceOpWithNewOp<cf::CondBranchOp>(
480 brOp, flag, dest, brOp.getInputs(), brOp.getElseSuccessor(),
484 inline void convertBranch(BlockReturnOp brOp,
Block *dest,
485 ConversionPatternRewriter &rewriter) {
486 rewriter.replaceOpWithNewOp<cf::BranchOp>(brOp, dest, brOp.getInputs());
489 template <
typename LevelInterfaceT>
491 convertBranchWrapper(LabelBranchingOpInterface branchOp,
Block *dest,
492 ConversionPatternRewriter &rewriter) {
493 auto cast = dyn_cast<LevelInterfaceT>(branchOp.getOperation());
496 auto sip = rewriter.saveInsertionPoint();
497 rewriter.setInsertionPoint(branchOp);
498 convertBranch(cast, dest, rewriter);
499 rewriter.restoreInsertionPoint(sip);
503 template <
typename... BranchInterfaceT>
504 LogicalResult convertBranchDispatch(LabelBranchingOpInterface branchOp,
505 ConversionPatternRewriter &rewriter) {
506 auto dest = getBlockFor(branchOp);
510 success((... || succeeded(convertBranchWrapper<BranchInterfaceT>(
511 branchOp, *dest, rewriter))));
513 return emitError(branchOp->getLoc(),
"No known converter for op ")
518 LogicalResult convertBranch(LabelBranchingOpInterface branchOp,
519 ConversionPatternRewriter &rewriter) {
520 return convertBranchDispatch<BlockReturnOp, BranchIfOp>(branchOp,
525 branch_to_dest_t branchToDest;
528 CFRewriterVisitor(func::FuncOp func) : func{func} {
529 func.walk([
this](LabelBranchingOpInterface branchOp) {
530 branchToDest.insert({branchOp, branchOp.getTarget()});
533 LogicalResult
rewrite(ConversionPatternRewriter &rewriter) {
534 llvm::SmallVector<LabelLevelOpInterface> nestingOps{
535 func.getOps<LabelLevelOpInterface>()};
536 for (
auto nestingOp : nestingOps)
537 if (
failed(inlineBlocks(nestingOp, rewriter)))
541 func->walk([
this, &rewriter](LabelBranchingOpInterface branchOp) {
542 if (
failed(convertBranch(branchOp, rewriter)))
546 return failure(res.wasInterrupted());
551 matchAndRewrite(FuncOp funcOp, FuncOp::Adaptor adaptor,
552 ConversionPatternRewriter &rewriter)
const override {
554 func::FuncOp::create(rewriter, funcOp->getLoc(), funcOp.getSymName(),
555 funcOp.getFunctionType());
556 rewriter.cloneRegionBefore(funcOp.getBody(), newFunc.getBody(),
557 newFunc.getBody().end());
558 Block *oldEntryBlock = &newFunc.getBody().front();
560 TypeConverter::SignatureConversion sC{oldEntryBlock->
getNumArguments()};
561 auto numArgs = blockArgTypes.size();
562 for (
size_t i = 0; i < numArgs; ++i) {
563 auto argType = dyn_cast<LocalRefType>(blockArgTypes[i]);
566 sC.addInputs(i, argType.getElementType());
569 rewriter.applySignatureConversion(oldEntryBlock, sC, getTypeConverter());
570 rewriter.replaceOp(funcOp, newFunc);
571 CFRewriterVisitor cfRewriter{newFunc};
572 return cfRewriter.rewrite(rewriter);
576struct WasmGlobalImportOpConverter : OpConversionPattern<GlobalImportOp> {
577 using OpConversionPattern::OpConversionPattern;
579 matchAndRewrite(GlobalImportOp gIOp, GlobalImportOp::Adaptor adaptor,
580 ConversionPatternRewriter &rewriter)
const override {
581 auto memrefGOp = rewriter.replaceOpWithNewOp<memref::GlobalOp>(
582 gIOp, gIOp.getSymNameAttr(), rewriter.getStringAttr(
"nested"),
583 TypeAttr::get(MemRefType::get({1}, gIOp.getType())), Attribute{},
586 memrefGOp.setConstant(!gIOp.getIsMutable());
591template <
typename CRTP,
typename OriginOpType>
592struct GlobalOpConverter : OpConversionPattern<GlobalOp> {
593 using OpConversionPattern::OpConversionPattern;
595 matchAndRewrite(GlobalOp globalOp, GlobalOp::Adaptor adaptor,
596 ConversionPatternRewriter &rewriter)
const override {
597 ReturnOp rop = globalOp.getInitTerminator();
599 if (rop->getNumOperands() != 1)
600 return rewriter.notifyMatchFailure(
601 globalOp,
"globalOp initializer should return one value exactly");
604 dyn_cast<OriginOpType>(rop->getOperand(0).getDefiningOp());
607 return rewriter.notifyMatchFailure(
608 globalOp,
"invalid initializer op type for this pattern");
610 return static_cast<CRTP
const *
>(
this)->handleInitializer(
611 globalOp, rewriter, initializerOp);
615struct WasmGlobalWithConstInitConversion
616 : GlobalOpConverter<WasmGlobalWithConstInitConversion, ConstOp> {
617 using GlobalOpConverter::GlobalOpConverter;
618 LogicalResult handleInitializer(GlobalOp globalOp,
619 ConversionPatternRewriter &rewriter,
620 ConstOp constInit)
const {
623 ArrayRef<Attribute>{constInit.getValueAttr()});
624 auto globalReplacement = rewriter.replaceOpWithNewOp<memref::GlobalOp>(
625 globalOp, globalOp.getSymNameAttr(), rewriter.getStringAttr(
"private"),
626 TypeAttr::get(MemRefType::get({1}, globalOp.getType())), initializer,
629 globalReplacement.setConstant(!globalOp.getIsMutable());
634struct WasmGlobalWithGetGlobalInitConversion
635 : GlobalOpConverter<WasmGlobalWithGetGlobalInitConversion, GlobalGetOp> {
636 using GlobalOpConverter::GlobalOpConverter;
637 LogicalResult handleInitializer(GlobalOp globalOp,
638 ConversionPatternRewriter &rewriter,
639 GlobalGetOp constInit)
const {
640 auto globalReplacement = rewriter.replaceOpWithNewOp<memref::GlobalOp>(
641 globalOp, globalOp.getSymNameAttr(), rewriter.getStringAttr(
"private"),
642 TypeAttr::get(MemRefType::get({1}, globalOp.getType())),
643 rewriter.getUnitAttr(),
646 globalReplacement.setConstant(!globalOp.getIsMutable());
647 auto loc = globalOp.getLoc();
648 auto initializerName = (globalOp.getSymName() +
"::initializer").str();
649 auto globalInitializer =
650 func::FuncOp::create(rewriter, loc, initializerName,
652 globalInitializer->setAttr(rewriter.getStringAttr(
"initializer"),
653 rewriter.getUnitAttr());
654 auto *initializerBody = globalInitializer.addEntryBlock();
655 auto sip = rewriter.saveInsertionPoint();
656 rewriter.setInsertionPointToStart(initializerBody);
657 auto srcGlobalPtr = memref::GetGlobalOp::create(
658 rewriter, loc, MemRefType::get({1}, constInit.getType()),
659 constInit.getGlobal());
661 memref::GetGlobalOp::create(rewriter, loc, globalReplacement.getType(),
662 globalReplacement.getSymName());
665 memref::LoadOp::create(rewriter, loc, srcGlobalPtr,
ValueRange{idx});
666 memref::StoreOp::create(rewriter, loc, loadSrc.getResult(),
668 func::ReturnOp::create(rewriter, loc);
669 rewriter.restoreInsertionPoint(sip);
674struct WasmGlobalSetOpConversion : OpConversionPattern<GlobalSetOp> {
675 using OpConversionPattern::OpConversionPattern;
677 matchAndRewrite(GlobalSetOp globalSetOp, GlobalSetOp::Adaptor adaptor,
678 ConversionPatternRewriter &rewriter)
const override {
679 auto loc = globalSetOp.getLoc();
680 auto globalPtr = memref::GetGlobalOp::create(
681 rewriter, loc, MemRefType::get({1}, adaptor.getValue().
getType()),
682 globalSetOp.getGlobal());
684 rewriter.replaceOpWithNewOp<memref::StoreOp>(
685 globalSetOp, adaptor.getValue(), globalPtr.getResult(),
691struct WasmMemoryOpConversion : OpConversionPattern<MemOp> {
692 using OpConversionPattern::OpConversionPattern;
695 matchAndRewrite(MemOp memOp, MemOp::Adaptor adaptor,
696 ConversionPatternRewriter &rewriter)
const override {
697 auto loc = memOp.getLoc();
699 MemRefType::get({ShapedType::kDynamic}, rewriter.getI8Type());
700 auto bufferPtrType = MemRefType::get({1}, bufferType);
701 auto memVisibility = memOp.getVisibility();
704 mlir::StringAttr visAttr;
706 visAttr = mlir::StringAttr::get(memOp->getContext(),
"public");
708 visAttr = mlir::StringAttr::get(memOp->getContext(),
"private");
710 visAttr = mlir::StringAttr::get(memOp->getContext(),
"nested");
712 auto memPtr = rewriter.replaceOpWithNewOp<memref::GlobalOp>(
713 memOp, memOp.getSymNameAttr(), visAttr, TypeAttr::get(bufferPtrType),
714 rewriter.getUnitAttr(),
715 UnitAttr{}, IntegerAttr{});
716 auto initializerName = (memPtr.getSymName() +
"::initializer").str();
717 auto memInitializer =
718 func::FuncOp::create(rewriter, loc, initializerName,
720 memInitializer->setAttr(rewriter.getStringAttr(
"initializer"),
721 rewriter.getUnitAttr());
722 auto *initializerBody = memInitializer.addEntryBlock();
723 auto sip = rewriter.saveInsertionPoint();
724 rewriter.setInsertionPointToStart(initializerBody);
725 auto memRefPtr = memref::GetGlobalOp::create(
726 rewriter, loc, MemRefType::get({1}, bufferType), memPtr.getSymName());
727 auto alloc = memref::AllocOp::create(
729 MemRefType::get({memOp.getLimits().getMin()}, rewriter.getI8Type()));
731 memref::CastOp::create(rewriter, loc, bufferType, alloc.getResult());
733 memref::StoreOp::create(rewriter, loc, castOp.getResult(),
734 memRefPtr.getResult(),
ValueRange{idx.getResult()});
735 func::ReturnOp::create(rewriter, loc);
736 rewriter.restoreInsertionPoint(sip);
737 func::CallOp::create(rewriter, loc, memInitializer);
742inline TypedAttr getInitializerAttr(
Type t) {
744 "This helper is intended to use with int and float types");
746 return IntegerAttr::get(t, 0);
748 return FloatAttr::get(t, 0.);
752struct WasmLocalConversion : OpConversionPattern<LocalOp> {
753 using OpConversionPattern::OpConversionPattern;
755 matchAndRewrite(LocalOp localOp, LocalOp::Adaptor adaptor,
756 ConversionPatternRewriter &rewriter)
const override {
757 auto alloca = rewriter.replaceOpWithNewOp<memref::AllocaOp>(
760 auto initializer = arith::ConstantOp::create(
761 rewriter, localOp->getLoc(),
762 getInitializerAttr(localOp.getResult().getType().getElementType()));
763 memref::StoreOp::create(rewriter, localOp->getLoc(),
764 initializer.getResult(), alloca.getResult());
769struct WasmLocalGetConversion : OpConversionPattern<LocalGetOp> {
770 using OpConversionPattern::OpConversionPattern;
772 matchAndRewrite(LocalGetOp localGetOp, LocalGetOp::Adaptor adaptor,
773 ConversionPatternRewriter &rewriter)
const override {
774 rewriter.replaceOpWithNewOp<memref::LoadOp>(
775 localGetOp, localGetOp.getResult().
getType(), adaptor.getLocalVar(),
781struct WasmLocalSetConversion : OpConversionPattern<LocalSetOp> {
782 using OpConversionPattern::OpConversionPattern;
784 matchAndRewrite(LocalSetOp localSetOp, LocalSetOp::Adaptor adaptor,
785 ConversionPatternRewriter &rewriter)
const override {
786 rewriter.replaceOpWithNewOp<memref::StoreOp>(
787 localSetOp, adaptor.getValue(), adaptor.getLocalVar(),
ValueRange{});
792struct WasmLocalTeeConversion : OpConversionPattern<LocalTeeOp> {
793 using OpConversionPattern::OpConversionPattern;
795 matchAndRewrite(LocalTeeOp localTeeOp, LocalTeeOp::Adaptor adaptor,
796 ConversionPatternRewriter &rewriter)
const override {
797 memref::StoreOp::create(rewriter, localTeeOp->getLoc(), adaptor.getValue(),
798 adaptor.getLocalVar());
799 rewriter.replaceOp(localTeeOp, adaptor.getValue());
804struct WasmReturnOpConversion : OpConversionPattern<ReturnOp> {
805 using OpConversionPattern::OpConversionPattern;
808 matchAndRewrite(ReturnOp returnOp, ReturnOp::Adaptor adaptor,
809 ConversionPatternRewriter &rewriter)
const override {
810 rewriter.replaceOpWithNewOp<func::ReturnOp>(returnOp,
811 adaptor.getOperands());
816struct WasmSelectOpConversion : OpConversionPattern<SelectOp> {
817 using OpConversionPattern::OpConversionPattern;
820 matchAndRewrite(SelectOp selectOp, SelectOp::Adaptor adaptor,
821 ConversionPatternRewriter &rewriter)
const override {
822 auto loc = selectOp.getLoc();
824 arith::ConstantOp::create(rewriter, loc, rewriter.getI32IntegerAttr(0));
825 auto flag = arith::CmpIOp::create(rewriter, loc, arith::CmpIPredicate::ne,
826 adaptor.getCondition(), zero.getResult());
827 rewriter.replaceOpWithNewOp<arith::SelectOp>(selectOp, flag.getResult(),
828 adaptor.getTrueValue(),
829 adaptor.getFalseValue());
834struct RaiseWasmMLIRPass :
public impl::RaiseWasmMLIRBase<RaiseWasmMLIRPass> {
835 void runOnOperation()
override {
837 target.addIllegalDialect<WasmSSADialect>();
838 target.addLegalDialect<arith::ArithDialect, BuiltinDialect,
839 cf::ControlFlowDialect, func::FuncDialect,
840 memref::MemRefDialect, math::MathDialect>();
843 tc.addConversion([](Type type) -> std::optional<Type> {
return type; });
844 tc.addConversion([](LocalRefType type) -> std::optional<Type> {
845 return MemRefType::get({}, type.getElementType());
847 tc.addTargetMaterialization([](OpBuilder &builder, MemRefType destType,
849 if (values.size() != 1 ||
850 values.front().
getType() != destType.getElementType())
852 auto localVar = memref::AllocaOp::create(builder, loc, destType);
853 memref::StoreOp::create(builder, loc, values.front(),
854 localVar.getResult());
855 return localVar.getResult();
859 llvm::DenseMap<StringAttr, StringAttr> idxSymToImportSym{};
860 auto *topOp = getOperation();
861 topOp->walk([&idxSymToImportSym,
this](ImportOpInterface importOp) {
862 auto const qualifiedImportName = importOp.getQualifiedImportName();
863 auto qualNameAttr = StringAttr::get(&
getContext(), qualifiedImportName);
864 idxSymToImportSym.insert(
865 std::make_pair(importOp.getSymbolName(), qualNameAttr));
868 if (
failed(applyFullConversion(topOp,
target, std::move(patterns))))
869 return signalPassFailure();
871 auto symTable = SymbolTable{topOp};
872 for (
auto &[oldName, newName] : idxSymToImportSym) {
873 if (
failed(symTable.rename(oldName, newName)))
874 return signalPassFailure();
890 WasmCallOpConversion,
891 WasmCeilOpConversion,
893 WasmConstOpConversion,
894 WasmConvertSOpConversion,
895 WasmConvertUOpConversion,
896 WasmCopySignOpConversion,
898 WasmDemoteOpConversion,
899 WasmDivFPOpConversion,
900 WasmDivSIOpConversion,
901 WasmDivUIOpConversion,
904 WasmExtendLowBitsOpConversion,
905 WasmExtendSOpConversion,
906 WasmExtendUOpConversion,
907 WasmFloorOpConversion,
908 WasmFuncImportOpConversion,
909 WasmFuncOpConversion,
911 WasmGeSIOpConversion,
912 WasmGeUIOpConversion,
913 WasmGlobalImportOpConverter,
914 WasmGlobalSetOpConversion,
915 WasmGlobalWithConstInitConversion,
916 WasmGlobalWithGetGlobalInitConversion,
918 WasmGtSIOpConversion,
919 WasmGtUIOpConversion,
921 WasmLeSIOpConversion,
922 WasmLeUIOpConversion,
924 WasmLocalGetConversion,
925 WasmLocalSetConversion,
926 WasmLocalTeeConversion,
928 WasmLtSIOpConversion,
929 WasmLtUIOpConversion,
931 WasmMemoryOpConversion,
937 WasmPopCntOpConversion,
938 WasmPromoteOpConversion,
939 WasmReinterpretOpConversion,
940 WasmRemSIOpConversion,
941 WasmRemUIOpConversion,
942 WasmReturnOpConversion,
943 WasmRotlOpConversion,
944 WasmRotrOpConversion,
945 WasmSelectOpConversion,
947 WasmShRSOpConversion,
948 WasmShRUOpConversion,
949 WasmSqrtOpConversion,
951 WasmTruncOpConversion,
952 WasmWrapOpConversion,
959 return std::make_unique<RaiseWasmMLIRPass>();
static Type getElementType(Type type)
Determine the element type of type.
std::unique_ptr< Pass > createRaiseWasmMLIRPass()
static void rewrite(DataFlowSolver &solver, MLIRContext *context, MutableArrayRef< Region > initialRegions)
Rewrite the given regions using the computing analysis.
ValueTypeRange< BlockArgListType > getArgumentTypes()
Return a range containing the types of the arguments for this block.
unsigned getNumArguments()
static DenseElementsAttr get(ShapedType type, ArrayRef< Attribute > values)
Constructs a dense elements attribute from an array of element values.
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.
@ Public
The symbol is public and may be referenced anywhere internal or external to the visible references in...
@ Private
The symbol is private and may only be referenced by SymbolRefAttrs local to the operations within the...
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
bool isFloat() const
Return true if this is an float type (with the specified width).
bool isInteger() const
Return true if this is an integer type (with the specified width).
bool isIntOrFloat() const
Return true if this is an integer (of any signedness) or a float type.
unsigned getIntOrFloatBitWidth() const
Return the bit width of an integer or a float type, assert failure on other types.
type_range getType() const
Location getLoc() const
Return the location of this value.
static WalkResult advance()
static WalkResult interrupt()
static ConstantIndexOp create(OpBuilder &builder, Location location, int64_t value)
Include the generated interface declarations.
void populateRaiseWasmMLIRConversionPatterns(TypeConverter &, RewritePatternSet &)
Collect a set of patterns to convert from the Wasm dialect to standard dialects.
Type getType(OpFoldResult ofr)
Returns the int type of the integer in ofr.
InFlightDiagnostic emitError(Location loc)
Utility method to emit an error message using this location.