brintos

brintos / llvm-project-archived public Read only

0
0
Text · 2.7 KiB · 3c363f3 Raw
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