[Mlir-commits] [mlir] [MLIR][WASM] Introduce full support for raising WASM MLIR to other dialects (PR #205990)
llvmlistbot at llvm.org
llvmlistbot at llvm.org
Thu Jun 25 23:48:35 PDT 2026
llvmorg-github-actions[bot] wrote:
<!--LLVM PR SUMMARY COMMENT-->
@llvm/pr-subscribers-mlir
Author: Ferdinand Lemaire (flemairen6)
<details>
<summary>Changes</summary>
Following https://github.com/llvm/llvm-project/pull/164562 where RaiseWasm was introduced.
This PR completes the support for rewriting the currently supported Wasm MLIR operators to arith, math, cf and memref.
---
Patch is 77.13 KiB, truncated to 20.00 KiB below, full version: https://github.com/llvm/llvm-project/pull/205990.diff
9 Files Affected:
- (modified) mlir/lib/Conversion/RaiseWasm/RaiseWasmMLIR.cpp (+439-1)
- (added) mlir/test/Conversion/RaiseWasm/wasm-blocks-to-cf.mlir (+366)
- (added) mlir/test/Conversion/RaiseWasm/wasm-comparisons-to-arith-cmp.mlir (+421)
- (added) mlir/test/Conversion/RaiseWasm/wasm-eqz-to-arith-cmp.mlir (+28)
- (added) mlir/test/Conversion/RaiseWasm/wasm-extend-to-arith-ext.mlir (+76)
- (added) mlir/test/Conversion/RaiseWasm/wasm-loop-to-cf.mlir (+206)
- (added) mlir/test/Conversion/RaiseWasm/wasm-memory-to-memref.mlir (+14)
- (added) mlir/test/Conversion/RaiseWasm/wasm-rotl-to-arith.mlir (+75)
- (added) mlir/test/Conversion/RaiseWasm/wasm-rotr-to-arith.mlir (+77)
``````````diff
diff --git a/mlir/lib/Conversion/RaiseWasm/RaiseWasmMLIR.cpp b/mlir/lib/Conversion/RaiseWasm/RaiseWasmMLIR.cpp
index 83bfde7032ef8..0e6e0c05b79b6 100644
--- a/mlir/lib/Conversion/RaiseWasm/RaiseWasmMLIR.cpp
+++ b/mlir/lib/Conversion/RaiseWasm/RaiseWasmMLIR.cpp
@@ -20,6 +20,7 @@
#include "mlir/Dialect/MemRef/IR/MemRef.h"
#include "mlir/Dialect/Vector/IR/VectorOps.h"
#include "mlir/Dialect/WasmSSA/IR/WasmSSA.h"
+#include "mlir/Dialect/WasmSSA/IR/WasmSSAInterfaces.h"
#include "mlir/IR/BuiltinAttributes.h"
#include "mlir/IR/BuiltinDialect.h"
#include "mlir/IR/ValueRange.h"
@@ -124,6 +125,162 @@ using WasmTruncOpConversion = OpMappingConversion<TruncOp, math::TruncOp>;
using WasmSqrtOpConversion = OpMappingConversion<SqrtOp, math::SqrtOp>;
using WasmWrapOpConversion = OpMappingConversion<WrapOp, arith::TruncIOp>;
+/// Lower a rotate to a series of bitwise operations. Intended for us
+/// in dialects that do not natively support rotate operations.
+///
+/// Result stays in the wasm dialect. It will then subsequently be lowered to
+/// the target dialect.
+///
+/// The rotate will be lowered to a pattern like so:
+///
+/// (val LHSShiftOp (bits & (width-1))) | (val RHSShiftOp (-bits & (width-1)))
+///
+/// Where LHSShiftOp and RHSShiftOp are shift operations. Concretely,
+///
+/// rotr = (val >> (bits & (width - 1))) | (val << (-bits & (width - 1)))
+/// rotl = (val << (bits & (width - 1))) | (val >> (-bits & (width - 1)))
+///
+/// Using this variant ensures that our rotate is defined in the target dialect.
+///
+/// \p SourceOp - Rotate operation to replace.
+/// \p LHSShiftOp - Shift operation to use on the left-hand side of the OR.
+/// \p RHSShiftOp - Shift operation to use on the right-hand side of the OR.
+template <typename SourceOp, typename LHSShiftOp, typename RHSShiftOp>
+struct RotateOpConversion : OpConversionPattern<SourceOp> {
+ using OpConversionPattern<SourceOp>::OpConversionPattern;
+
+ LogicalResult
+ matchAndRewrite(SourceOp srcOp, typename SourceOp::Adaptor adaptor,
+ ConversionPatternRewriter &rewriter) const override {
+ const Type ty = srcOp->getResultTypes()[0];
+ const Location loc = srcOp->getLoc();
+ const Value val = adaptor.getVal();
+ const Value bits = adaptor.getBits();
+ const unsigned width = ty.getIntOrFloatBitWidth();
+
+ // Materialize (width - 1) for use in both sides of the expression.
+ auto cstWidthMinusOne =
+ ConstOp::create(rewriter, loc, IntegerAttr::get(ty, width - 1));
+
+ // Form the left-hand side of the OR:
+ // (val (lhs shift op) (bits & (width - 1)))
+ auto orLHS = LHSShiftOp::create(
+ rewriter, loc, val, AndOp::create(rewriter, loc, bits, cstWidthMinusOne));
+
+ // Form the right-hand side of the OR:
+ // (val (rhs shift op) (-bits & (width - 1)))
+ auto orRHS = RHSShiftOp::create(
+ rewriter, loc, val,
+ // (-bits & (width - 1))
+ AndOp::create(
+ rewriter, loc,
+ // 0 - bits == -bits
+ SubOp::create(rewriter, loc, ConstOp::create(rewriter, loc, IntegerAttr::get(ty, 0)), bits),
+ cstWidthMinusOne));
+
+ // OR together the two shifts and replace the rotate with the new
+ // expression.
+ rewriter.replaceOpWithNewOp<OrOp>(srcOp, orLHS, orRHS);
+ return success();
+ }
+};
+
+using WasmRotrOpConversion = RotateOpConversion<RotrOp, ShRUOp, ShLOp>;
+using WasmRotlOpConversion = RotateOpConversion<RotlOp, ShLOp, ShRUOp>;
+
+template <typename SourceOp, typename TargetOp, typename AttrType,
+ typename ValType, ValType flag>
+struct ComparisonOpConversion : OpConversionPattern<SourceOp> {
+ using OpConversionPattern<SourceOp>::OpConversionPattern;
+
+ LogicalResult
+ matchAndRewrite(SourceOp srcOp, typename SourceOp::Adaptor adaptor,
+ ConversionPatternRewriter &rewriter) const override {
+ auto cmpRes =
+ TargetOp::create(rewriter, srcOp.getLoc(), rewriter.getI1Type(),
+ AttrType::get(rewriter.getContext(), flag),
+ adaptor.getLhs(), adaptor.getRhs())
+ .getResult();
+ rewriter.replaceOpWithNewOp<arith::ExtUIOp>(srcOp, rewriter.getI32Type(),
+ cmpRes);
+
+ return success();
+ }
+};
+
+template <typename SourceOp, arith::CmpFPredicate compFlag>
+using FPComparisonConversion =
+ ComparisonOpConversion<SourceOp, arith::CmpFOp, arith::CmpFPredicateAttr,
+ arith::CmpFPredicate, compFlag>;
+
+template <typename SourceOp, arith::CmpIPredicate compFlag>
+using IntComparisonConversion =
+ ComparisonOpConversion<SourceOp, arith::CmpIOp, arith::CmpIPredicateAttr,
+ arith::CmpIPredicate, compFlag>;
+
+using WasmLtSIOpConversion =
+ IntComparisonConversion<LtSIOp, arith::CmpIPredicate::slt>;
+using WasmLeSIOpConversion =
+ IntComparisonConversion<LeSIOp, arith::CmpIPredicate::sle>;
+using WasmGtSIOpConversion =
+ IntComparisonConversion<GtSIOp, arith::CmpIPredicate::sgt>;
+using WasmGeSIOpConversion =
+ IntComparisonConversion<GeSIOp, arith::CmpIPredicate::sge>;
+using WasmLtUIOpConversion =
+ IntComparisonConversion<LtUIOp, arith::CmpIPredicate::ult>;
+using WasmLeUIOpConversion =
+ IntComparisonConversion<LeUIOp, arith::CmpIPredicate::ule>;
+using WasmGtUIOpConversion =
+ IntComparisonConversion<GtUIOp, arith::CmpIPredicate::ugt>;
+using WasmGeUIOpConversion =
+ IntComparisonConversion<GeUIOp, arith::CmpIPredicate::uge>;
+using WasmLtOpConversion =
+ FPComparisonConversion<LtOp, arith::CmpFPredicate::OLT>;
+using WasmLeOpConversion =
+ FPComparisonConversion<LeOp, arith::CmpFPredicate::OLE>;
+using WasmGtOpConversion =
+ FPComparisonConversion<GtOp, arith::CmpFPredicate::OGT>;
+using WasmGeOpConversion =
+ FPComparisonConversion<GeOp, arith::CmpFPredicate::OGE>;
+
+template <typename SourceOp, arith::CmpIPredicate IntFlag,
+ arith::CmpFPredicate FloatFlag>
+struct IntFpComparisonOpConversion : OpConversionPattern<SourceOp> {
+ using OpConversionPattern<SourceOp>::OpConversionPattern;
+
+ LogicalResult
+ matchAndRewrite(SourceOp srcOp, typename SourceOp::Adaptor adaptor,
+ ConversionPatternRewriter &rewriter) const override {
+ Value comparisonResult;
+ if (srcOp.getLhs().getType().isInteger())
+ comparisonResult =
+ arith::CmpIOp::create(rewriter, srcOp.getLoc(), rewriter.getI1Type(),
+ arith::CmpIPredicateAttr::get(rewriter.getContext(), IntFlag),
+ adaptor.getLhs(), adaptor.getRhs())
+ .getResult();
+ else if (srcOp.getLhs().getType().isFloat())
+ comparisonResult =
+ arith::CmpFOp::create(rewriter, srcOp.getLoc(), rewriter.getI1Type(),
+ arith::CmpFPredicateAttr::get(rewriter.getContext(), FloatFlag),
+ adaptor.getLhs(), adaptor.getRhs())
+ .getResult();
+ else
+ return rewriter.notifyMatchFailure(
+ srcOp.getLoc(), "Unsupported datatype for comparison OP.");
+
+ rewriter.replaceOpWithNewOp<arith::ExtUIOp>(srcOp, rewriter.getI32Type(),
+ comparisonResult);
+ return success();
+ }
+};
+
+using WasmEqOpConversion =
+ IntFpComparisonOpConversion<EqOp, arith::CmpIPredicate::eq,
+ arith::CmpFPredicate::OEQ>;
+using WasmNeOpConversion =
+ IntFpComparisonOpConversion<NeOp, arith::CmpIPredicate::ne,
+ arith::CmpFPredicate::ONE>;
+
struct WasmCallOpConversion : OpConversionPattern<FuncCallOp> {
using OpConversionPattern::OpConversionPattern;
@@ -148,6 +305,44 @@ struct WasmConstOpConversion : OpConversionPattern<ConstOp> {
}
};
+struct WasmEqzOpConversion : OpConversionPattern<EqzOp> {
+ using OpConversionPattern::OpConversionPattern;
+
+ LogicalResult
+ matchAndRewrite(EqzOp eqzOp, EqzOp::Adaptor adaptor,
+ ConversionPatternRewriter &rewriter) const override {
+ auto loc = eqzOp->getLoc();
+ auto zero =
+ arith::ConstantOp::create(rewriter, loc, rewriter.getIntegerAttr(adaptor.getInput().getType(), 0))
+ .getResult();
+ auto cmpRes = arith::CmpIOp::create(rewriter, loc, rewriter.getI1Type(),
+ arith::CmpIPredicateAttr::get(
+ rewriter.getContext(), arith::CmpIPredicate::eq),
+ adaptor.getInput(), zero)
+ .getResult();
+ rewriter.replaceOpWithNewOp<arith::ExtUIOp>(eqzOp, rewriter.getI32Type(),
+ cmpRes);
+
+ return success();
+ }
+};
+
+struct WasmExtendLowBitsOpConversion : OpConversionPattern<ExtendLowBitsSOp> {
+ using OpConversionPattern::OpConversionPattern;
+
+ LogicalResult
+ matchAndRewrite(ExtendLowBitsSOp extendLowBytesSOp,
+ ExtendLowBitsSOp::Adaptor adaptor,
+ ConversionPatternRewriter &rewriter) const override {
+ auto truncWidth = extendLowBytesSOp.getBitsToTake().getInt();
+ auto truncation = arith::TruncIOp::create(rewriter, extendLowBytesSOp->getLoc(), rewriter.getIntegerType(truncWidth), adaptor.getInput());
+ rewriter.replaceOpWithNewOp<arith::ExtSIOp>(
+ extendLowBytesSOp, extendLowBytesSOp.getResult().getType(),
+ truncation.getResult());
+ return success();
+ }
+};
+
struct WasmFuncImportOpConversion : OpConversionPattern<FuncImportOp> {
using OpConversionPattern::OpConversionPattern;
@@ -163,6 +358,183 @@ struct WasmFuncImportOpConversion : OpConversionPattern<FuncImportOp> {
struct WasmFuncOpConversion : OpConversionPattern<FuncOp> {
using OpConversionPattern::OpConversionPattern;
+ ///
+ /// Control flow conversion needs shared state for tracking which block
+ /// corresponds to which operation at which level.
+ ///
+ /// This class handles such tracking and performs the conversion of control
+ /// flow related ops contained in a function.
+ class CFRewriterVisitor {
+ private:
+ using branch_to_dest_t =
+ llvm::DenseMap<LabelBranchingOpInterface, Block *>;
+ Value getCompResultAsI1(Value compResult,
+ ConversionPatternRewriter &rewriter) {
+ auto testValue = arith::ConstantOp::create(rewriter, compResult.getLoc(), rewriter.getI32IntegerAttr(0));
+ auto flag = arith::CmpIOp::create(rewriter, compResult.getLoc(), rewriter.getIntegerType(1), arith::CmpIPredicate::ne, compResult, testValue)
+ .getResult();
+ return flag;
+ }
+
+ void replaceNestLevelWithBranch(BlockOp blockOp,
+ llvm::ArrayRef<Block *> regionsToEntry,
+ ConversionPatternRewriter &rewriter) {
+ rewriter.replaceOpWithNewOp<cf::BranchOp>(blockOp, regionsToEntry[0],
+ blockOp->getOperands());
+ }
+
+ void replaceNestLevelWithBranch(LoopOp loopOp,
+ llvm::ArrayRef<Block *> regionsToEntry,
+ ConversionPatternRewriter &rewriter) {
+ rewriter.replaceOpWithNewOp<cf::BranchOp>(loopOp, regionsToEntry[0],
+ loopOp->getOperands());
+ }
+
+ void replaceNestLevelWithBranch(IfOp ifOp,
+ llvm::ArrayRef<Block *> regionsToEntry,
+ ConversionPatternRewriter &rewriter) {
+ Block *falseDest =
+ regionsToEntry.size() == 2 ? regionsToEntry[1] : ifOp.getTarget();
+ auto flag = getCompResultAsI1(ifOp.getCondition(), rewriter);
+ rewriter.replaceOpWithNewOp<cf::CondBranchOp>(
+ ifOp, flag, regionsToEntry[0], ifOp.getInputs(), falseDest,
+ ifOp.getInputs());
+ }
+
+ template <typename LevelType>
+ LogicalResult
+ replaceNestLevelWithBranchWrapper(LabelLevelOpInterface nestingOp,
+ llvm::ArrayRef<Block *> regionsToEntry,
+ ConversionPatternRewriter &rewriter) {
+ auto cast = dyn_cast<LevelType>(nestingOp.getOperation());
+ if (!cast)
+ return failure();
+ replaceNestLevelWithBranch(cast, regionsToEntry, rewriter);
+ return success();
+ }
+
+ template <typename... LevelTypes>
+ LogicalResult inlineNestDispatcher(LabelLevelOpInterface nestingOp,
+ ConversionPatternRewriter &rewriter) {
+ auto sip = rewriter.saveInsertionPoint();
+ Block *blockSuccessor = nestingOp->getSuccessor(0);
+ llvm::SmallVector<Block *, 2> regionEntries;
+ LLVM_DEBUG(llvm::dbgs()
+ << "Starting inlining blocks for " << nestingOp << "\n";);
+ for (auto ®ion : nestingOp->getRegions()) {
+ if (region.empty())
+ continue;
+ regionEntries.push_back(®ion.front());
+ /// Inline blocks of nested ops
+ llvm::SmallVector<LabelLevelOpInterface> nestedOps{
+ region.getOps<LabelLevelOpInterface>()};
+ for (auto nestedOp : nestedOps) {
+ LLVM_DEBUG(llvm::dbgs() << " Found nested op: " << nestedOp);
+ if (failed(inlineBlocks(nestedOp, rewriter)))
+ return failure();
+ }
+ rewriter.inlineRegionBefore(region, blockSuccessor);
+ }
+ LLVM_DEBUG(llvm::dbgs() << "End of region inlining\n");
+ LLVM_DEBUG(llvm::dbgs() << "Replacing initial op with branching\n");
+ rewriter.setInsertionPoint(nestingOp);
+ auto res = success(
+ (... || succeeded(replaceNestLevelWithBranchWrapper<LevelTypes>(
+ nestingOp, regionEntries, rewriter))));
+ rewriter.restoreInsertionPoint(sip);
+ if (failed(res))
+ return emitError(nestingOp->getLoc(),
+ "Unable to inline the operation regions.");
+ return success();
+ }
+
+ /// Take a nesting level defining op and inline it in the parent region.
+ LogicalResult inlineBlocks(LabelLevelOpInterface nestingOp,
+ ConversionPatternRewriter &rewriter) {
+ return inlineNestDispatcher<BlockOp, IfOp, LoopOp>(nestingOp, rewriter);
+ }
+
+ llvm::FailureOr<Block *> getBlockFor(LabelBranchingOpInterface branchOp) {
+ auto destIter = branchToDest.find(branchOp);
+ if (destIter == branchToDest.end())
+ return branchOp->emitError("No indexed label op for this operation: ")
+ << branchOp;
+ return destIter->second;
+ }
+
+ inline void convertBranch(BranchIfOp brOp, Block *dest,
+ ConversionPatternRewriter &rewriter) {
+ auto flag = getCompResultAsI1(brOp.getCondition(), rewriter);
+ rewriter.replaceOpWithNewOp<cf::CondBranchOp>(
+ brOp, flag, dest, brOp.getInputs(), brOp.getElseSuccessor(),
+ ValueRange{});
+ }
+
+ inline void convertBranch(BlockReturnOp brOp, Block *dest,
+ ConversionPatternRewriter &rewriter) {
+ rewriter.replaceOpWithNewOp<cf::BranchOp>(brOp, dest, brOp.getInputs());
+ }
+
+ template <typename LevelInterfaceT>
+ inline LogicalResult
+ convertBranchWrapper(LabelBranchingOpInterface branchOp, Block *dest,
+ ConversionPatternRewriter &rewriter) {
+ auto cast = dyn_cast<LevelInterfaceT>(branchOp.getOperation());
+ if (!cast)
+ return failure();
+ auto sip = rewriter.saveInsertionPoint();
+ rewriter.setInsertionPoint(branchOp);
+ convertBranch(cast, dest, rewriter);
+ rewriter.restoreInsertionPoint(sip);
+ return success();
+ }
+
+ template <typename... BranchInterfaceT>
+ LogicalResult convertBranchDispatch(LabelBranchingOpInterface branchOp,
+ ConversionPatternRewriter &rewriter) {
+ auto dest = getBlockFor(branchOp);
+ if (failed(dest))
+ return failure();
+ auto res =
+ success((... || succeeded(convertBranchWrapper<BranchInterfaceT>(
+ branchOp, *dest, rewriter))));
+ if (failed(res))
+ return emitError(branchOp->getLoc(), "No known converter for op ")
+ << branchOp;
+ return res;
+ }
+
+ LogicalResult convertBranch(LabelBranchingOpInterface branchOp,
+ ConversionPatternRewriter &rewriter) {
+ return convertBranchDispatch<BlockReturnOp, BranchIfOp>(branchOp,
+ rewriter);
+ }
+
+ func::FuncOp func;
+ branch_to_dest_t branchToDest;
+
+ public:
+ CFRewriterVisitor(func::FuncOp func) : func{func} {
+ func.walk([this](LabelBranchingOpInterface branchOp) {
+ branchToDest.insert({branchOp, branchOp.getTarget()});
+ });
+ }
+ LogicalResult rewrite(ConversionPatternRewriter &rewriter) {
+ llvm::SmallVector<LabelLevelOpInterface> nestingOps{
+ func.getOps<LabelLevelOpInterface>()};
+ for (auto nestingOp : nestingOps)
+ if (failed(inlineBlocks(nestingOp, rewriter)))
+ return failure();
+
+ auto res =
+ func->walk([this, &rewriter](LabelBranchingOpInterface branchOp) {
+ if (failed(convertBranch(branchOp, rewriter)))
+ return WalkResult::interrupt();
+ return WalkResult::advance();
+ });
+ return failure(res.wasInterrupted());
+ }
+ };
LogicalResult
matchAndRewrite(FuncOp funcOp, FuncOp::Adaptor adaptor,
@@ -185,7 +557,8 @@ struct WasmFuncOpConversion : OpConversionPattern<FuncOp> {
rewriter.applySignatureConversion(oldEntryBlock, sC, getTypeConverter());
rewriter.replaceOp(funcOp, newFunc);
- return success();
+ CFRewriterVisitor cfRewriter{newFunc};
+ return cfRewriter.rewrite(rewriter);
}
};
@@ -287,6 +660,52 @@ struct WasmGlobalWithGetGlobalInitConversion
}
};
+struct WasmMemoryOpConversion : OpConversionPattern<MemOp> {
+ using OpConversionPattern::OpConversionPattern;
+
+ LogicalResult
+ matchAndRewrite(MemOp memOp, MemOp::Adaptor adaptor,
+ ConversionPatternRewriter &rewriter) const override {
+ auto loc = memOp.getLoc();
+ auto bufferType =
+ MemRefType::get({ShapedType::kDynamic}, rewriter.getI8Type());
+ auto bufferPtrType = MemRefType::get({1}, bufferType);
+ auto memVisibility = memOp.getVisibility();
+ // Convert to StringAttr since memref::GlobalOp expects visibility as a string attribute.
+ mlir::StringAttr visAttr;
+ if (memVisibility == mlir::SymbolTable::Visibility::Public)
+ visAttr = mlir::StringAttr::get(memOp->getContext(), "public");
+ else if (memVisibility == mlir::SymbolTable::Visibility::Private)
+ visAttr = mlir::StringAttr::get(memOp->getContext(), "private");
+ else
+ visAttr = mlir::StringAttr::get(memOp->getContext(), "nested");
+
+ auto memPtr = rewriter.replaceOpWithNewOp<memref::GlobalOp>(
+ memOp, memOp.getSymNameAttr(), visAttr,
+ TypeAttr::get(bufferPtrType), /*initialValue*/ rewriter.getUnitAttr(),
+ /*constant*/ UnitAttr{}, /*alignment*/ IntegerAttr{});
+ auto initializerName = (memPtr.getSymName() + "::initializer").str();
+ auto memInitializer = func::FuncOp::create(rewriter, loc, initializerName,
+ FunctionType::get(getContext(), {}, {}));
+ memInitializer->setAttr(rewriter.getStringAttr("initializer"),
+ rewriter.getUnitAttr());
+ auto *initializerBody = memInitializer.addEntryBlock();
+ auto sip = rewriter.saveInsertionPoint();
+ rewriter.setInsertionPointToStart(initializerBody);
+ auto memRefPtr = memref::GetGlobalOp::create(rewriter, loc, MemRefType::get({1}, bufferType), memPtr.getSymName());
+ auto alloc = memref::AllocOp::create(rewriter, loc, MemRefType::get({memOp.getLimits().getMin()}, rewriter.getI8Type()));
+ auto castOp =
+ memref::CastOp::create(rewriter, loc, bufferType, alloc.getResult());
+ auto idx = arith::ConstantIndexOp::create(rewriter, loc, 0);
+ memref::StoreOp::create(rewriter, loc, castOp.getResult(),
+ memRefPtr.getResult(), ValueRange{idx.getResult()});
+ func...
[truncated]
``````````
</details>
https://github.com/llvm/llvm-project/pull/205990
More information about the Mlir-commits
mailing list