================ @@ -0,0 +1,494 @@ +//===- ComplexLowering.cpp - Expand complex multiply and divide -----------===// +// +// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions. +// See https://llvm.org/LICENSE.txt for license information. +// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception +// +//===----------------------------------------------------------------------===// +// +// This file implements a pass that replaces cir.complex.mul and +// cir.complex.div with the arithmetic each one expands to, which for the full +// complex range is a call to a runtime helper such as __mulsc3 or __divsc3. +// +//===----------------------------------------------------------------------===// + +#include "PassDetail.h" +#include "mlir/IR/Location.h" +#include "mlir/IR/Value.h" +#include "clang/AST/ASTContext.h" +#include "clang/Basic/LangOptions.h" +#include "clang/Basic/TargetInfo.h" +#include "clang/CIR/Dialect/Builder/CIRBaseBuilder.h" +#include "clang/CIR/Dialect/IR/CIRDialect.h" +#include "clang/CIR/Dialect/IR/CIROpsEnums.h" +#include "clang/CIR/Dialect/IR/CIRTypes.h" +#include "clang/CIR/Dialect/Passes.h" +#include "clang/CIR/Dialect/Transforms/CIRTransformUtils.h" +#include "clang/CIR/MissingFeatures.h" +#include "llvm/ADT/APFloat.h" +#include "llvm/ADT/StringRef.h" +#include "llvm/Support/ErrorHandling.h" + +#include <memory> + +using namespace mlir; +using namespace cir; + +namespace mlir { +#define GEN_PASS_DEF_COMPLEXLOWERING +#include "clang/CIR/Dialect/Passes.h.inc" +} // namespace mlir + +namespace { +struct ComplexLoweringPass + : public impl::ComplexLoweringBase<ComplexLoweringPass> { + ComplexLoweringPass() = default; + + void runOnOperation() override; + + void lowerComplexDivOp(cir::ComplexDivOp op); + void lowerComplexMulOp(cir::ComplexMulOp op); + + void setASTContext(clang::ASTContext *c) { astCtx = c; } + + /// Read by the promoted-range division path, which asks the target for the + /// semantics of a higher-precision element type. + clang::ASTContext *astCtx = nullptr; + + mlir::ModuleOp mlirModule; +}; + +} // namespace + +static mlir::Value buildComplexBinOpLibCall( + mlir::ModuleOp mlirModule, CIRBaseBuilderTy &builder, + llvm::StringRef (*libFuncNameGetter)(llvm::APFloat::Semantics), + mlir::Location loc, cir::ComplexType ty, mlir::Value lhsReal, + mlir::Value lhsImag, mlir::Value rhsReal, mlir::Value rhsImag) { + cir::FPTypeInterface elementTy = + mlir::cast<cir::FPTypeInterface>(ty.getElementType()); + + llvm::StringRef libFuncName = libFuncNameGetter( + llvm::APFloat::SemanticsToEnum(elementTy.getFloatSemantics())); + llvm::SmallVector<mlir::Type, 4> libFuncInputTypes(4, elementTy); + + cir::FuncType libFuncTy = cir::FuncType::get(libFuncInputTypes, ty); + + // Insert a declaration for the runtime function to be used in Complex + // multiplication and division when needed + cir::FuncOp libFunc; + { + mlir::OpBuilder::InsertionGuard ipGuard{builder}; + builder.setInsertionPointToStart(mlirModule.getBody()); + libFunc = cir::buildRuntimeFunction(builder, mlirModule, libFuncName, loc, + libFuncTy); + } + + cir::CallOp call = + builder.createCallOp(loc, libFunc, {lhsReal, lhsImag, rhsReal, rhsImag}); + return call.getResult(); +} + +static llvm::StringRef +getComplexDivLibCallName(llvm::APFloat::Semantics semantics) { + switch (semantics) { + case llvm::APFloat::S_IEEEhalf: + return "__divhc3"; + case llvm::APFloat::S_IEEEsingle: + return "__divsc3"; + case llvm::APFloat::S_IEEEdouble: + return "__divdc3"; + case llvm::APFloat::S_PPCDoubleDouble: + return "__divtc3"; + case llvm::APFloat::S_x87DoubleExtended: + return "__divxc3"; + case llvm::APFloat::S_IEEEquad: + return "__divtc3"; + default: + llvm_unreachable("unsupported floating point type"); + } +} + +static mlir::Value +buildAlgebraicComplexDiv(CIRBaseBuilderTy &builder, mlir::Location loc, + mlir::Value lhsReal, mlir::Value lhsImag, + mlir::Value rhsReal, mlir::Value rhsImag) { + // (a+bi) / (c+di) = ((ac+bd)/(cc+dd)) + ((bc-ad)/(cc+dd))i + mlir::Value &a = lhsReal; + mlir::Value &b = lhsImag; + mlir::Value &c = rhsReal; + mlir::Value &d = rhsImag; + + // The element type of the complex (lhs/rhs) determines whether floating + // point or integer ops are needed. + bool isFP = cir::isFPOrVectorOfFPType(a.getType()); + auto mul = [&](mlir::Location l, mlir::Value x, mlir::Value y) { + return isFP ? builder.createFMul(l, x, y) : builder.createMul(l, x, y); + }; + auto add = [&](mlir::Location l, mlir::Value x, mlir::Value y) { + return isFP ? builder.createFAdd(l, x, y) : builder.createAdd(l, x, y); + }; + auto sub = [&](mlir::Location l, mlir::Value x, mlir::Value y) { + return isFP ? builder.createFSub(l, x, y) : builder.createSub(l, x, y); + }; + auto div = [&](mlir::Location l, mlir::Value x, mlir::Value y) { + return isFP ? builder.createFDiv(l, x, y) : builder.createDiv(l, x, y); + }; + + mlir::Value ac = mul(loc, a, c); // a*c + mlir::Value bd = mul(loc, b, d); // b*d + mlir::Value cc = mul(loc, c, c); // c*c + mlir::Value dd = mul(loc, d, d); // d*d + mlir::Value acbd = add(loc, ac, bd); // ac+bd + mlir::Value ccdd = add(loc, cc, dd); // cc+dd + mlir::Value resultReal = div(loc, acbd, ccdd); + + mlir::Value bc = mul(loc, b, c); // b*c + mlir::Value ad = mul(loc, a, d); // a*d + mlir::Value bcad = sub(loc, bc, ad); // bc-ad + mlir::Value resultImag = div(loc, bcad, ccdd); + return builder.createComplexCreate(loc, resultReal, resultImag); +} + +static mlir::Value +buildRangeReductionComplexDiv(CIRBaseBuilderTy &builder, mlir::Location loc, + mlir::Value lhsReal, mlir::Value lhsImag, + mlir::Value rhsReal, mlir::Value rhsImag) { + // Implements Smith's algorithm for complex division. + // SMITH, R. L. Algorithm 116: Complex division. Commun. ACM 5, 8 (1962). + + // Let: + // - lhs := a+bi + // - rhs := c+di + // - result := lhs / rhs = e+fi + // + // The algorithm pseudocode looks like follows: + // if fabs(c) >= fabs(d): + // r := d / c + // tmp := c + r*d + // e = (a + b*r) / tmp + // f = (b - a*r) / tmp + // else: + // r := c / d + // tmp := d + r*c + // e = (a*r + b) / tmp + // f = (b*r - a) / tmp + + mlir::Value &a = lhsReal; + mlir::Value &b = lhsImag; + mlir::Value &c = rhsReal; + mlir::Value &d = rhsImag; + + // Smith's algorithm is only used for floating-point complex division. + assert(cir::isFPOrVectorOfFPType(a.getType()) && + "range-reduction complex divide expects floating-point operands"); + + auto trueBranchBuilder = [&](mlir::OpBuilder &, mlir::Location) { + mlir::Value r = builder.createFDiv(loc, d, c); // r := d / c + mlir::Value rd = builder.createFMul(loc, r, d); // r*d + mlir::Value tmp = builder.createFAdd(loc, c, rd); // tmp := c + r*d + + mlir::Value br = builder.createFMul(loc, b, r); // b*r + mlir::Value abr = builder.createFAdd(loc, a, br); // a + b*r + mlir::Value e = builder.createFDiv(loc, abr, tmp); + + mlir::Value ar = builder.createFMul(loc, a, r); // a*r + mlir::Value bar = builder.createFSub(loc, b, ar); // b - a*r + mlir::Value f = builder.createFDiv(loc, bar, tmp); + + mlir::Value result = builder.createComplexCreate(loc, e, f); + builder.createYield(loc, result); + }; + + auto falseBranchBuilder = [&](mlir::OpBuilder &, mlir::Location) { + mlir::Value r = builder.createFDiv(loc, c, d); // r := c / d + mlir::Value rc = builder.createFMul(loc, r, c); // r*c + mlir::Value tmp = builder.createFAdd(loc, d, rc); // tmp := d + r*c + + mlir::Value ar = builder.createFMul(loc, a, r); // a*r + mlir::Value arb = builder.createFAdd(loc, ar, b); // a*r + b + mlir::Value e = builder.createFDiv(loc, arb, tmp); + + mlir::Value br = builder.createFMul(loc, b, r); // b*r + mlir::Value bra = builder.createFSub(loc, br, a); // b*r - a + mlir::Value f = builder.createFDiv(loc, bra, tmp); + + mlir::Value result = builder.createComplexCreate(loc, e, f); + builder.createYield(loc, result); + }; + + auto cFabs = cir::FAbsOp::create(builder, loc, c); + auto dFabs = cir::FAbsOp::create(builder, loc, d); + cir::CmpOp cmpResult = + builder.createCompare(loc, cir::CmpOpKind::ge, cFabs, dFabs); + auto ternary = cir::TernaryOp::create(builder, loc, cmpResult, + trueBranchBuilder, falseBranchBuilder); + + return ternary.getResult(); +} + +static mlir::Type higherPrecisionElementTypeForComplexArithmetic( + mlir::MLIRContext &context, clang::ASTContext &cc, + CIRBaseBuilderTy &builder, mlir::Type elementType) { + + auto getHigherPrecisionFPType = [&context](mlir::Type type) -> mlir::Type { + if (mlir::isa<cir::FP16Type>(type)) + return cir::SingleType::get(&context); + + if (mlir::isa<cir::SingleType>(type) || mlir::isa<cir::BF16Type>(type)) + return cir::DoubleType::get(&context); ---------------- andykaylor wrote:
This isn't reliable. Not all targets support double-precision floating-point values. https://github.com/llvm/llvm-project/pull/216498 _______________________________________________ cfe-commits mailing list [email protected] https://lists.llvm.org/cgi-bin/mailman/listinfo/cfe-commits
