26#define GEN_PASS_DEF_TOSADOWNGRADE1P1TO1P0PASS
27#include "mlir/Dialect/Tosa/Transforms/Passes.h.inc"
40 LogicalResult matchAndRewrite(tosa::CastOp op,
41 PatternRewriter &rewriter)
const override {
42 const Value input = op.getInput();
49 const bool isFp32ToBool =
50 inputElemType == f32Type && outputElemType == i1Type;
51 const bool isBoolToFp32 =
52 inputElemType == i1Type && outputElemType == f32Type;
54 if (!isFp32ToBool && !isBoolToFp32)
56 "expected cast between bool and f32");
58 const Type outputType = op.getType();
60 const Type intermediateType = cast<TensorType>(outputType).clone(i8Type);
62 auto inner = tosa::CastOp::create(rewriter, op.getLoc(), intermediateType,
65 tosa::CastOp::create(rewriter, op.getLoc(), outputType,
66 inner.getOutput(),
false);
67 rewriter.
replaceOp(op, outer.getOutput());
76 LogicalResult matchAndRewrite(tosa::GatherOp op,
77 PatternRewriter &rewriter)
const override {
78 const Value values = op.getValues();
79 const Value
indices = op.getIndices();
81 const Type valuesType = values.
getType();
82 const Type resultType = op.getType();
89 op,
"expected values of bool type and indices of i32 type");
92 const Type valuesI8Type = cast<TensorType>(valuesType).clone(i8Type);
93 const Type resultI8Type = cast<TensorType>(resultType).clone(i8Type);
95 auto valuesToI8 = tosa::CastOp::create(rewriter, op.getLoc(), valuesI8Type,
97 auto gatherI8 = tosa::GatherOp::create(rewriter, op.getLoc(), resultI8Type,
98 valuesToI8.getOutput(),
indices);
100 tosa::CastOp::create(rewriter, op.getLoc(), resultType,
101 gatherI8.getOutput(),
false);
102 rewriter.
replaceOp(op, i8ToBool.getOutput());
111 LogicalResult matchAndRewrite(tosa::ScatterOp op,
112 PatternRewriter &rewriter)
const override {
113 const Value valuesIn = op.getValuesIn();
114 const Value
indices = op.getIndices();
116 const Type valuesInType = valuesIn.
getType();
117 const Type i1Type = rewriter.
getI1Type();
122 op,
"expected values of bool type and indices of i32 type");
124 const Value input = op.getInput();
125 const Type inputType = input.
getType();
126 const Type resultType = op.getType();
128 const Type i8Type = rewriter.
getI8Type();
129 const Type valuesInI8Type = cast<TensorType>(valuesInType).clone(i8Type);
130 const Type inputI8Type = cast<TensorType>(inputType).clone(i8Type);
131 const Type resultI8Type = cast<TensorType>(resultType).clone(i8Type);
134 tosa::CastOp::create(rewriter, op.getLoc(), valuesInI8Type, valuesIn,
136 auto inputToI8 = tosa::CastOp::create(rewriter, op.getLoc(), inputI8Type,
138 auto scatterI8 = tosa::ScatterOp::create(
139 rewriter, op.getLoc(), resultI8Type, valuesInToI8.getOutput(),
indices,
140 inputToI8.getOutput());
141 auto i8ToBool = tosa::CastOp::create(rewriter, op.getLoc(), resultType,
142 scatterI8.getValuesOut(),
144 rewriter.
replaceOp(op, i8ToBool.getOutput());
149static LogicalResult isMatMulTTypeCompatibleForDowngrade(tosa::MatMulTOp op) {
152 const Type outputElementType =
155 if (aElementType != bElementType)
158 if (isa<BlockScaledType>(aElementType) || isa<BlockScaledType>(bElementType))
161 if ((aElementType.
isF16() && outputElementType.
isF16()) ||
162 (aElementType.
isF16() && outputElementType.
isF32()) ||
163 (aElementType.
isF32() && outputElementType.
isF32()) ||
164 (aElementType.
isBF16() && outputElementType.
isF32()) ||
167 (isa<Float8E5M2Type>(aElementType) && outputElementType.
isF16()) ||
168 (isa<Float8E4M3FNType>(aElementType) && outputElementType.
isF16()))
178 LogicalResult matchAndRewrite(tosa::MatMulTOp op,
179 PatternRewriter &rewriter)
const override {
180 if (
failed(isMatMulTTypeCompatibleForDowngrade(op)))
182 op,
"expected 1.0-compatible matmul_t element types");
184 const Type aType = op.getA().getType();
185 const Type bType = op.getB().getType();
186 const ShapeAdaptor aShape(aType);
187 const ShapeAdaptor bShape(bType);
188 if (!aShape.hasRank() || !bShape.hasRank())
191 const int64_t dSize = bShape.getDimSize(0);
192 const int64_t nSize = aShape.getDimSize(0);
197 if (ShapedType::isDynamic(dSize) ||
198 (dSize == 1 && ShapedType::isDynamic(nSize)))
200 op,
"expected known batch size for broadcast");
202 const int64_t wSize = bShape.getDimSize(1);
203 const int64_t cSize = bShape.getDimSize(2);
204 const Location loc = op.getLoc();
205 const RankedTensorType transposedBType =
206 cast<RankedTensorType>(bType).clone({dSize, cSize, wSize});
208 tosa::TransposeOp::create(rewriter, loc, transposedBType, op.getB(),
210 Value matMulB = transpose.getOutput();
213 if (dSize == 1 && nSize != 1) {
214 const RankedTensorType tiledBType =
215 cast<RankedTensorType>(bType).clone({nSize, cSize, wSize});
218 tosa::TileOp::create(rewriter, loc, tiledBType, matMulB, multiples);
219 matMulB =
tile.getOutput();
222 auto matmul = tosa::MatMulOp::create(rewriter, loc, op.getType(), op.getA(),
223 matMulB, op.getAZp(), op.getBZp());
224 rewriter.
replaceOp(op, matmul.getOutput());
229struct TosaDowngrade1p1To1p0Pass
230 :
public tosa::impl::TosaDowngrade1p1To1p0PassBase<
231 TosaDowngrade1p1To1p0Pass> {
234 void runOnOperation()
override {
236 func::FuncOp func = getOperation();
238 RewritePatternSet patterns(&context);
239 patterns.add<BoolFp32CastRewrite, BoolGatherRewrite, BoolScatterRewrite,
240 MatMulTRewrite>(&context);
241 FrozenRewritePatternSet frozenPatterns(std::move(patterns));
244 return signalPassFailure();
DenseI32ArrayAttr getDenseI32ArrayAttr(ArrayRef< int32_t > values)
virtual void replaceOp(Operation *op, ValueRange newValues)
Replace the results of the given (original) operation with the specified list of values (replacements...
std::enable_if_t<!std::is_convertible< CallbackT, Twine >::value, LogicalResult > notifyMatchFailure(Location loc, CallbackT &&reasonCallback)
Used to notify the listener that the IR failed to be rewritten because of a match failure,...
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
bool isInteger() const
Return true if this is an integer type (with the specified width).
Type getType() const
Return the type of this value.
Type getStorageElementTypeOrSelf(Type type)
Value getTosaConstShape(ImplicitLocOpBuilder &builder, llvm::ArrayRef< int64_t > shape)
Include the generated interface declarations.
LogicalResult applyPatternsGreedily(Region ®ion, const FrozenRewritePatternSet &patterns, GreedyRewriteConfig config=GreedyRewriteConfig(), bool *changed=nullptr)
Rewrite ops in the given region, which must be isolated from above, by repeatedly applying the highes...
Type getElementTypeOrSelf(Type type)
Return the element type or return the type itself.
SmallVector< Loops, 8 > tile(ArrayRef< scf::ForOp > forOps, ArrayRef< Value > sizes, ArrayRef< scf::ForOp > targets)
Performs tiling fo imperfectly nested loops (with interchange) by strip-mining the forOps by sizes an...
OpRewritePattern is a wrapper around RewritePattern that allows for matching and rewriting against an...
OpRewritePattern(MLIRContext *context, PatternBenefit benefit=1, ArrayRef< StringRef > generatedNames={})
Patterns must specify the root operation name they match against, and can also specify the benefit of...