brintos

brintos / llvm-project-archived public Read only

0
0
Text · 7.3 KiB · 12adfe1 Raw
171 lines · cpp
1//=== ComplexToROCDLLibraryCalls.cpp - convert from Complex to ROCDL calls ===//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/Conversion/ComplexToROCDLLibraryCalls/ComplexToROCDLLibraryCalls.h"10#include "mlir/Dialect/Arith/IR/Arith.h"11#include "mlir/Dialect/Complex/IR/Complex.h"12#include "mlir/Dialect/Func/IR/FuncOps.h"13#include "mlir/IR/PatternMatch.h"14#include "mlir/IR/TypeUtilities.h"15#include "mlir/Transforms/DialectConversion.h"16 17namespace mlir {18#define GEN_PASS_DEF_CONVERTCOMPLEXTOROCDLLIBRARYCALLS19#include "mlir/Conversion/Passes.h.inc"20} // namespace mlir21 22using namespace mlir;23 24namespace {25 26template <typename Op, typename FloatTy>27// Pattern to convert Complex ops to ROCDL function calls.28struct ComplexOpToROCDLLibraryCalls : public OpRewritePattern<Op> {29  using OpRewritePattern<Op>::OpRewritePattern;30  ComplexOpToROCDLLibraryCalls(MLIRContext *context, StringRef funcName,31                               PatternBenefit benefit = 1)32      : OpRewritePattern<Op>(context, benefit), funcName(funcName) {}33 34  LogicalResult matchAndRewrite(Op op, PatternRewriter &rewriter) const final {35    Operation *symTable = SymbolTable::getNearestSymbolTable(op);36    Type resType = op.getType();37    if (auto complexType = dyn_cast<ComplexType>(resType))38      resType = complexType.getElementType();39    if (!isa<FloatTy>(resType))40      return failure();41 42    auto opFunc = dyn_cast_or_null<SymbolOpInterface>(43        SymbolTable::lookupSymbolIn(symTable, funcName));44    if (!opFunc) {45      OpBuilder::InsertionGuard guard(rewriter);46      rewriter.setInsertionPointToStart(&symTable->getRegion(0).front());47      auto funcTy = FunctionType::get(48          rewriter.getContext(), op->getOperandTypes(), op->getResultTypes());49      opFunc = func::FuncOp::create(rewriter, rewriter.getUnknownLoc(),50                                    funcName, funcTy);51      opFunc.setPrivate();52    }53    rewriter.replaceOpWithNewOp<func::CallOp>(op, funcName, op.getType(),54                                              op->getOperands());55    return success();56  }57 58private:59  std::string funcName;60};61 62// Rewrite complex.pow(z, w) -> complex.exp(w * complex.log(z))63struct PowOpToROCDLLibraryCalls : public OpRewritePattern<complex::PowOp> {64  using OpRewritePattern<complex::PowOp>::OpRewritePattern;65 66  LogicalResult matchAndRewrite(complex::PowOp op,67                                PatternRewriter &rewriter) const final {68    Location loc = op.getLoc();69    auto fastmath = op.getFastmathAttr();70    Value logBase =71        complex::LogOp::create(rewriter, loc, op.getLhs(), fastmath);72    Value mul =73        complex::MulOp::create(rewriter, loc, op.getRhs(), logBase, fastmath);74    Value exp = complex::ExpOp::create(rewriter, loc, mul, fastmath);75    rewriter.replaceOp(op, exp);76    return success();77  }78};79 80// Rewrite complex.powi(z, n) -> complex.pow(z, complex(float(n), 0))81struct PowiOpToROCDLLibraryCalls : public OpRewritePattern<complex::PowiOp> {82  using OpRewritePattern<complex::PowiOp>::OpRewritePattern;83 84  LogicalResult matchAndRewrite(complex::PowiOp op,85                                PatternRewriter &rewriter) const final {86    auto complexType = cast<ComplexType>(getElementTypeOrSelf(op.getType()));87    Type elementType = complexType.getElementType();88 89    Type exponentType = op.getRhs().getType();90    Type exponentFloatType = elementType;91    if (auto shapedType = dyn_cast<ShapedType>(exponentType))92      exponentFloatType = shapedType.cloneWith(std::nullopt, elementType);93 94    Location loc = op.getLoc();95    Value exponentReal =96        arith::SIToFPOp::create(rewriter, loc, exponentFloatType, op.getRhs());97    Value zeroImag = arith::ConstantOp::create(98        rewriter, loc, rewriter.getZeroAttr(exponentFloatType));99    Value exponent = complex::CreateOp::create(100        rewriter, loc, op.getLhs().getType(), exponentReal, zeroImag);101 102    rewriter.replaceOpWithNewOp<complex::PowOp>(op, op.getType(), op.getLhs(),103                                                exponent, op.getFastmathAttr());104    return success();105  }106};107} // namespace108 109void mlir::populateComplexToROCDLLibraryCallsConversionPatterns(110    RewritePatternSet &patterns) {111  patterns.add<PowiOpToROCDLLibraryCalls>(patterns.getContext());112  patterns.add<PowOpToROCDLLibraryCalls>(patterns.getContext());113  patterns.add<ComplexOpToROCDLLibraryCalls<complex::AbsOp, Float32Type>>(114      patterns.getContext(), "__ocml_cabs_f32");115  patterns.add<ComplexOpToROCDLLibraryCalls<complex::AbsOp, Float64Type>>(116      patterns.getContext(), "__ocml_cabs_f64");117  patterns.add<ComplexOpToROCDLLibraryCalls<complex::CosOp, Float32Type>>(118      patterns.getContext(), "__ocml_ccos_f32");119  patterns.add<ComplexOpToROCDLLibraryCalls<complex::CosOp, Float64Type>>(120      patterns.getContext(), "__ocml_ccos_f64");121  patterns.add<ComplexOpToROCDLLibraryCalls<complex::ExpOp, Float32Type>>(122      patterns.getContext(), "__ocml_cexp_f32");123  patterns.add<ComplexOpToROCDLLibraryCalls<complex::ExpOp, Float64Type>>(124      patterns.getContext(), "__ocml_cexp_f64");125  patterns.add<ComplexOpToROCDLLibraryCalls<complex::LogOp, Float32Type>>(126      patterns.getContext(), "__ocml_clog_f32");127  patterns.add<ComplexOpToROCDLLibraryCalls<complex::LogOp, Float64Type>>(128      patterns.getContext(), "__ocml_clog_f64");129  patterns.add<ComplexOpToROCDLLibraryCalls<complex::SinOp, Float32Type>>(130      patterns.getContext(), "__ocml_csin_f32");131  patterns.add<ComplexOpToROCDLLibraryCalls<complex::SinOp, Float64Type>>(132      patterns.getContext(), "__ocml_csin_f64");133  patterns.add<ComplexOpToROCDLLibraryCalls<complex::SqrtOp, Float32Type>>(134      patterns.getContext(), "__ocml_csqrt_f32");135  patterns.add<ComplexOpToROCDLLibraryCalls<complex::SqrtOp, Float64Type>>(136      patterns.getContext(), "__ocml_csqrt_f64");137  patterns.add<ComplexOpToROCDLLibraryCalls<complex::TanOp, Float32Type>>(138      patterns.getContext(), "__ocml_ctan_f32");139  patterns.add<ComplexOpToROCDLLibraryCalls<complex::TanOp, Float64Type>>(140      patterns.getContext(), "__ocml_ctan_f64");141  patterns.add<ComplexOpToROCDLLibraryCalls<complex::TanhOp, Float32Type>>(142      patterns.getContext(), "__ocml_ctanh_f32");143  patterns.add<ComplexOpToROCDLLibraryCalls<complex::TanhOp, Float64Type>>(144      patterns.getContext(), "__ocml_ctanh_f64");145}146 147namespace {148struct ConvertComplexToROCDLLibraryCallsPass149    : public impl::ConvertComplexToROCDLLibraryCallsBase<150          ConvertComplexToROCDLLibraryCallsPass> {151  void runOnOperation() override;152};153} // namespace154 155void ConvertComplexToROCDLLibraryCallsPass::runOnOperation() {156  Operation *op = getOperation();157 158  RewritePatternSet patterns(&getContext());159  populateComplexToROCDLLibraryCallsConversionPatterns(patterns);160 161  ConversionTarget target(getContext());162  target.addLegalDialect<arith::ArithDialect, func::FuncDialect>();163  target.addLegalOp<complex::CreateOp, complex::MulOp>();164  target.addIllegalOp<complex::AbsOp, complex::CosOp, complex::ExpOp,165                      complex::LogOp, complex::PowOp, complex::PowiOp,166                      complex::SinOp, complex::SqrtOp, complex::TanOp,167                      complex::TanhOp>();168  if (failed(applyPartialConversion(op, target, std::move(patterns))))169    signalPassFailure();170}171