brintos

brintos / llvm-project-archived public Read only

0
0
Text · 4.7 KiB · c904556 Raw
96 lines · cpp
1//===- NamedToElementwise.cpp - convert linalg named op into elementwise --===//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// This file implements rewriting those linalg named ops that are essentially10// elementwise e.g. `linalg.exp`, to `linalg.elementwise`. This allows further11// optimization on `linalg.elementwise` such as folding transpose, broadcast.12//13//===----------------------------------------------------------------------===//14 15#include "mlir/Dialect/Linalg/IR/Linalg.h"16#include "mlir/Dialect/Linalg/Passes.h"17#include "mlir/Dialect/Linalg/Transforms/Transforms.h"18#include "mlir/IR/PatternMatch.h"19#include "mlir/Transforms/GreedyPatternRewriteDriver.h"20#include "llvm/ADT/SmallVector.h"21#include "llvm/ADT/TypeSwitch.h"22 23using namespace mlir;24using namespace mlir::linalg;25 26#define DEBUG_TYPE "linalg-named-to-elementwise"27 28namespace {29ElementwiseKind getKind(Operation *op) {30  return llvm::TypeSwitch<Operation *, ElementwiseKind>(op)31      .Case([](SelectOp) { return ElementwiseKind::select; })32      .Case([](AddOp) { return ElementwiseKind::add; })33      .Case([](SubOp) { return ElementwiseKind::sub; })34      .Case([](MulOp) { return ElementwiseKind::mul; })35      .Case([](DivOp) { return ElementwiseKind::div; })36      .Case([](DivUnsignedOp) { return ElementwiseKind::div_unsigned; })37      .Case([](PowFOp) { return ElementwiseKind::powf; })38      .Case([](ExpOp) { return ElementwiseKind::exp; })39      .Case([](LogOp) { return ElementwiseKind::log; })40      .Case([](AbsOp) { return ElementwiseKind::abs; })41      .Case([](CeilOp) { return ElementwiseKind::ceil; })42      .Case([](FloorOp) { return ElementwiseKind::floor; })43      .Case([](NegFOp) { return ElementwiseKind::negf; })44      .Case([](ReciprocalOp) { return ElementwiseKind::reciprocal; })45      .Case([](RoundOp) { return ElementwiseKind::round; })46      .Case([](SqrtOp) { return ElementwiseKind::sqrt; })47      .Case([](RsqrtOp) { return ElementwiseKind::rsqrt; })48      .Case([](SquareOp) { return ElementwiseKind::square; })49      .Case([](TanhOp) { return ElementwiseKind::tanh; })50      .Case([](ErfOp) { return ElementwiseKind::erf; })51      .DefaultUnreachable("unhandled case in named to elementwise");52}53 54template <typename NamedOpTy>55struct NamedToElementwisePattern : public OpRewritePattern<NamedOpTy> {56  using OpRewritePattern<NamedOpTy>::OpRewritePattern;57 58  LogicalResult matchAndRewrite(NamedOpTy op,59                                PatternRewriter &rewriter) const override {60    SmallVector<NamedAttribute> attrs;61    auto kindAttr = ElementwiseKindAttr::get(op.getContext(), getKind(op));62    attrs.push_back(rewriter.getNamedAttr("kind", kindAttr));63    attrs.push_back(64        rewriter.getNamedAttr("indexing_maps", op.getIndexingMaps()));65 66    rewriter.replaceOpWithNewOp<ElementwiseOp>(op, op.getDpsInputs(),67                                               op.getDpsInits(), attrs);68    return success();69  }70};71} // namespace72 73void mlir::linalg::populateLinalgNamedToElementwisePatterns(74    RewritePatternSet &patterns) {75  patterns.add<NamedToElementwisePattern<SelectOp>>(patterns.getContext());76  patterns.add<NamedToElementwisePattern<AddOp>>(patterns.getContext());77  patterns.add<NamedToElementwisePattern<SubOp>>(patterns.getContext());78  patterns.add<NamedToElementwisePattern<MulOp>>(patterns.getContext());79  patterns.add<NamedToElementwisePattern<DivOp>>(patterns.getContext());80  patterns.add<NamedToElementwisePattern<DivUnsignedOp>>(patterns.getContext());81  patterns.add<NamedToElementwisePattern<PowFOp>>(patterns.getContext());82  patterns.add<NamedToElementwisePattern<ExpOp>>(patterns.getContext());83  patterns.add<NamedToElementwisePattern<LogOp>>(patterns.getContext());84  patterns.add<NamedToElementwisePattern<AbsOp>>(patterns.getContext());85  patterns.add<NamedToElementwisePattern<CeilOp>>(patterns.getContext());86  patterns.add<NamedToElementwisePattern<FloorOp>>(patterns.getContext());87  patterns.add<NamedToElementwisePattern<NegFOp>>(patterns.getContext());88  patterns.add<NamedToElementwisePattern<ReciprocalOp>>(patterns.getContext());89  patterns.add<NamedToElementwisePattern<RoundOp>>(patterns.getContext());90  patterns.add<NamedToElementwisePattern<SqrtOp>>(patterns.getContext());91  patterns.add<NamedToElementwisePattern<RsqrtOp>>(patterns.getContext());92  patterns.add<NamedToElementwisePattern<SquareOp>>(patterns.getContext());93  patterns.add<NamedToElementwisePattern<TanhOp>>(patterns.getContext());94  patterns.add<NamedToElementwisePattern<ErfOp>>(patterns.getContext());95}96