Skip to content

Commit 0e83d42

Browse files
authored
[compiler] support cpu codegen of l2_norm (#516)
as title
1 parent 9a4e01f commit 0e83d42

2 files changed

Lines changed: 109 additions & 0 deletions

File tree

compiler/lib/Dialect/mhlo/Transforms/DecomposeMhloCustomCallOps.cpp

Lines changed: 105 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -116,6 +116,108 @@ struct DecomposeByteIRSoftmax : public OpRewritePattern<mhlo::CustomCallOp> {
116116
}
117117
};
118118

119+
struct DecomposeByteIRL2Norm : public OpRewritePattern<mhlo::CustomCallOp> {
120+
using OpRewritePattern<mhlo::CustomCallOp>::OpRewritePattern;
121+
122+
LogicalResult matchAndRewrite(mhlo::CustomCallOp op,
123+
PatternRewriter &rewriter) const override {
124+
if (op.getCallTargetName() != getL2NormName()) {
125+
return failure();
126+
}
127+
128+
Value operand = op.getOperand(0);
129+
RankedTensorType inType = cast<RankedTensorType>(operand.getType());
130+
mlir::FloatType fpType = cast<FloatType>(inType.getElementType());
131+
132+
DictionaryAttr byteirAttrs =
133+
cast<DictionaryAttr>(op->getAttr(getCustomCallAttrName()));
134+
if (!byteirAttrs)
135+
return failure();
136+
auto axisAttr = cast<ArrayAttr>(byteirAttrs.get("axis"));
137+
if (axisAttr.size() != 1) {
138+
return op->emitError("only support 1 axis");
139+
}
140+
auto axis = cast<IntegerAttr>(axisAttr[0]).getInt();
141+
142+
auto epsAttr = cast<FloatAttr>(byteirAttrs.get("epsilon"));
143+
APFloat eps = epsAttr.getValue();
144+
bool losesInfo;
145+
auto status = eps.convert(fpType.getFloatSemantics(),
146+
APFloat::rmNearestTiesToEven, &losesInfo);
147+
if (losesInfo) {
148+
op->emitRemark("loses info when eps convert to input type");
149+
}
150+
epsAttr = rewriter.getFloatAttr(fpType, eps);
151+
152+
bool epsOutsideSqrt = false;
153+
if (byteirAttrs.contains("eps_outside_sqrt")) {
154+
epsOutsideSqrt =
155+
cast<BoolAttr>(byteirAttrs.get("eps_outside_sqrt")).getValue();
156+
}
157+
158+
Value pow2 = rewriter.create<mhlo::MulOp>(op.getLoc(), operand, operand);
159+
Value reduce;
160+
{
161+
SmallVector<int64_t> reduceResultShape(inType.getShape());
162+
reduceResultShape.erase(reduceResultShape.begin() + axis);
163+
RankedTensorType reduceResultType =
164+
RankedTensorType::get(reduceResultShape, fpType);
165+
166+
Value initValue = rewriter.create<mhlo::ConstantOp>(
167+
op.getLoc(), DenseElementsAttr::get(
168+
RankedTensorType::get({}, fpType),
169+
{APFloat::getZero(fpType.getFloatSemantics())}));
170+
auto reduceOp = rewriter.create<mhlo::ReduceOp>(
171+
op.getLoc(), reduceResultType, pow2, initValue,
172+
rewriter.getI64TensorAttr({axis}));
173+
174+
Block &block = reduceOp.getBody().emplaceBlock();
175+
OpBuilder::InsertionGuard guard(rewriter);
176+
rewriter.setInsertionPointToStart(&block);
177+
auto blockValArgumentType =
178+
RankedTensorType::get({}, inType.getElementType());
179+
block.addArgument(blockValArgumentType, op->getLoc());
180+
block.addArgument(blockValArgumentType, op->getLoc());
181+
auto *firstValArg = block.args_begin();
182+
auto *secondValArg = std::next(firstValArg);
183+
Value result = rewriter.create<mhlo::AddOp>(op->getLoc(), *firstValArg,
184+
*secondValArg);
185+
rewriter.create<mhlo::ReturnOp>(op->getLoc(), result);
186+
187+
reduce = reduceOp.getResults()[0];
188+
}
189+
190+
Value epsValue = rewriter.create<mhlo::ConstantOp>(
191+
op.getLoc(),
192+
DenseElementsAttr::get(RankedTensorType::get({}, fpType), epsAttr));
193+
epsValue = rewriter.create<mhlo::DynamicBroadcastInDimOp>(
194+
op.getLoc(), reduce.getType(), epsValue,
195+
rewriter.create<shape::ShapeOfOp>(op.getLoc(), reduce),
196+
rewriter.getI64TensorAttr({}));
197+
Value sqrt;
198+
if (epsOutsideSqrt) {
199+
sqrt = rewriter.create<mhlo::SqrtOp>(op.getLoc(), reduce);
200+
sqrt = rewriter.create<mhlo::MaxOp>(op.getLoc(), sqrt, epsValue);
201+
} else {
202+
sqrt = rewriter.create<mhlo::MaxOp>(op.getLoc(), reduce, epsValue);
203+
sqrt = rewriter.create<mhlo::SqrtOp>(op.getLoc(), sqrt);
204+
}
205+
206+
SmallVector broadcastDim =
207+
llvm::to_vector(llvm::seq<int64_t>(0, inType.getRank()));
208+
broadcastDim.erase(broadcastDim.begin() + axis);
209+
Value broadcast = rewriter.create<mhlo::DynamicBroadcastInDimOp>(
210+
op->getLoc(), inType, sqrt,
211+
rewriter.create<shape::ShapeOfOp>(op.getLoc(), operand),
212+
rewriter.getI64TensorAttr(broadcastDim));
213+
214+
Value result =
215+
rewriter.create<mhlo::DivOp>(op->getLoc(), operand, broadcast);
216+
rewriter.replaceOp(op, result);
217+
return success();
218+
}
219+
};
220+
119221
struct DecomposeByteIRArgMaxMin : public OpRewritePattern<mhlo::CustomCallOp> {
120222
DecomposeByteIRArgMaxMin(MLIRContext *context, llvm::StringRef customCallName)
121223
: OpRewritePattern<mhlo::CustomCallOp>(context),
@@ -340,6 +442,9 @@ struct DecomposeMhloCustomCallOpsPass
340442
if (!legalOpsSet.contains(getSoftmaxName())) {
341443
patterns.add<DecomposeByteIRSoftmax>(context);
342444
}
445+
if (!legalOpsSet.contains(getL2NormName())) {
446+
patterns.add<DecomposeByteIRL2Norm>(context);
447+
}
343448
if (!legalOpsSet.contains(getArgMaxName())) {
344449
patterns.add<DecomposeByteIRArgMaxMin>(context, getArgMaxName());
345450
}
Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,4 @@
1+
func.func @byteir.l2_norm(%arg0: tensor<10x128xf32>) -> tensor<10x128xf32> {
2+
%0 = stablehlo.custom_call @byteir.l2_norm(%arg0) {byteir_attrs = {axis = [1], eps_outside_sqrt = true, epsilon = 9.9999999999999998E-13 : f64}} : (tensor<10x128xf32>) -> tensor<10x128xf32>
3+
return %0 : tensor<10x128xf32>
4+
}

0 commit comments

Comments
 (0)