[clang] [CIR] Lower complex mul and div before callconv lowering (PR #216498)
Andy Kaylor via cfe-commits
cfe-commits at lists.llvm.org
Mon Aug 17 12:47:56 PDT 2026
================
@@ -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
More information about the cfe-commits
mailing list