83 lines · cpp
1//===- ShapeToShapeLowering.cpp - Prepare for lowering to Standard --------===//2//3// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.4// See https://llvm.org/LICENSE.txt for license information.5// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception6//7//===----------------------------------------------------------------------===//8 9#include "mlir/Dialect/Shape/Transforms/Passes.h"10 11#include "mlir/Dialect/Arith/IR/Arith.h"12#include "mlir/Dialect/Func/IR/FuncOps.h"13#include "mlir/Dialect/Shape/IR/Shape.h"14#include "mlir/IR/Builders.h"15#include "mlir/IR/PatternMatch.h"16#include "mlir/Transforms/DialectConversion.h"17 18namespace mlir {19#define GEN_PASS_DEF_SHAPETOSHAPELOWERINGPASS20#include "mlir/Dialect/Shape/Transforms/Passes.h.inc"21} // namespace mlir22 23using namespace mlir;24using namespace mlir::shape;25 26namespace {27/// Converts `shape.num_elements` to `shape.reduce`.28struct NumElementsOpConverter : public OpRewritePattern<NumElementsOp> {29public:30 using OpRewritePattern::OpRewritePattern;31 32 LogicalResult matchAndRewrite(NumElementsOp op,33 PatternRewriter &rewriter) const final;34};35} // namespace36 37LogicalResult38NumElementsOpConverter::matchAndRewrite(NumElementsOp op,39 PatternRewriter &rewriter) const {40 auto loc = op.getLoc();41 Type valueType = op.getResult().getType();42 Value init = op->getDialect()43 ->materializeConstant(rewriter, rewriter.getIndexAttr(1),44 valueType, loc)45 ->getResult(0);46 ReduceOp reduce = ReduceOp::create(rewriter, loc, op.getShape(), init);47 48 // Generate reduce operator.49 Block *body = reduce.getBody();50 OpBuilder b = OpBuilder::atBlockEnd(body);51 Value product = MulOp::create(b, loc, valueType, body->getArgument(1),52 body->getArgument(2));53 shape::YieldOp::create(b, loc, product);54 55 rewriter.replaceOp(op, reduce.getResult());56 return success();57}58 59namespace {60struct ShapeToShapeLowering61 : public impl::ShapeToShapeLoweringPassBase<ShapeToShapeLowering> {62 void runOnOperation() override;63};64} // namespace65 66void ShapeToShapeLowering::runOnOperation() {67 MLIRContext &ctx = getContext();68 69 RewritePatternSet patterns(&ctx);70 populateShapeRewritePatterns(patterns);71 72 ConversionTarget target(getContext());73 target.addLegalDialect<arith::ArithDialect, ShapeDialect>();74 target.addIllegalOp<NumElementsOp>();75 if (failed(mlir::applyPartialConversion(getOperation(), target,76 std::move(patterns))))77 signalPassFailure();78}79 80void mlir::populateShapeRewritePatterns(RewritePatternSet &patterns) {81 patterns.add<NumElementsOpConverter>(patterns.getContext());82}83