Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions mlir/include/mlir/Conversion/Passes.h
Original file line number Diff line number Diff line change
Expand Up @@ -81,6 +81,7 @@
#include "mlir/Conversion/VectorToSCF/VectorToSCF.h"
#include "mlir/Conversion/VectorToSPIRV/VectorToSPIRVPass.h"
#include "mlir/Conversion/VectorToXeGPU/VectorToXeGPU.h"
#include "mlir/Conversion/WasmSSAMLIRToEmbedder/WasmSSAMLIRToEmbedder.h"

namespace mlir {

Expand Down
15 changes: 13 additions & 2 deletions mlir/include/mlir/Conversion/Passes.td
Original file line number Diff line number Diff line change
Expand Up @@ -1503,8 +1503,19 @@ def RaiseWasmMLIR : Pass<"raise-wasm-mlir"> {
let summary = "Convert Wasm dialect to a group of dialect as a bridge to LLVM MLIR conversion";
let dependentDialects = [
"func::FuncDialect", "arith::ArithDialect", "cf::ControlFlowDialect",
"memref::MemRefDialect", "vector::VectorDialect", "wasmssa::WasmSSADialect",
"math::MathDialect"
"LLVM::LLVMDialect", "math::MathDialect", "memref::MemRefDialect",
"vector::VectorDialect", "wasmssa::WasmSSADialect"
];
}

//===----------------------------------------------------------------------===//
// WasmSSAMLIRToEmbedder
//===----------------------------------------------------------------------===//

def WasmSSAMLIRToEmbedder : Pass<"wasm-mlir-to-embedder"> {
let summary = "Convert Wasm operations that needs to interact with an embedder to embedder specific constructs.";
let dependentDialects = [
"func::FuncDialect", "LLVM::LLVMDialect", "wasmssa::WasmSSADialect",
];
}

Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,30 @@
//===- WasmSSAMLIRToEmbedder.h - Convert wasm mlir ops to embedder specific constructs -*- C++ -*-===//
//
// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.
// See https://llvm.org/LICENSE.txt for license information.
// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
//
//===----------------------------------------------------------------------===//

#ifndef MLIR_CONVERSION_WASMSSAMLIRTOEMBEDDER_WASMSSAMLIRTOEMBEDDER_H
#define MLIR_CONVERSION_WASMSSAMLIRTOEMBEDDER_WASMSSAMLIRTOEMBEDDER_H

#include "mlir/IR/PatternMatch.h"
#include "mlir/Transforms/DialectConversion.h"

namespace mlir {
class Pass;
class RewritePatternSet;

#define GEN_PASS_DECL_WASMSSAMLIRTOEMBEDDER
#include "mlir/Conversion/Passes.h.inc"

/// Collect a set of patterns to convert from the Wasm dialect to embedder constructs.
void populateWasmSSAMLIRToEmbedderConversionPatterns(TypeConverter&, RewritePatternSet &);

/// Create a pass to convert ops from WasmDialect to embedder specific constructs
std::unique_ptr<Pass> createWasmSSAMLIRToEmbedderPass();

} // namespace mlir

#endif // MLIR_CONVERSION_WASMSSAMLIRTOEMBEDDER_WASMSSAMLIRTOEMBEDDER_H
Original file line number Diff line number Diff line change
Expand Up @@ -183,4 +183,4 @@ def WasmSSAConstantExprInterface :
];
}

#endif // WEBASSEMBLYSSA_INTERFACES
#endif // WEBASSEMBLY_INTERFACES
Original file line number Diff line number Diff line change
Expand Up @@ -344,6 +344,11 @@ def WasmSSA_TableImportOp : WasmSSA_Op<"import_table", [Symbol, WasmSSAImportOpI
"wasmssa::TableType":$type)>];
}

def WasmSSA_TrapOp : WasmSSA_Op<"trap", []> {
let summary = "Trap placeholder operation";
let assemblyFormat = "attr-dict";
}

def WasmSSA_ReturnOp : WasmSSA_Op<"return", [Terminator]> {
let summary = "Return from the current function frame";
let arguments = (ins Variadic<WasmSSA_ValType>: $operands);
Expand All @@ -356,8 +361,8 @@ def WasmSSA_ReturnOp : WasmSSA_Op<"return", [Terminator]> {
// ---- Numeric ops

class WasmSSA_BinaryNumericalOp<string mnemonic, string summaryStr,
list<Type> validOpTypes> :
WasmSSA_Op<mnemonic, [AllTypesMatch<["lhs", "rhs", "result"]>]> {
list<Type> validOpTypes, list<Trait> traits = []> :
WasmSSA_Op<mnemonic, !listconcat([AllTypesMatch<["lhs", "rhs", "result"]>], traits)> {
let summary = summaryStr;
let arguments = (ins AnyTypeOf<validOpTypes>:$lhs, AnyTypeOf<validOpTypes>:$rhs);
let results = (outs AnyTypeOf<validOpTypes>:$result);
Expand Down
1 change: 1 addition & 0 deletions mlir/lib/Conversion/CMakeLists.txt
Original file line number Diff line number Diff line change
Expand Up @@ -75,3 +75,4 @@ add_subdirectory(VectorToLLVM)
add_subdirectory(VectorToSCF)
add_subdirectory(VectorToSPIRV)
add_subdirectory(VectorToXeGPU)
add_subdirectory(WasmSSAMLIRToEmbedder)
143 changes: 116 additions & 27 deletions mlir/lib/Conversion/RaiseWasm/RaiseWasmMLIR.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -14,12 +14,15 @@
#include "mlir/Conversion/RaiseWasm/RaiseWasmMLIR.h"


#include "llvm/Support/Casting.h"
#include "llvm/Support/LogicalResult.h"
#include "mlir/Dialect/Arith/IR/Arith.h"
#include "mlir/Dialect/ControlFlow/IR/ControlFlowOps.h"
#include "mlir/Dialect/Func/IR/FuncOps.h"
#include "mlir/Dialect/Math/IR/Math.h"
#include "mlir/Dialect/MemRef/IR/MemRef.h"
#include "mlir/Dialect/LLVMIR/LLVMTypes.h"
#include "mlir/Dialect/LLVMIR/LLVMDialect.h"
#include "mlir/Dialect/Vector/IR/VectorOps.h"
#include "mlir/Dialect/WebAssemblySSA/IR/WebAssemblySSA.h"
#include "mlir/IR/BuiltinDialect.h"
Expand All @@ -40,33 +43,117 @@ using namespace mlir::wasmssa;

namespace {

template <typename SourceOp, typename TargetIntOp, typename TargetFPOp>
struct IntFPDispatchMappingConversion : OpConversionPattern<SourceOp> {
template <typename SourceOp, typename TargetOp>
struct FpArithBinaryOpConversion : OpConversionPattern<SourceOp> {
using OpConversionPattern<SourceOp>::OpConversionPattern;

LogicalResult
matchAndRewrite(SourceOp srcOp, typename SourceOp::Adaptor adaptor,
ConversionPatternRewriter &rewriter) const override {
Type type = srcOp.getRhs().getType();
if (type.isInteger()) {
rewriter.replaceOpWithNewOp<TargetIntOp>(srcOp, srcOp->getResultTypes(),
adaptor.getOperands());
return success();
}
if (!type.isFloat())
return failure();
rewriter.replaceOpWithNewOp<TargetFPOp>(srcOp, srcOp->getResultTypes(),

// If Add / Sub / Mul are using their integer variant, it can generate
// an integer overflow trap. Special handling for those cases is necessary.
if (!srcOp.getRhs().getType().isFloat())
return rewriter.notifyMatchFailure(
srcOp->getLoc(),
"This pattern only handles operations on integer operands");
rewriter.replaceOpWithNewOp<TargetOp>(srcOp, srcOp->getResultTypes(),
adaptor.getOperands());
return success();
}
};

using WasmAddOpConversion =
IntFPDispatchMappingConversion<AddOp, arith::AddIOp, arith::AddFOp>;
using WasmMulOpConversion =
IntFPDispatchMappingConversion<MulOp, arith::MulIOp, arith::MulFOp>;
using WasmSubOpConversion =
IntFPDispatchMappingConversion<SubOp, arith::SubIOp, arith::SubFOp>;
using WasmAddFPOpConversion =
FpArithBinaryOpConversion<AddOp, arith::AddFOp>;
using WasmMulFpOpConversion =
FpArithBinaryOpConversion<MulOp, arith::MulFOp>;
using WasmSubFpOpConversion =
FpArithBinaryOpConversion<SubOp, arith::SubFOp>;

template <typename SourceOp, typename IntrinsicOp>
struct TrapIntArithOpConversion : OpConversionPattern<SourceOp> {
using OpConversionPattern<SourceOp>::OpConversionPattern;

LogicalResult
matchAndRewrite(SourceOp wasmOp, typename SourceOp::Adaptor adaptor,
ConversionPatternRewriter &rewriter) const override {
if (!wasmOp.getRhs().getType().isInteger())
return rewriter.notifyMatchFailure(
wasmOp->getLoc(),
"This pattern only handles operations on integer operands");
auto loc = wasmOp.getLoc();
auto addResType = wasmOp.getResult().getType();
auto overflowMarkerType = rewriter.getIntegerType(1);
auto intrinsicResType = LLVM::LLVMStructType::getLiteral(
this->getContext(), {addResType, overflowMarkerType});
auto intrinsicCall = rewriter.create<IntrinsicOp>(
loc, intrinsicResType, wasmOp.getLhs(), wasmOp.getRhs());
auto result =
cast<TypedValue<LLVM::LLVMStructType>>(intrinsicCall.getResult());
auto addResult = rewriter.create<LLVM::ExtractValueOp>(
loc, addResType, result,
rewriter.getDenseI64ArrayAttr({0}));
auto overflowMarker = rewriter.create<LLVM::ExtractValueOp>(
loc, overflowMarkerType, result,
rewriter.getDenseI64ArrayAttr({1}));
auto *curBlock = wasmOp->getBlock();
auto callTrap = rewriter.create<TrapOp>(loc);
Block *trapBlock =
rewriter.splitBlock(curBlock, Block::iterator(callTrap.getOperation()));
Block *normalBlock =
rewriter.splitBlock(trapBlock, ++Block::iterator{callTrap});
rewriter.setInsertionPointAfter(callTrap);
rewriter.create<cf::BranchOp>(loc, normalBlock);
rewriter.setInsertionPointAfter(overflowMarker);
rewriter.create<cf::CondBranchOp>(loc, overflowMarker.getRes(),
trapBlock, ValueRange{}, normalBlock,
ValueRange{});
rewriter.replaceOp(wasmOp, addResult.getRes());
return success();
};
};

using WasmAddIOpConversion =
TrapIntArithOpConversion<AddOp, LLVM::UAddWithOverflowOp>;
using WasmMulIOpConversion =
TrapIntArithOpConversion<MulOp, LLVM::UMulWithOverflowOp>;
using WasmSubIOpConversion =
TrapIntArithOpConversion<SubOp, LLVM::USubWithOverflowOp>;

template<typename WasmOpType, typename ArithOpName>
struct DivOpConversion : OpConversionPattern<WasmOpType> {
using OpConversionPattern<WasmOpType>::OpConversionPattern;
LogicalResult
matchAndRewrite(WasmOpType divOp, typename WasmOpType::Adaptor adaptor,
ConversionPatternRewriter &rewriter) const override {
auto loc = divOp.getLoc();
auto addResType = divOp.getResult().getType();
auto divByZeroMarkerType = rewriter.getIntegerType(1);
auto zero = rewriter.create<arith::ConstantOp>(
loc, rewriter.getIntegerAttr(addResType, 0));
auto isDividerZero = rewriter.create<arith::CmpIOp>(
loc, divByZeroMarkerType, arith::CmpIPredicate::eq,
adaptor.getRhs(), zero.getResult());
auto *curBlock = divOp->getBlock();
auto callTrap = rewriter.create<TrapOp>(loc);
Block *trapBlock =
rewriter.splitBlock(curBlock, Block::iterator(callTrap.getOperation()));
Block *normalBlock =
rewriter.splitBlock(trapBlock, ++Block::iterator{callTrap});
rewriter.setInsertionPointAfter(callTrap);
rewriter.create<cf::BranchOp>(loc, normalBlock);
rewriter.setInsertionPointAfter(isDividerZero);
rewriter.create<cf::CondBranchOp>(loc, isDividerZero.getResult(),
trapBlock, ValueRange{}, normalBlock,
ValueRange{});
rewriter.setInsertionPointToStart(normalBlock);
rewriter.replaceOpWithNewOp<ArithOpName>(divOp, adaptor.getLhs(), adaptor.getRhs());
return success();
};
};

using WasmDivSIConversion = DivOpConversion<DivSIOp, arith::DivSIOp>;
using WasmDivUIConversion = DivOpConversion<DivUIOp, arith::DivUIOp>;

/// Convert a k-ary source operation \p SourceOp into an operation \p TargetOp.
/// Both \p SourceOp and \p TargetOp must have the same number of operands.
Expand All @@ -85,16 +172,14 @@ struct OpMappingConversion : OpConversionPattern<SourceOp> {

using WasmAndOpConversion = OpMappingConversion<AndOp, arith::AndIOp>;
using WasmCeilOpConversion = OpMappingConversion<CeilOp, math::CeilOp>;
using WasmDivFPOpConversion = OpMappingConversion<DivOp, arith::DivFOp>;
/// TODO: SIToFP and UIToFP don't allow specification of the floating point
/// rounding mode
using WasmConvertSOpConversion =
OpMappingConversion<ConvertSOp, arith::SIToFPOp>;
using WasmConvertUOpConversion =
OpMappingConversion<ConvertUOp, arith::UIToFPOp>;
using WasmDemoteOpConversion = OpMappingConversion<DemoteOp, arith::TruncFOp>;
using WasmDivFPOpConversion = OpMappingConversion<DivOp, arith::DivFOp>;
using WasmDivSIOpConversion = OpMappingConversion<DivSIOp, arith::DivSIOp>;
using WasmDivUIOpConversion = OpMappingConversion<DivUIOp, arith::DivUIOp>;
using WasmExtendSOpConversion =
OpMappingConversion<ExtendSI32Op, arith::ExtSIOp>;
using WasmExtendUOpConversion =
Expand Down Expand Up @@ -792,10 +877,11 @@ struct WasmReturnOpConversion : OpConversionPattern<ReturnOp> {
struct RaiseWasmMLIRPass : public impl::RaiseWasmMLIRBase<RaiseWasmMLIRPass> {
void runOnOperation() override {
ConversionTarget target{getContext()};
target.addIllegalDialect<WasmSSADialect>();
target.addLegalDialect<arith::ArithDialect, BuiltinDialect,
cf::ControlFlowDialect, func::FuncDialect,
memref::MemRefDialect, math::MathDialect>();
memref::MemRefDialect, LLVM::LLVMDialect,
math::MathDialect>();
target.addLegalOp<TrapOp>();
RewritePatternSet patterns(&getContext());
TypeConverter tc{};
tc.addConversion([](Type type) -> std::optional<Type> { return type; });
Expand Down Expand Up @@ -843,7 +929,8 @@ void mlir::populateRaiseWasmMLIRConversionPatterns(
patternSet
.add<
WasmAbsOpConversion,
WasmAddOpConversion,
WasmAddFPOpConversion,
WasmAddIOpConversion,
WasmAndOpConversion,
WasmCallOpConversion,
WasmCeilOpConversion,
Expand All @@ -855,8 +942,8 @@ void mlir::populateRaiseWasmMLIRConversionPatterns(
WasmCtzOpConversion,
WasmDemoteOpConversion,
WasmDivFPOpConversion,
WasmDivSIOpConversion,
WasmDivUIOpConversion,
WasmDivSIConversion,
WasmDivUIConversion,
WasmEqOpConversion,
WasmEqzOpConversion,
WasmExtendLowBitsOpConversion,
Expand Down Expand Up @@ -887,7 +974,8 @@ void mlir::populateRaiseWasmMLIRConversionPatterns(
WasmMaxOpConversion,
WasmMemoryOpConversion,
WasmMinOpConversion,
WasmMulOpConversion,
WasmMulFpOpConversion,
WasmMulIOpConversion,
WasmNeOpConversion,
WasmNegOpConversion,
WasmOrOpConversion,
Expand All @@ -903,7 +991,8 @@ void mlir::populateRaiseWasmMLIRConversionPatterns(
WasmShRSOpConversion,
WasmShRUOpConversion,
WasmSqrtOpConversion,
WasmSubOpConversion,
WasmSubFpOpConversion,
WasmSubIOpConversion,
WasmTruncOpConversion,
WasmWrapOpConversion,
WasmXOrOpConversion
Expand Down
14 changes: 14 additions & 0 deletions mlir/lib/Conversion/WasmSSAMLIRToEmbedder/CMakeLists.txt
Original file line number Diff line number Diff line change
@@ -0,0 +1,14 @@
add_mlir_conversion_library(MLIRWasmSSAToEmbedder
WasmSSAMLIRToEmbedder.cpp

ADDITIONAL_HEADER_DIRS
${MLIR_MAIN_INCLUDE_DIR}/mlir/Conversion/WasmSSAMLIRToEmbedder

DEPENDS
MLIRConversionPassIncGen

LINK_LIBS PUBLIC
MLIRFuncDialect
MLIRTransforms
MLIRWasmSSADialect
)
Loading