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 OpTy>
44static typename OpTy::Properties getDefaultProperties(
MLIRContext *context) {
45 typename OpTy::Properties properties{};
46 if constexpr (!std::is_same_v<typename OpTy::Properties, EmptyProperties>)
47 OpTy::populateDefaultProperties(
48 OperationName(OpTy::getOperationName(), context), properties);
52template <
typename SourceOp,
typename TargetIntOp,
typename TargetFPOp>
53struct IntFPDispatchMappingConversion : OpConversionPattern<SourceOp> {
54 using OpConversionPattern<SourceOp>::OpConversionPattern;
57 matchAndRewrite(SourceOp srcOp,
typename SourceOp::Adaptor adaptor,
58 ConversionPatternRewriter &rewriter)
const override {
59 Type type = srcOp.getRhs().getType();
61 rewriter.replaceOpWithNewOp<TargetIntOp>(
62 srcOp, srcOp->getResultTypes(), adaptor.getOperands(),
63 getDefaultProperties<TargetIntOp>(rewriter.getContext()),
64 ArrayRef<NamedAttribute>{});
69 rewriter.replaceOpWithNewOp<TargetFPOp>(
70 srcOp, srcOp->getResultTypes(), adaptor.getOperands(),
71 getDefaultProperties<TargetFPOp>(rewriter.getContext()),
72 ArrayRef<NamedAttribute>{});
77using WasmAddOpConversion =
78 IntFPDispatchMappingConversion<AddOp, arith::AddIOp, arith::AddFOp>;
79using WasmMulOpConversion =
80 IntFPDispatchMappingConversion<MulOp, arith::MulIOp, arith::MulFOp>;
81using WasmSubOpConversion =
82 IntFPDispatchMappingConversion<SubOp, arith::SubIOp, arith::SubFOp>;
86template <
typename SourceOp,
typename TargetOp>
87struct OpMappingConversion : OpConversionPattern<SourceOp> {
88 using OpConversionPattern<SourceOp>::OpConversionPattern;
91 matchAndRewrite(SourceOp srcOp,
typename SourceOp::Adaptor adaptor,
92 ConversionPatternRewriter &rewriter)
const override {
93 rewriter.replaceOpWithNewOp<TargetOp>(
94 srcOp, srcOp->getResultTypes(), adaptor.getOperands(),
95 getDefaultProperties<TargetOp>(rewriter.getContext()),
96 ArrayRef<NamedAttribute>{});
101using WasmAndOpConversion = OpMappingConversion<AndOp, arith::AndIOp>;
102using WasmCeilOpConversion = OpMappingConversion<CeilOp, math::CeilOp>;
105using WasmConvertSOpConversion =
106 OpMappingConversion<ConvertSOp, arith::SIToFPOp>;
107using WasmConvertUOpConversion =
108 OpMappingConversion<ConvertUOp, arith::UIToFPOp>;
109using WasmDemoteOpConversion = OpMappingConversion<DemoteOp, arith::TruncFOp>;
110using WasmDivFPOpConversion = OpMappingConversion<DivOp, arith::DivFOp>;
111using WasmDivSIOpConversion = OpMappingConversion<DivSIOp, arith::DivSIOp>;
112using WasmDivUIOpConversion = OpMappingConversion<DivUIOp, arith::DivUIOp>;
113using WasmExtendSOpConversion =
114 OpMappingConversion<ExtendSI32Op, arith::ExtSIOp>;
115using WasmExtendUOpConversion =
116 OpMappingConversion<ExtendUI32Op, arith::ExtUIOp>;
117using WasmFloorOpConversion = OpMappingConversion<FloorOp, math::FloorOp>;
118using WasmMaxOpConversion = OpMappingConversion<MaxOp, arith::MaximumFOp>;
119using WasmMinOpConversion = OpMappingConversion<MinOp, arith::MinimumFOp>;
120using WasmOrOpConversion = OpMappingConversion<OrOp, arith::OrIOp>;
121using WasmPromoteOpConversion = OpMappingConversion<PromoteOp, arith::ExtFOp>;
122using WasmRemSIOpConversion = OpMappingConversion<RemSIOp, arith::RemSIOp>;
123using WasmRemUIOpConversion = OpMappingConversion<RemUIOp, arith::RemUIOp>;
124using WasmReinterpretOpConversion =
125 OpMappingConversion<ReinterpretOp, arith::BitcastOp>;
126using WasmShLOpConversion = OpMappingConversion<ShLOp, arith::ShLIOp>;
127using WasmShRSOpConversion = OpMappingConversion<ShRSOp, arith::ShRSIOp>;
128using WasmShRUOpConversion = OpMappingConversion<ShRUOp, arith::ShRUIOp>;
129using WasmXOrOpConversion = OpMappingConversion<XOrOp, arith::XOrIOp>;
130using WasmNegOpConversion = OpMappingConversion<NegOp, arith::NegFOp>;
131using WasmCopySignOpConversion =
132 OpMappingConversion<CopySignOp, math::CopySignOp>;
133using WasmClzOpConversion =
134 OpMappingConversion<ClzOp, math::CountLeadingZerosOp>;
135using WasmCtzOpConversion =
136 OpMappingConversion<CtzOp, math::CountTrailingZerosOp>;
137using WasmPopCntOpConversion = OpMappingConversion<PopCntOp, math::CtPopOp>;
138using WasmAbsOpConversion = OpMappingConversion<AbsOp, math::AbsFOp>;
139using WasmTruncOpConversion = OpMappingConversion<TruncOp, math::TruncOp>;
140using WasmSqrtOpConversion = OpMappingConversion<SqrtOp, math::SqrtOp>;
141using WasmWrapOpConversion = OpMappingConversion<WrapOp, arith::TruncIOp>;
163template <
typename SourceOp,
typename LHSShiftOp,
typename RHSShiftOp>
164struct RotateOpConversion : OpConversionPattern<SourceOp> {
165 using OpConversionPattern<SourceOp>::OpConversionPattern;
168 matchAndRewrite(SourceOp srcOp,
typename SourceOp::Adaptor adaptor,
169 ConversionPatternRewriter &rewriter)
const override {
170 const Type ty = srcOp->getResultTypes()[0];
171 const Location loc = srcOp->getLoc();
172 const Value val = adaptor.getVal();
173 const Value bits = adaptor.getBits();
177 auto cstWidthMinusOne =
178 ConstOp::create(rewriter, loc, IntegerAttr::get(ty, width - 1));
182 auto orLHS = LHSShiftOp::create(
184 AndOp::create(rewriter, loc, bits, cstWidthMinusOne));
188 auto orRHS = RHSShiftOp::create(
191 AndOp::create(rewriter, loc,
193 SubOp::create(rewriter, loc,
194 ConstOp::create(rewriter, loc,
195 IntegerAttr::get(ty, 0)),
201 rewriter.replaceOpWithNewOp<OrOp>(srcOp, orLHS, orRHS);
206using WasmRotrOpConversion = RotateOpConversion<RotrOp, ShRUOp, ShLOp>;
207using WasmRotlOpConversion = RotateOpConversion<RotlOp, ShLOp, ShRUOp>;
209template <
typename SourceOp,
typename TargetOp,
typename AttrType,
210 typename ValType, ValType flag>
211struct ComparisonOpConversion : OpConversionPattern<SourceOp> {
212 using OpConversionPattern<SourceOp>::OpConversionPattern;
215 matchAndRewrite(SourceOp srcOp,
typename SourceOp::Adaptor adaptor,
216 ConversionPatternRewriter &rewriter)
const override {
218 TargetOp::create(rewriter, srcOp.getLoc(), rewriter.getI1Type(),
219 AttrType::get(rewriter.getContext(), flag),
220 adaptor.getLhs(), adaptor.getRhs())
222 rewriter.replaceOpWithNewOp<arith::ExtUIOp>(srcOp, rewriter.getI32Type(),
229template <
typename SourceOp, arith::CmpFPredicate compFlag>
230using FPComparisonConversion =
231 ComparisonOpConversion<SourceOp, arith::CmpFOp, arith::CmpFPredicateAttr,
232 arith::CmpFPredicate, compFlag>;
234template <
typename SourceOp, arith::CmpIPredicate compFlag>
235using IntComparisonConversion =
236 ComparisonOpConversion<SourceOp, arith::CmpIOp, arith::CmpIPredicateAttr,
237 arith::CmpIPredicate, compFlag>;
239using WasmLtSIOpConversion =
240 IntComparisonConversion<LtSIOp, arith::CmpIPredicate::slt>;
241using WasmLeSIOpConversion =
242 IntComparisonConversion<LeSIOp, arith::CmpIPredicate::sle>;
243using WasmGtSIOpConversion =
244 IntComparisonConversion<GtSIOp, arith::CmpIPredicate::sgt>;
245using WasmGeSIOpConversion =
246 IntComparisonConversion<GeSIOp, arith::CmpIPredicate::sge>;
247using WasmLtUIOpConversion =
248 IntComparisonConversion<LtUIOp, arith::CmpIPredicate::ult>;
249using WasmLeUIOpConversion =
250 IntComparisonConversion<LeUIOp, arith::CmpIPredicate::ule>;
251using WasmGtUIOpConversion =
252 IntComparisonConversion<GtUIOp, arith::CmpIPredicate::ugt>;
253using WasmGeUIOpConversion =
254 IntComparisonConversion<GeUIOp, arith::CmpIPredicate::uge>;
255using WasmLtOpConversion =
256 FPComparisonConversion<LtOp, arith::CmpFPredicate::OLT>;
257using WasmLeOpConversion =
258 FPComparisonConversion<LeOp, arith::CmpFPredicate::OLE>;
259using WasmGtOpConversion =
260 FPComparisonConversion<GtOp, arith::CmpFPredicate::OGT>;
261using WasmGeOpConversion =
262 FPComparisonConversion<GeOp, arith::CmpFPredicate::OGE>;
264template <
typename SourceOp, arith::CmpIPredicate IntFlag,
265 arith::CmpFPredicate FloatFlag>
266struct IntFpComparisonOpConversion : OpConversionPattern<SourceOp> {
267 using OpConversionPattern<SourceOp>::OpConversionPattern;
270 matchAndRewrite(SourceOp srcOp,
typename SourceOp::Adaptor adaptor,
271 ConversionPatternRewriter &rewriter)
const override {
272 Value comparisonResult;
273 if (srcOp.getLhs().getType().isInteger())
275 arith::CmpIOp::create(
276 rewriter, srcOp.getLoc(), rewriter.getI1Type(),
277 arith::CmpIPredicateAttr::get(rewriter.getContext(), IntFlag),
278 adaptor.getLhs(), adaptor.getRhs())
280 else if (srcOp.getLhs().getType().isFloat())
282 arith::CmpFOp::create(
283 rewriter, srcOp.getLoc(), rewriter.getI1Type(),
284 arith::CmpFPredicateAttr::get(rewriter.getContext(), FloatFlag),
285 adaptor.getLhs(), adaptor.getRhs())
288 return rewriter.notifyMatchFailure(
289 srcOp.getLoc(),
"Unsupported datatype for comparison OP.");
291 rewriter.replaceOpWithNewOp<arith::ExtUIOp>(srcOp, rewriter.getI32Type(),
297using WasmEqOpConversion =
298 IntFpComparisonOpConversion<EqOp, arith::CmpIPredicate::eq,
299 arith::CmpFPredicate::OEQ>;
300using WasmNeOpConversion =
301 IntFpComparisonOpConversion<NeOp, arith::CmpIPredicate::ne,
302 arith::CmpFPredicate::ONE>;
304struct WasmCallOpConversion : OpConversionPattern<FuncCallOp> {
305 using OpConversionPattern::OpConversionPattern;
308 matchAndRewrite(FuncCallOp funcCallOp, FuncCallOp::Adaptor adaptor,
309 ConversionPatternRewriter &rewriter)
const override {
310 rewriter.replaceOpWithNewOp<func::CallOp>(
311 funcCallOp, funcCallOp.getCallee(), funcCallOp.getResults().getTypes(),
312 funcCallOp.getOperands());
317struct WasmConstOpConversion : OpConversionPattern<ConstOp> {
318 using OpConversionPattern::OpConversionPattern;
321 matchAndRewrite(ConstOp constOp, ConstOp::Adaptor adaptor,
322 ConversionPatternRewriter &rewriter)
const override {
323 rewriter.replaceOpWithNewOp<arith::ConstantOp>(constOp, constOp.getValue());
328struct WasmEqzOpConversion : OpConversionPattern<EqzOp> {
329 using OpConversionPattern::OpConversionPattern;
332 matchAndRewrite(EqzOp eqzOp, EqzOp::Adaptor adaptor,
333 ConversionPatternRewriter &rewriter)
const override {
334 auto loc = eqzOp->getLoc();
335 auto zero = arith::ConstantOp::create(
337 rewriter.getIntegerAttr(adaptor.getInput().getType(), 0))
339 auto cmpRes = arith::CmpIOp::create(
340 rewriter, loc, rewriter.getI1Type(),
341 arith::CmpIPredicateAttr::get(rewriter.getContext(),
342 arith::CmpIPredicate::eq),
343 adaptor.getInput(), zero)
345 rewriter.replaceOpWithNewOp<arith::ExtUIOp>(eqzOp, rewriter.getI32Type(),
352struct WasmExtendLowBitsOpConversion : OpConversionPattern<ExtendLowBitsSOp> {
353 using OpConversionPattern::OpConversionPattern;
356 matchAndRewrite(ExtendLowBitsSOp extendLowBytesSOp,
357 ExtendLowBitsSOp::Adaptor adaptor,
358 ConversionPatternRewriter &rewriter)
const override {
359 auto truncWidth = extendLowBytesSOp.getBitsToTake().getInt();
360 auto truncation = arith::TruncIOp::create(
361 rewriter, extendLowBytesSOp->getLoc(),
362 rewriter.getIntegerType(truncWidth), adaptor.getInput());
363 rewriter.replaceOpWithNewOp<arith::ExtSIOp>(
364 extendLowBytesSOp, extendLowBytesSOp.getResult().
getType(),
365 truncation.getResult());
370struct WasmFuncImportOpConversion : OpConversionPattern<FuncImportOp> {
371 using OpConversionPattern::OpConversionPattern;
374 matchAndRewrite(FuncImportOp funcImportOp, FuncImportOp::Adaptor,
375 ConversionPatternRewriter &rewriter)
const override {
376 auto nFunc = rewriter.replaceOpWithNewOp<func::FuncOp>(
377 funcImportOp, funcImportOp.getSymName(), funcImportOp.getType());
378 nFunc.setVisibility(SymbolTable::Visibility::Private);
383struct WasmFuncOpConversion : OpConversionPattern<FuncOp> {
384 using OpConversionPattern::OpConversionPattern;
391 class CFRewriterVisitor {
393 using branch_to_dest_t = llvm::DenseMap<LabelBranchingOpInterface, Block *>;
394 Value getCompResultAsI1(Value compResult,
395 ConversionPatternRewriter &rewriter) {
396 auto testValue = arith::ConstantOp::create(rewriter, compResult.
getLoc(),
397 rewriter.getI32IntegerAttr(0));
398 auto flag = arith::CmpIOp::create(
399 rewriter, compResult.
getLoc(), rewriter.getIntegerType(1),
400 arith::CmpIPredicate::ne, compResult, testValue)
405 void replaceNestLevelWithBranch(BlockOp blockOp,
406 llvm::ArrayRef<Block *> regionsToEntry,
407 ConversionPatternRewriter &rewriter) {
408 rewriter.replaceOpWithNewOp<cf::BranchOp>(blockOp, regionsToEntry[0],
409 blockOp->getOperands());
412 void replaceNestLevelWithBranch(LoopOp loopOp,
413 llvm::ArrayRef<Block *> regionsToEntry,
414 ConversionPatternRewriter &rewriter) {
415 rewriter.replaceOpWithNewOp<cf::BranchOp>(loopOp, regionsToEntry[0],
416 loopOp->getOperands());
419 void replaceNestLevelWithBranch(IfOp ifOp,
420 llvm::ArrayRef<Block *> regionsToEntry,
421 ConversionPatternRewriter &rewriter) {
423 regionsToEntry.size() == 2 ? regionsToEntry[1] : ifOp.getTarget();
424 auto flag = getCompResultAsI1(ifOp.getCondition(), rewriter);
425 rewriter.replaceOpWithNewOp<cf::CondBranchOp>(
426 ifOp, flag, regionsToEntry[0], ifOp.getInputs(), falseDest,
430 template <
typename LevelType>
432 replaceNestLevelWithBranchWrapper(LabelLevelOpInterface nestingOp,
433 llvm::ArrayRef<Block *> regionsToEntry,
434 ConversionPatternRewriter &rewriter) {
435 auto cast = dyn_cast<LevelType>(nestingOp.getOperation());
438 replaceNestLevelWithBranch(cast, regionsToEntry, rewriter);
442 template <
typename... LevelTypes>
443 LogicalResult inlineNestDispatcher(LabelLevelOpInterface nestingOp,
444 ConversionPatternRewriter &rewriter) {
445 auto sip = rewriter.saveInsertionPoint();
446 Block *blockSuccessor = nestingOp->getSuccessor(0);
447 llvm::SmallVector<Block *, 2> regionEntries;
448 LLVM_DEBUG(llvm::dbgs()
449 <<
"Starting inlining blocks for " << nestingOp <<
"\n";);
450 for (
auto ®ion : nestingOp->getRegions()) {
453 regionEntries.push_back(®ion.front());
455 llvm::SmallVector<LabelLevelOpInterface> nestedOps{
456 region.getOps<LabelLevelOpInterface>()};
457 for (
auto nestedOp : nestedOps) {
458 LLVM_DEBUG(llvm::dbgs() <<
" Found nested op: " << nestedOp);
459 if (
failed(inlineBlocks(nestedOp, rewriter)))
462 rewriter.inlineRegionBefore(region, blockSuccessor);
464 LLVM_DEBUG(llvm::dbgs() <<
"End of region inlining\n");
465 LLVM_DEBUG(llvm::dbgs() <<
"Replacing initial op with branching\n");
466 rewriter.setInsertionPoint(nestingOp);
468 (... || succeeded(replaceNestLevelWithBranchWrapper<LevelTypes>(
469 nestingOp, regionEntries, rewriter))));
470 rewriter.restoreInsertionPoint(sip);
473 "Unable to inline the operation regions.");
478 LogicalResult inlineBlocks(LabelLevelOpInterface nestingOp,
479 ConversionPatternRewriter &rewriter) {
480 return inlineNestDispatcher<BlockOp, IfOp, LoopOp>(nestingOp, rewriter);
483 llvm::FailureOr<Block *> getBlockFor(LabelBranchingOpInterface branchOp) {
484 auto destIter = branchToDest.find(branchOp);
485 if (destIter == branchToDest.end())
486 return branchOp->emitError(
"No indexed label op for this operation: ")
488 return destIter->second;
491 inline void convertBranch(BranchIfOp brOp,
Block *dest,
492 ConversionPatternRewriter &rewriter) {
493 auto flag = getCompResultAsI1(brOp.getCondition(), rewriter);
494 rewriter.replaceOpWithNewOp<cf::CondBranchOp>(
495 brOp, flag, dest, brOp.getInputs(), brOp.getElseSuccessor(),
499 inline void convertBranch(BlockReturnOp brOp,
Block *dest,
500 ConversionPatternRewriter &rewriter) {
501 rewriter.replaceOpWithNewOp<cf::BranchOp>(brOp, dest, brOp.getInputs());
504 template <
typename LevelInterfaceT>
506 convertBranchWrapper(LabelBranchingOpInterface branchOp,
Block *dest,
507 ConversionPatternRewriter &rewriter) {
508 auto cast = dyn_cast<LevelInterfaceT>(branchOp.getOperation());
511 auto sip = rewriter.saveInsertionPoint();
512 rewriter.setInsertionPoint(branchOp);
513 convertBranch(cast, dest, rewriter);
514 rewriter.restoreInsertionPoint(sip);
518 template <
typename... BranchInterfaceT>
519 LogicalResult convertBranchDispatch(LabelBranchingOpInterface branchOp,
520 ConversionPatternRewriter &rewriter) {
521 auto dest = getBlockFor(branchOp);
525 success((... || succeeded(convertBranchWrapper<BranchInterfaceT>(
526 branchOp, *dest, rewriter))));
528 return emitError(branchOp->getLoc(),
"No known converter for op ")
533 LogicalResult convertBranch(LabelBranchingOpInterface branchOp,
534 ConversionPatternRewriter &rewriter) {
535 return convertBranchDispatch<BlockReturnOp, BranchIfOp>(branchOp,
540 branch_to_dest_t branchToDest;
543 CFRewriterVisitor(func::FuncOp func) : func{func} {
544 func.walk([
this](LabelBranchingOpInterface branchOp) {
545 branchToDest.insert({branchOp, branchOp.getTarget()});
548 LogicalResult
rewrite(ConversionPatternRewriter &rewriter) {
549 llvm::SmallVector<LabelLevelOpInterface> nestingOps{
550 func.getOps<LabelLevelOpInterface>()};
551 for (
auto nestingOp : nestingOps)
552 if (
failed(inlineBlocks(nestingOp, rewriter)))
556 func->walk([
this, &rewriter](LabelBranchingOpInterface branchOp) {
557 if (
failed(convertBranch(branchOp, rewriter)))
561 return failure(res.wasInterrupted());
566 matchAndRewrite(FuncOp funcOp, FuncOp::Adaptor adaptor,
567 ConversionPatternRewriter &rewriter)
const override {
569 func::FuncOp::create(rewriter, funcOp->getLoc(), funcOp.getSymName(),
570 funcOp.getFunctionType());
571 rewriter.cloneRegionBefore(funcOp.getBody(), newFunc.getBody(),
572 newFunc.getBody().end());
573 Block *oldEntryBlock = &newFunc.getBody().front();
575 TypeConverter::SignatureConversion sC{oldEntryBlock->
getNumArguments()};
576 auto numArgs = blockArgTypes.size();
577 for (
size_t i = 0; i < numArgs; ++i) {
578 auto argType = dyn_cast<LocalRefType>(blockArgTypes[i]);
581 sC.addInputs(i, argType.getElementType());
584 rewriter.applySignatureConversion(oldEntryBlock, sC, getTypeConverter());
585 rewriter.replaceOp(funcOp, newFunc);
586 CFRewriterVisitor cfRewriter{newFunc};
587 return cfRewriter.rewrite(rewriter);
591struct WasmGlobalImportOpConverter : OpConversionPattern<GlobalImportOp> {
592 using OpConversionPattern::OpConversionPattern;
594 matchAndRewrite(GlobalImportOp gIOp, GlobalImportOp::Adaptor adaptor,
595 ConversionPatternRewriter &rewriter)
const override {
596 auto memrefGOp = rewriter.replaceOpWithNewOp<memref::GlobalOp>(
597 gIOp, gIOp.getSymNameAttr(), rewriter.getStringAttr(
"nested"),
598 TypeAttr::get(MemRefType::get({1}, gIOp.getType())), Attribute{},
601 memrefGOp.setConstant(!gIOp.getIsMutable());
606template <
typename CRTP,
typename OriginOpType>
607struct GlobalOpConverter : OpConversionPattern<GlobalOp> {
608 using OpConversionPattern::OpConversionPattern;
610 matchAndRewrite(GlobalOp globalOp, GlobalOp::Adaptor adaptor,
611 ConversionPatternRewriter &rewriter)
const override {
612 ReturnOp rop = globalOp.getInitTerminator();
614 if (rop->getNumOperands() != 1)
615 return rewriter.notifyMatchFailure(
616 globalOp,
"globalOp initializer should return one value exactly");
619 dyn_cast<OriginOpType>(rop->getOperand(0).getDefiningOp());
622 return rewriter.notifyMatchFailure(
623 globalOp,
"invalid initializer op type for this pattern");
625 return static_cast<CRTP
const *
>(
this)->handleInitializer(
626 globalOp, rewriter, initializerOp);
630struct WasmGlobalWithConstInitConversion
631 : GlobalOpConverter<WasmGlobalWithConstInitConversion, ConstOp> {
632 using GlobalOpConverter::GlobalOpConverter;
633 LogicalResult handleInitializer(GlobalOp globalOp,
634 ConversionPatternRewriter &rewriter,
635 ConstOp constInit)
const {
638 ArrayRef<Attribute>{constInit.getValueAttr()});
639 auto globalReplacement = rewriter.replaceOpWithNewOp<memref::GlobalOp>(
640 globalOp, globalOp.getSymNameAttr(), rewriter.getStringAttr(
"private"),
641 TypeAttr::get(MemRefType::get({1}, globalOp.getType())), initializer,
644 globalReplacement.setConstant(!globalOp.getIsMutable());
649struct WasmGlobalWithGetGlobalInitConversion
650 : GlobalOpConverter<WasmGlobalWithGetGlobalInitConversion, GlobalGetOp> {
651 using GlobalOpConverter::GlobalOpConverter;
652 LogicalResult handleInitializer(GlobalOp globalOp,
653 ConversionPatternRewriter &rewriter,
654 GlobalGetOp constInit)
const {
655 auto globalReplacement = rewriter.replaceOpWithNewOp<memref::GlobalOp>(
656 globalOp, globalOp.getSymNameAttr(), rewriter.getStringAttr(
"private"),
657 TypeAttr::get(MemRefType::get({1}, globalOp.getType())),
658 rewriter.getUnitAttr(),
661 globalReplacement.setConstant(!globalOp.getIsMutable());
662 auto loc = globalOp.getLoc();
663 auto initializerName = (globalOp.getSymName() +
"::initializer").str();
664 auto globalInitializer =
665 func::FuncOp::create(rewriter, loc, initializerName,
667 globalInitializer->setAttr(rewriter.getStringAttr(
"initializer"),
668 rewriter.getUnitAttr());
669 auto *initializerBody = globalInitializer.addEntryBlock();
670 auto sip = rewriter.saveInsertionPoint();
671 rewriter.setInsertionPointToStart(initializerBody);
672 auto srcGlobalPtr = memref::GetGlobalOp::create(
673 rewriter, loc, MemRefType::get({1}, constInit.getType()),
674 constInit.getGlobal());
676 memref::GetGlobalOp::create(rewriter, loc, globalReplacement.getType(),
677 globalReplacement.getSymName());
680 memref::LoadOp::create(rewriter, loc, srcGlobalPtr,
ValueRange{idx});
681 memref::StoreOp::create(rewriter, loc, loadSrc.getResult(),
683 func::ReturnOp::create(rewriter, loc);
684 rewriter.restoreInsertionPoint(sip);
689struct WasmGlobalSetOpConversion : OpConversionPattern<GlobalSetOp> {
690 using OpConversionPattern::OpConversionPattern;
692 matchAndRewrite(GlobalSetOp globalSetOp, GlobalSetOp::Adaptor adaptor,
693 ConversionPatternRewriter &rewriter)
const override {
694 auto loc = globalSetOp.getLoc();
695 auto globalPtr = memref::GetGlobalOp::create(
696 rewriter, loc, MemRefType::get({1}, adaptor.getValue().
getType()),
697 globalSetOp.getGlobal());
699 rewriter.replaceOpWithNewOp<memref::StoreOp>(
700 globalSetOp, adaptor.getValue(), globalPtr.getResult(),
706struct WasmMemoryOpConversion : OpConversionPattern<MemOp> {
707 using OpConversionPattern::OpConversionPattern;
710 matchAndRewrite(MemOp memOp, MemOp::Adaptor adaptor,
711 ConversionPatternRewriter &rewriter)
const override {
712 auto loc = memOp.getLoc();
714 MemRefType::get({ShapedType::kDynamic}, rewriter.getI8Type());
715 auto bufferPtrType = MemRefType::get({1}, bufferType);
716 auto memVisibility = memOp.getVisibility();
719 mlir::StringAttr visAttr;
721 visAttr = mlir::StringAttr::get(memOp->getContext(),
"public");
723 visAttr = mlir::StringAttr::get(memOp->getContext(),
"private");
725 visAttr = mlir::StringAttr::get(memOp->getContext(),
"nested");
727 auto memPtr = rewriter.replaceOpWithNewOp<memref::GlobalOp>(
728 memOp, memOp.getSymNameAttr(), visAttr, TypeAttr::get(bufferPtrType),
729 rewriter.getUnitAttr(),
730 UnitAttr{}, IntegerAttr{});
731 auto initializerName = (memPtr.getSymName() +
"::initializer").str();
732 auto memInitializer =
733 func::FuncOp::create(rewriter, loc, initializerName,
735 memInitializer->setAttr(rewriter.getStringAttr(
"initializer"),
736 rewriter.getUnitAttr());
737 auto *initializerBody = memInitializer.addEntryBlock();
738 auto sip = rewriter.saveInsertionPoint();
739 rewriter.setInsertionPointToStart(initializerBody);
740 auto memRefPtr = memref::GetGlobalOp::create(
741 rewriter, loc, MemRefType::get({1}, bufferType), memPtr.getSymName());
742 auto alloc = memref::AllocOp::create(
744 MemRefType::get({memOp.getLimits().getMin()}, rewriter.getI8Type()));
746 memref::CastOp::create(rewriter, loc, bufferType, alloc.getResult());
748 memref::StoreOp::create(rewriter, loc, castOp.getResult(),
749 memRefPtr.getResult(),
ValueRange{idx.getResult()});
750 func::ReturnOp::create(rewriter, loc);
751 rewriter.restoreInsertionPoint(sip);
752 func::CallOp::create(rewriter, loc, memInitializer);
757inline TypedAttr getInitializerAttr(
Type t) {
759 "This helper is intended to use with int and float types");
761 return IntegerAttr::get(t, 0);
763 return FloatAttr::get(t, 0.);
767struct WasmLocalConversion : OpConversionPattern<LocalOp> {
768 using OpConversionPattern::OpConversionPattern;
770 matchAndRewrite(LocalOp localOp, LocalOp::Adaptor adaptor,
771 ConversionPatternRewriter &rewriter)
const override {
772 auto alloca = rewriter.replaceOpWithNewOp<memref::AllocaOp>(
775 auto initializer = arith::ConstantOp::create(
776 rewriter, localOp->getLoc(),
777 getInitializerAttr(localOp.getResult().getType().getElementType()));
778 memref::StoreOp::create(rewriter, localOp->getLoc(),
779 initializer.getResult(), alloca.getResult());
784struct WasmLocalGetConversion : OpConversionPattern<LocalGetOp> {
785 using OpConversionPattern::OpConversionPattern;
787 matchAndRewrite(LocalGetOp localGetOp, LocalGetOp::Adaptor adaptor,
788 ConversionPatternRewriter &rewriter)
const override {
789 rewriter.replaceOpWithNewOp<memref::LoadOp>(
790 localGetOp, localGetOp.getResult().
getType(), adaptor.getLocalVar(),
796struct WasmLocalSetConversion : OpConversionPattern<LocalSetOp> {
797 using OpConversionPattern::OpConversionPattern;
799 matchAndRewrite(LocalSetOp localSetOp, LocalSetOp::Adaptor adaptor,
800 ConversionPatternRewriter &rewriter)
const override {
801 rewriter.replaceOpWithNewOp<memref::StoreOp>(
802 localSetOp, adaptor.getValue(), adaptor.getLocalVar(),
ValueRange{});
807struct WasmLocalTeeConversion : OpConversionPattern<LocalTeeOp> {
808 using OpConversionPattern::OpConversionPattern;
810 matchAndRewrite(LocalTeeOp localTeeOp, LocalTeeOp::Adaptor adaptor,
811 ConversionPatternRewriter &rewriter)
const override {
812 memref::StoreOp::create(rewriter, localTeeOp->getLoc(), adaptor.getValue(),
813 adaptor.getLocalVar());
814 rewriter.replaceOp(localTeeOp, adaptor.getValue());
819struct WasmReturnOpConversion : OpConversionPattern<ReturnOp> {
820 using OpConversionPattern::OpConversionPattern;
823 matchAndRewrite(ReturnOp returnOp, ReturnOp::Adaptor adaptor,
824 ConversionPatternRewriter &rewriter)
const override {
825 rewriter.replaceOpWithNewOp<func::ReturnOp>(returnOp,
826 adaptor.getOperands());
831struct WasmSelectOpConversion : OpConversionPattern<SelectOp> {
832 using OpConversionPattern::OpConversionPattern;
835 matchAndRewrite(SelectOp selectOp, SelectOp::Adaptor adaptor,
836 ConversionPatternRewriter &rewriter)
const override {
837 auto loc = selectOp.getLoc();
839 arith::ConstantOp::create(rewriter, loc, rewriter.getI32IntegerAttr(0));
840 auto flag = arith::CmpIOp::create(rewriter, loc, arith::CmpIPredicate::ne,
841 adaptor.getCondition(), zero.getResult());
842 rewriter.replaceOpWithNewOp<arith::SelectOp>(selectOp, flag.getResult(),
843 adaptor.getTrueValue(),
844 adaptor.getFalseValue());
850 void runOnOperation()
override {
852 target.addIllegalDialect<WasmSSADialect>();
853 target.addLegalDialect<arith::ArithDialect, BuiltinDialect,
854 cf::ControlFlowDialect, func::FuncDialect,
855 memref::MemRefDialect, math::MathDialect>();
858 tc.addConversion([](Type type) -> std::optional<Type> {
return type; });
859 tc.addConversion([](LocalRefType type) -> std::optional<Type> {
860 return MemRefType::get({}, type.getElementType());
862 tc.addTargetMaterialization([](OpBuilder &builder, MemRefType destType,
864 if (values.size() != 1 ||
865 values.front().
getType() != destType.getElementType())
867 auto localVar = memref::AllocaOp::create(builder, loc, destType);
868 memref::StoreOp::create(builder, loc, values.front(),
869 localVar.getResult());
870 return localVar.getResult();
874 llvm::DenseMap<StringAttr, StringAttr> idxSymToImportSym{};
875 auto *topOp = getOperation();
876 topOp->walk([&idxSymToImportSym,
this](ImportOpInterface importOp) {
877 auto const qualifiedImportName = importOp.getQualifiedImportName();
878 auto qualNameAttr = StringAttr::get(&
getContext(), qualifiedImportName);
879 idxSymToImportSym.insert(
880 std::make_pair(importOp.getSymbolName(), qualNameAttr));
883 if (
failed(applyFullConversion(topOp,
target, std::move(patterns))))
884 return signalPassFailure();
886 auto symTable = SymbolTable{topOp};
887 for (
auto &[oldName, newName] : idxSymToImportSym) {
888 if (
failed(symTable.rename(oldName, newName)))
889 return signalPassFailure();
905 WasmCallOpConversion,
906 WasmCeilOpConversion,
908 WasmConstOpConversion,
909 WasmConvertSOpConversion,
910 WasmConvertUOpConversion,
911 WasmCopySignOpConversion,
913 WasmDemoteOpConversion,
914 WasmDivFPOpConversion,
915 WasmDivSIOpConversion,
916 WasmDivUIOpConversion,
919 WasmExtendLowBitsOpConversion,
920 WasmExtendSOpConversion,
921 WasmExtendUOpConversion,
922 WasmFloorOpConversion,
923 WasmFuncImportOpConversion,
924 WasmFuncOpConversion,
926 WasmGeSIOpConversion,
927 WasmGeUIOpConversion,
928 WasmGlobalImportOpConverter,
929 WasmGlobalSetOpConversion,
930 WasmGlobalWithConstInitConversion,
931 WasmGlobalWithGetGlobalInitConversion,
933 WasmGtSIOpConversion,
934 WasmGtUIOpConversion,
936 WasmLeSIOpConversion,
937 WasmLeUIOpConversion,
939 WasmLocalGetConversion,
940 WasmLocalSetConversion,
941 WasmLocalTeeConversion,
943 WasmLtSIOpConversion,
944 WasmLtUIOpConversion,
946 WasmMemoryOpConversion,
952 WasmPopCntOpConversion,
953 WasmPromoteOpConversion,
954 WasmReinterpretOpConversion,
955 WasmRemSIOpConversion,
956 WasmRemUIOpConversion,
957 WasmReturnOpConversion,
958 WasmRotlOpConversion,
959 WasmRotrOpConversion,
960 WasmSelectOpConversion,
962 WasmShRSOpConversion,
963 WasmShRUOpConversion,
964 WasmSqrtOpConversion,
966 WasmTruncOpConversion,
967 WasmWrapOpConversion,
974 return std::make_unique<RaiseWasmMLIRPass>();
std::unique_ptr< Pass > createRaiseWasmMLIRPass()
static void rewrite(DataFlowSolver &solver, MLIRContext *context, MutableArrayRef< Region > initialRegions)
Rewrite the given regions using the computing analysis.
static Type getElementType(Type type, ArrayRef< int32_t > indices, function_ref< InFlightDiagnostic(StringRef)> emitErrorFn)
Walks the given type hierarchy with the given indices, potentially down to component granularity,...
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 is the top-level object for a collection of MLIR operations.
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.