@@ -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+
119221struct 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 }
0 commit comments