[clang] [CIR] Lower complex mul and div before callconv lowering (PR #216498)
Adam Smith via cfe-commits
cfe-commits at lists.llvm.org
Mon Aug 17 11:49:08 PDT 2026
https://github.com/adams381 updated https://github.com/llvm/llvm-project/pull/216498
>From a725f0b935aef8482472662bf714e42be71379c9 Mon Sep 17 00:00:00 2001
From: Adam Smith <adams at nvidia.com>
Date: Sat, 15 Aug 2026 09:57:19 -0700
Subject: [PATCH 1/3] [CIR] Lower complex mul and div before callconv lowering
Complex division returns a wrong imaginary part. CIR declares `__divsc3` as
returning `{ float, float }` where classic CodeGen coerces the return to
`<2 x float>`, so the caller reads the two halves out of two registers while the
callee packs both into one. Dividing 3+4i by 1+2i gives 2.2 and 4.0 instead of
2.2 and -0.4, with no diagnostic. The helper call is synthesized by
LoweringPrepare, which runs after the calling-convention pass, so it is never
classified.
Expanding `cir.complex.mul` and `cir.complex.div` in their own pass, scheduled
before CallConvLowering, is enough to fix it and leaves the rest of the
lowering-prepare work where it is. Only the float helpers change shape, since
double and long double already returned a record that matched classic.
Assisted-by: Cursor / claude-opus-5
---
clang/include/clang/CIR/Dialect/Passes.h | 2 +
clang/include/clang/CIR/Dialect/Passes.td | 19 ++++
.../Dialect/Transforms/LoweringPrepare.cpp | 106 +++++++++++++-----
clang/lib/CIR/Lowering/CIRPasses.cpp | 6 +
.../CodeGen/complex-compound-assignment.cpp | 36 ++++--
.../complex-libcall-abi-global-init.cpp | 30 +++++
clang/test/CIR/CodeGen/complex-libcall-abi.c | 76 +++++++++++++
clang/test/CIR/CodeGen/complex-mul-div.cpp | 29 +++--
8 files changed, 262 insertions(+), 42 deletions(-)
create mode 100644 clang/test/CIR/CodeGen/complex-libcall-abi-global-init.cpp
create mode 100644 clang/test/CIR/CodeGen/complex-libcall-abi.c
diff --git a/clang/include/clang/CIR/Dialect/Passes.h b/clang/include/clang/CIR/Dialect/Passes.h
index 674f2180c27ab..c668607dced34 100644
--- a/clang/include/clang/CIR/Dialect/Passes.h
+++ b/clang/include/clang/CIR/Dialect/Passes.h
@@ -42,6 +42,8 @@ std::unique_ptr<Pass> createCallConvLoweringPass(
std::unique_ptr<Pass> createHoistAllocasPass();
std::unique_ptr<Pass> createLoweringPreparePass();
std::unique_ptr<Pass> createLoweringPreparePass(clang::ASTContext *astCtx);
+std::unique_ptr<Pass> createComplexLoweringPass();
+std::unique_ptr<Pass> createComplexLoweringPass(clang::ASTContext *astCtx);
std::unique_ptr<Pass> createGotoSolverPass();
std::unique_ptr<Pass> createIdiomRecognizerPass();
std::unique_ptr<Pass> createLibOptPass();
diff --git a/clang/include/clang/CIR/Dialect/Passes.td b/clang/include/clang/CIR/Dialect/Passes.td
index b983cfe59112a..c0a79d6fde66c 100644
--- a/clang/include/clang/CIR/Dialect/Passes.td
+++ b/clang/include/clang/CIR/Dialect/Passes.td
@@ -184,6 +184,25 @@ def LoweringPrepare : Pass<"cir-lowering-prepare"> {
let dependentDialects = ["cir::CIRDialect"];
}
+def ComplexLowering : Pass<"cir-complex-lowering", "mlir::ModuleOp"> {
+ let summary = "Expand complex multiplication and division";
+ let description = [{
+ This pass 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`.
+
+ It runs before the calling-convention pass rather than alongside the rest
+ of the lowering-prepare work, because a call created after that pass has
+ run never gets its return type coerced, which places the real and
+ imaginary halves of the result in the wrong registers.
+
+ Only these two operations need the earlier position, so `cir.complex.conj`
+ and the complex casts are still expanded by `cir-lowering-prepare`.
+ }];
+ let constructor = "mlir::createComplexLoweringPass()";
+ let dependentDialects = ["cir::CIRDialect"];
+}
+
def IdiomRecognizer : Pass<"cir-idiom-recognizer", "mlir::ModuleOp"> {
let summary = "Raise calls to C/C++ libraries to CIR operations";
let description = [{
diff --git a/clang/lib/CIR/Dialect/Transforms/LoweringPrepare.cpp b/clang/lib/CIR/Dialect/Transforms/LoweringPrepare.cpp
index 28ba14962cc7b..a71aba6c1460f 100644
--- a/clang/lib/CIR/Dialect/Transforms/LoweringPrepare.cpp
+++ b/clang/lib/CIR/Dialect/Transforms/LoweringPrepare.cpp
@@ -44,6 +44,7 @@ using namespace mlir;
using namespace cir;
namespace mlir {
+#define GEN_PASS_DEF_COMPLEXLOWERING
#define GEN_PASS_DEF_LOWERINGPREPARE
#include "clang/CIR/Dialect/Passes.h.inc"
} // namespace mlir
@@ -88,8 +89,6 @@ struct LoweringPreparePass
void runOnOp(mlir::Operation *op);
void lowerCastOp(cir::CastOp op);
void lowerComplexConjOp(cir::ComplexConjOp op);
- void lowerComplexDivOp(cir::ComplexDivOp op);
- void lowerComplexMulOp(cir::ComplexMulOp op);
void lowerGetGlobalOp(cir::GetGlobalOp op);
void lowerGlobalOp(cir::GlobalOp op);
void lowerThreeWayCmpOp(cir::CmpThreeWayOp op);
@@ -505,6 +504,26 @@ struct LoweringPreparePass
void setASTContext(clang::ASTContext *c) { astCtx = c; }
};
+/// Expand `cir.complex.mul` and `cir.complex.div`. See the pass description
+/// in Passes.td for why this cannot run with the rest of LoweringPrepare.
+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
cir::GlobalOp LoweringPreparePass::getOrCreateRuntimeVariable(
@@ -525,9 +544,13 @@ cir::GlobalOp LoweringPreparePass::getOrCreateRuntimeVariable(
return g;
}
-cir::FuncOp LoweringPreparePass::buildRuntimeFunction(
- mlir::OpBuilder &builder, llvm::StringRef name, mlir::Location loc,
- cir::FuncType type, cir::GlobalLinkageKind linkage) {
+/// Declare `name` in `mlirModule` if it is not already declared there, and
+/// return the declaration. Free-standing so that ComplexLoweringPass can
+/// reach it without a LoweringPreparePass instance.
+static cir::FuncOp buildRuntimeFunction(
+ mlir::OpBuilder &builder, mlir::ModuleOp mlirModule, llvm::StringRef name,
+ mlir::Location loc, cir::FuncType type,
+ cir::GlobalLinkageKind linkage = cir::GlobalLinkageKind::ExternalLinkage) {
cir::FuncOp f = dyn_cast_or_null<FuncOp>(SymbolTable::lookupNearestSymbolFrom(
mlirModule, StringAttr::get(mlirModule->getContext(), name)));
if (!f) {
@@ -542,6 +565,12 @@ cir::FuncOp LoweringPreparePass::buildRuntimeFunction(
return f;
}
+cir::FuncOp LoweringPreparePass::buildRuntimeFunction(
+ mlir::OpBuilder &builder, llvm::StringRef name, mlir::Location loc,
+ cir::FuncType type, cir::GlobalLinkageKind linkage) {
+ return ::buildRuntimeFunction(builder, mlirModule, name, loc, type, linkage);
+}
+
static mlir::Value lowerScalarToComplexCast(mlir::MLIRContext &ctx,
cir::CastOp op) {
cir::CIRBaseBuilderTy builder(ctx);
@@ -628,7 +657,7 @@ void LoweringPreparePass::lowerCastOp(cir::CastOp op) {
}
static mlir::Value buildComplexBinOpLibCall(
- LoweringPreparePass &pass, CIRBaseBuilderTy &builder,
+ 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) {
@@ -646,8 +675,9 @@ static mlir::Value buildComplexBinOpLibCall(
cir::FuncOp libFunc;
{
mlir::OpBuilder::InsertionGuard ipGuard{builder};
- builder.setInsertionPointToStart(pass.mlirModule.getBody());
- libFunc = pass.buildRuntimeFunction(builder, libFuncName, loc, libFuncTy);
+ builder.setInsertionPointToStart(mlirModule.getBody());
+ libFunc =
+ buildRuntimeFunction(builder, mlirModule, libFuncName, loc, libFuncTy);
}
cir::CallOp call =
@@ -864,7 +894,7 @@ static mlir::Type higherPrecisionElementTypeForComplexArithmetic(
}
static mlir::Value
-lowerComplexDiv(LoweringPreparePass &pass, CIRBaseBuilderTy &builder,
+lowerComplexDiv(mlir::ModuleOp mlirModule, CIRBaseBuilderTy &builder,
mlir::Location loc, cir::ComplexDivOp op, mlir::Value lhsReal,
mlir::Value lhsImag, mlir::Value rhsReal, mlir::Value rhsImag,
mlir::MLIRContext &mlirCx, clang::ASTContext &cc) {
@@ -876,9 +906,9 @@ lowerComplexDiv(LoweringPreparePass &pass, CIRBaseBuilderTy &builder,
rhsReal, rhsImag);
if (range == cir::ComplexRangeKind::Full)
- return buildComplexBinOpLibCall(pass, builder, &getComplexDivLibCallName,
- loc, complexTy, lhsReal, lhsImag, rhsReal,
- rhsImag);
+ return buildComplexBinOpLibCall(mlirModule, builder,
+ &getComplexDivLibCallName, loc, complexTy,
+ lhsReal, lhsImag, rhsReal, rhsImag);
if (range == cir::ComplexRangeKind::Promoted) {
mlir::Type originalElementType = complexTy.getElementType();
@@ -918,7 +948,7 @@ lowerComplexDiv(LoweringPreparePass &pass, CIRBaseBuilderTy &builder,
rhsImag);
}
-void LoweringPreparePass::lowerComplexDivOp(cir::ComplexDivOp op) {
+void ComplexLoweringPass::lowerComplexDivOp(cir::ComplexDivOp op) {
cir::CIRBaseBuilderTy builder(getContext());
builder.setInsertionPointAfter(op);
mlir::Location loc = op.getLoc();
@@ -930,7 +960,7 @@ void LoweringPreparePass::lowerComplexDivOp(cir::ComplexDivOp op) {
mlir::Value rhsImag = builder.createComplexImag(loc, rhs);
mlir::Value loweredResult =
- lowerComplexDiv(*this, builder, loc, op, lhsReal, lhsImag, rhsReal,
+ lowerComplexDiv(mlirModule, builder, loc, op, lhsReal, lhsImag, rhsReal,
rhsImag, getContext(), *astCtx);
op.replaceAllUsesWith(loweredResult);
op.erase();
@@ -956,7 +986,7 @@ getComplexMulLibCallName(llvm::APFloat::Semantics semantics) {
}
}
-static mlir::Value lowerComplexMul(LoweringPreparePass &pass,
+static mlir::Value lowerComplexMul(mlir::ModuleOp mlirModule,
CIRBaseBuilderTy &builder,
mlir::Location loc, cir::ComplexMulOp op,
mlir::Value lhsReal, mlir::Value lhsImag,
@@ -1004,8 +1034,8 @@ static mlir::Value lowerComplexMul(LoweringPreparePass &pass,
builder, loc, resultRealAndImagAreNaN,
[&](mlir::OpBuilder &, mlir::Location) {
mlir::Value libCallResult = buildComplexBinOpLibCall(
- pass, builder, &getComplexMulLibCallName, loc, complexTy,
- lhsReal, lhsImag, rhsReal, rhsImag);
+ mlirModule, builder, &getComplexMulLibCallName, loc,
+ complexTy, lhsReal, lhsImag, rhsReal, rhsImag);
builder.createYield(loc, libCallResult);
},
[&](mlir::OpBuilder &, mlir::Location) {
@@ -1014,7 +1044,7 @@ static mlir::Value lowerComplexMul(LoweringPreparePass &pass,
.getResult();
}
-void LoweringPreparePass::lowerComplexMulOp(cir::ComplexMulOp op) {
+void ComplexLoweringPass::lowerComplexMulOp(cir::ComplexMulOp op) {
cir::CIRBaseBuilderTy builder(getContext());
builder.setInsertionPointAfter(op);
mlir::Location loc = op.getLoc();
@@ -1024,12 +1054,28 @@ void LoweringPreparePass::lowerComplexMulOp(cir::ComplexMulOp op) {
mlir::Value lhsImag = builder.createComplexImag(loc, lhs);
mlir::Value rhsReal = builder.createComplexReal(loc, rhs);
mlir::Value rhsImag = builder.createComplexImag(loc, rhs);
- mlir::Value loweredResult = lowerComplexMul(*this, builder, loc, op, lhsReal,
- lhsImag, rhsReal, rhsImag);
+ mlir::Value loweredResult = lowerComplexMul(
+ mlirModule, builder, loc, op, lhsReal, lhsImag, rhsReal, rhsImag);
op.replaceAllUsesWith(loweredResult);
op.erase();
}
+void ComplexLoweringPass::runOnOperation() {
+ mlirModule = cast<mlir::ModuleOp>(getOperation());
+
+ llvm::SmallVector<mlir::Operation *> opsToTransform;
+ mlirModule->walk([&](mlir::Operation *op) {
+ if (mlir::isa<cir::ComplexMulOp, cir::ComplexDivOp>(op))
+ opsToTransform.push_back(op);
+ });
+
+ for (mlir::Operation *o : opsToTransform)
+ if (auto complexDiv = mlir::dyn_cast<cir::ComplexDivOp>(o))
+ lowerComplexDivOp(complexDiv);
+ else
+ lowerComplexMulOp(mlir::cast<cir::ComplexMulOp>(o));
+}
+
void LoweringPreparePass::lowerComplexConjOp(cir::ComplexConjOp op) {
mlir::Location loc = op.getLoc();
CIRBaseBuilderTy builder(getContext());
@@ -2297,10 +2343,6 @@ void LoweringPreparePass::runOnOp(mlir::Operation *op) {
lowerCastOp(cast);
} else if (auto complexConj = mlir::dyn_cast<cir::ComplexConjOp>(op)) {
lowerComplexConjOp(complexConj);
- } else if (auto complexDiv = mlir::dyn_cast<cir::ComplexDivOp>(op)) {
- lowerComplexDivOp(complexDiv);
- } else if (auto complexMul = mlir::dyn_cast<cir::ComplexMulOp>(op)) {
- lowerComplexMulOp(complexMul);
} else if (auto glob = mlir::dyn_cast<cir::GlobalOp>(op)) {
lowerGlobalOp(glob);
if (auto regAttr = glob->getAttrOfType<CUDAVarRegistrationInfoAttr>(
@@ -2936,9 +2978,8 @@ void LoweringPreparePass::runOnOperation() {
op->walk([&](mlir::Operation *op) {
if (mlir::isa<cir::ArrayCtor, cir::ArrayDtor, cir::CastOp,
- cir::ComplexConjOp, cir::ComplexMulOp, cir::ComplexDivOp,
- cir::DynamicCastOp, cir::FuncOp, cir::CallOp,
- cir::GetGlobalOp, cir::GlobalOp, cir::StoreOp,
+ cir::ComplexConjOp, cir::DynamicCastOp, cir::FuncOp,
+ cir::CallOp, cir::GetGlobalOp, cir::GlobalOp, cir::StoreOp,
cir::CmpThreeWayOp, cir::LocalInitOp, cir::StdOpInterface>(
op))
opsToTransform.push_back(op);
@@ -2965,3 +3006,14 @@ mlir::createLoweringPreparePass(clang::ASTContext *astCtx) {
pass->setASTContext(astCtx);
return std::move(pass);
}
+
+std::unique_ptr<Pass> mlir::createComplexLoweringPass() {
+ return std::make_unique<ComplexLoweringPass>();
+}
+
+std::unique_ptr<Pass>
+mlir::createComplexLoweringPass(clang::ASTContext *astCtx) {
+ auto pass = std::make_unique<ComplexLoweringPass>();
+ pass->setASTContext(astCtx);
+ return std::move(pass);
+}
diff --git a/clang/lib/CIR/Lowering/CIRPasses.cpp b/clang/lib/CIR/Lowering/CIRPasses.cpp
index 1d1fdaf42aaa4..1d1501e3719ee 100644
--- a/clang/lib/CIR/Lowering/CIRPasses.cpp
+++ b/clang/lib/CIR/Lowering/CIRPasses.cpp
@@ -106,6 +106,12 @@ runCIRToCIRPasses(mlir::ModuleOp theModule, mlir::MLIRContext &mlirContext,
pm.addPass(mlir::createTargetLoweringPass());
pm.addPass(mlir::createCXXABILoweringPass());
+ // Complex multiplication and division synthesize calls to runtime helpers
+ // such as __mulsc3 and __divsc3, so they must be expanded before
+ // CallConvLowering classifies calls. The rest of the lowering-prepare work
+ // stays after it.
+ pm.addPass(mlir::createComplexLoweringPass(&astContext));
+
if (enableCallConvLowering) {
// CallConvLowering rewrites signatures and call sites using the classifier,
// so it must run after CXXABILowering has lowered C++ ABI types to plain
diff --git a/clang/test/CIR/CodeGen/complex-compound-assignment.cpp b/clang/test/CIR/CodeGen/complex-compound-assignment.cpp
index 8e58c51c570a9..e56717f064fa6 100644
--- a/clang/test/CIR/CodeGen/complex-compound-assignment.cpp
+++ b/clang/test/CIR/CodeGen/complex-compound-assignment.cpp
@@ -415,7 +415,10 @@ void foo7() {
// CIR: %[[CONST_FALSE:.*]] = cir.const #false
// CIR: %[[SELECT_CONDITION:.*]] = cir.select if %[[IS_C_REAL_NAN]] then %[[IS_C_IMAG_NAN]] else %[[CONST_FALSE]] : (!cir.bool, !cir.bool, !cir.bool) -> !cir.bool
// CIR: %[[RESULT:.*]] = cir.ternary(%[[SELECT_CONDITION]], true {
-// CIR: %[[LIBC_COMPLEX:.*]] = cir.call @__mulsc3(%[[B_REAL]], %[[B_IMAG]], %[[A_REAL]], %[[A_IMAG]]) : (!cir.float, !cir.float, !cir.float, !cir.float) -> !cir.complex<!cir.float>
+// CIR: %[[LIBC_COERCED:.*]] = cir.call @__mulsc3(%[[B_REAL]], %[[B_IMAG]], %[[A_REAL]], %[[A_IMAG]]) : (!cir.float, !cir.float, !cir.float, !cir.float) -> !cir.vector<2 x !cir.float>
+// CIR: cir.store %[[LIBC_COERCED]], %[[COERCE_SLOT:.*]] : !cir.vector<2 x !cir.float>, !cir.ptr<!cir.vector<2 x !cir.float>>
+// CIR: %[[COERCE_PTR:.*]] = cir.cast bitcast %[[COERCE_SLOT]] : !cir.ptr<!cir.vector<2 x !cir.float>> -> !cir.ptr<!cir.complex<!cir.float>>
+// CIR: %[[LIBC_COMPLEX:.*]] = cir.load %[[COERCE_PTR]] : !cir.ptr<!cir.complex<!cir.float>>, !cir.complex<!cir.float>
// CIR: cir.yield %[[LIBC_COMPLEX]] : !cir.complex<!cir.float>
// CIR: }, false {
// CIR: cir.yield %[[COMPLEX]] : !cir.complex<!cir.float>
@@ -443,7 +446,9 @@ void foo7() {
// LLVM: %[[SELECT_CONDITION:.*]] = and i1 %[[IS_C_REAL_NAN]], %[[IS_C_IMAG_NAN]]
// LLVM: br i1 %[[SELECT_CONDITION]], label %[[THEN_LABEL:.*]], label %[[ELSE_LABEL:.*]]
// LLVM: [[THEN_LABEL]]:
-// LLVM: %[[LIBC_COMPLEX:.*]] = call { float, float } @__mulsc3(float %[[B_REAL]], float %[[B_IMAG]], float %[[A_REAL]], float %[[A_IMAG]])
+// LLVM: %[[LIBC_COERCED:.*]] = call <2 x float> @__mulsc3(float %[[B_REAL]], float %[[B_IMAG]], float %[[A_REAL]], float %[[A_IMAG]])
+// LLVM: store <2 x float> %[[LIBC_COERCED]], ptr %[[COERCE_SLOT:.*]], align 8
+// LLVM: %[[LIBC_COMPLEX:.*]] = load { float, float }, ptr %[[COERCE_SLOT]], align 4
// LLVM: br label %[[PHI_BRANCH:.*]]
// LLVM: [[ELSE_LABEL]]:
// LLVM: br label %[[PHI_BRANCH:]]
@@ -548,7 +553,10 @@ void foo10() {
// CIR: %[[A_IMAG:.*]] = cir.complex.imag %[[TMP_A]] : !cir.complex<!cir.float> -> !cir.float
// CIR: %[[B_REAL:.*]] = cir.complex.real %[[TMP_B]] : !cir.complex<!cir.float> -> !cir.float
// CIR: %[[B_IMAG:.*]] = cir.complex.imag %[[TMP_B]] : !cir.complex<!cir.float> -> !cir.float
-// CIR: %[[RESULT:.*]] = cir.call @__divsc3(%[[A_REAL]], %[[A_IMAG]], %[[B_REAL]], %[[B_IMAG]]) : (!cir.float, !cir.float, !cir.float, !cir.float) -> !cir.complex<!cir.float>
+// CIR: %[[COERCED:.*]] = cir.call @__divsc3(%[[A_REAL]], %[[A_IMAG]], %[[B_REAL]], %[[B_IMAG]]) : (!cir.float, !cir.float, !cir.float, !cir.float) -> !cir.vector<2 x !cir.float>
+// CIR: cir.store %[[COERCED]], %[[COERCE_SLOT:.*]] : !cir.vector<2 x !cir.float>, !cir.ptr<!cir.vector<2 x !cir.float>>
+// CIR: %[[COERCE_PTR:.*]] = cir.cast bitcast %[[COERCE_SLOT]] : !cir.ptr<!cir.vector<2 x !cir.float>> -> !cir.ptr<!cir.complex<!cir.float>>
+// CIR: %[[RESULT:.*]] = cir.load %[[COERCE_PTR]] : !cir.ptr<!cir.complex<!cir.float>>, !cir.complex<!cir.float>
// CIR: cir.store{{.*}} %[[RESULT]], %[[A_ADDR]] : !cir.complex<!cir.float>, !cir.ptr<!cir.complex<!cir.float>>
// LLVM: %[[A_ADDR:.*]] = alloca { float, float }, align 4
@@ -559,7 +567,9 @@ void foo10() {
// LLVM: %[[A_IMAG:.*]] = extractvalue { float, float } %[[TMP_A]], 1
// LLVM: %[[B_REAL:.*]] = extractvalue { float, float } %[[TMP_B]], 0
// LLVM: %[[B_IMAG:.*]] = extractvalue { float, float } %[[TMP_B]], 1
-// LLVM: %[[RESULT:.*]] = call { float, float } @__divsc3(float %[[A_REAL]], float %[[A_IMAG]], float %[[B_REAL]], float %[[B_IMAG]])
+// LLVM: %[[COERCED:.*]] = call <2 x float> @__divsc3(float %[[A_REAL]], float %[[A_IMAG]], float %[[B_REAL]], float %[[B_IMAG]])
+// LLVM: store <2 x float> %[[COERCED]], ptr %[[COERCE_SLOT:.*]], align 8
+// LLVM: %[[RESULT:.*]] = load { float, float }, ptr %[[COERCE_SLOT]], align 4
// LLVM: store { float, float } %[[RESULT]], ptr %[[A_ADDR]], align 4
// OGCG: %[[A_ADDR:.*]] = alloca { float, float }, align 4
@@ -725,7 +735,10 @@ void foo13() {
// CIR: %[[A_IMAG_F32:.*]] = cir.complex.imag %[[A_COMPLEX_F32]] : !cir.complex<!cir.float> -> !cir.float
// CIR: %[[B_REAL_F32:.*]] = cir.complex.real %[[B_COMPLEX_F32]] : !cir.complex<!cir.float> -> !cir.float
// CIR: %[[B_IMAG_F32:.*]] = cir.complex.imag %[[B_COMPLEX_F32]] : !cir.complex<!cir.float> -> !cir.float
-// CIR: %[[DIV_A_B:.*]] = cir.call @__divsc3(%[[A_REAL_F32]], %[[A_IMAG_F32]], %[[B_REAL_F32]], %[[B_IMAG_F32]]) : (!cir.float, !cir.float, !cir.float, !cir.float) -> !cir.complex<!cir.float>
+// CIR: %[[DIV_A_B_COERCED:.*]] = cir.call @__divsc3(%[[A_REAL_F32]], %[[A_IMAG_F32]], %[[B_REAL_F32]], %[[B_IMAG_F32]]) : (!cir.float, !cir.float, !cir.float, !cir.float) -> !cir.vector<2 x !cir.float>
+// CIR: cir.store %[[DIV_A_B_COERCED]], %[[SLOT_AB:.*]] : !cir.vector<2 x !cir.float>, !cir.ptr<!cir.vector<2 x !cir.float>>
+// CIR: %[[SLOT_AB_PTR:.*]] = cir.cast bitcast %[[SLOT_AB]] : !cir.ptr<!cir.vector<2 x !cir.float>> -> !cir.ptr<!cir.complex<!cir.float>>
+// CIR: %[[DIV_A_B:.*]] = cir.load %[[SLOT_AB_PTR]] : !cir.ptr<!cir.complex<!cir.float>>, !cir.complex<!cir.float>
// CIR: %[[TMP_B:.*]] = cir.load{{.*}} %[[B_ADDR]] : !cir.ptr<!cir.complex<!cir.f16>>, !cir.complex<!cir.f16>
// CIR: %[[B_REAL:.*]] = cir.complex.real %[[TMP_B]] : !cir.complex<!cir.f16> -> !cir.f16
// CIR: %[[B_IMAG:.*]] = cir.complex.imag %[[TMP_B]] : !cir.complex<!cir.f16> -> !cir.f16
@@ -736,7 +749,10 @@ void foo13() {
// CIR: %[[B_IMAG_F32:.*]] = cir.complex.imag %[[B_COMPLEX_F32]] : !cir.complex<!cir.float> -> !cir.float
// CIR: %[[DIV_AB_REAL:.*]] = cir.complex.real %[[DIV_A_B]] : !cir.complex<!cir.float> -> !cir.float
// CIR: %[[DIV_AB_IMAG:.*]] = cir.complex.imag %[[DIV_A_B]] : !cir.complex<!cir.float> -> !cir.float
-// CIR: %[[RESULT:.*]] = cir.call @__divsc3(%[[B_REAL_F32]], %[[B_IMAG_F32]], %[[DIV_AB_REAL]], %[[DIV_AB_IMAG]]) : (!cir.float, !cir.float, !cir.float, !cir.float) -> !cir.complex<!cir.float>
+// CIR: %[[RESULT_COERCED:.*]] = cir.call @__divsc3(%[[B_REAL_F32]], %[[B_IMAG_F32]], %[[DIV_AB_REAL]], %[[DIV_AB_IMAG]]) : (!cir.float, !cir.float, !cir.float, !cir.float) -> !cir.vector<2 x !cir.float>
+// CIR: cir.store %[[RESULT_COERCED]], %[[SLOT_R:.*]] : !cir.vector<2 x !cir.float>, !cir.ptr<!cir.vector<2 x !cir.float>>
+// CIR: %[[SLOT_R_PTR:.*]] = cir.cast bitcast %[[SLOT_R]] : !cir.ptr<!cir.vector<2 x !cir.float>> -> !cir.ptr<!cir.complex<!cir.float>>
+// CIR: %[[RESULT:.*]] = cir.load %[[SLOT_R_PTR]] : !cir.ptr<!cir.complex<!cir.float>>, !cir.complex<!cir.float>
// CIR: %[[RESULT_REAL_F32:.*]] = cir.complex.real %[[RESULT]] : !cir.complex<!cir.float> -> !cir.float
// CIR: %[[RESULT_IMAG_F32:.*]] = cir.complex.imag %[[RESULT]] : !cir.complex<!cir.float> -> !cir.float
// CIR: %[[RESULT_REAL_F16:.*]] = cir.cast floating %[[RESULT_REAL_F32]] : !cir.float -> !cir.f16
@@ -760,7 +776,9 @@ void foo13() {
// LLVM: %[[B_IMAG_F32:.*]] = fpext half %[[B_IMAG]] to float
// LLVM: %[[TMP_B_COMPLEX_F32:.*]] = insertvalue { float, float } {{.*}}, float %[[B_REAL_F32]], 0
// LLVM: %[[B_COMPLEX_F32:.*]] = insertvalue { float, float } %[[TMP_B_COMPLEX_F32]], float %[[B_IMAG_F32]], 1
-// LLVM: %[[DIV_A_B:.*]] = call { float, float } @__divsc3(float %[[A_REAL_F32]], float %[[A_IMAG_F32]], float %[[B_REAL_F32]], float %[[B_IMAG_F32]])
+// LLVM: %[[DIV_A_B_COERCED:.*]] = call <2 x float> @__divsc3(float %[[A_REAL_F32]], float %[[A_IMAG_F32]], float %[[B_REAL_F32]], float %[[B_IMAG_F32]])
+// LLVM: store <2 x float> %[[DIV_A_B_COERCED]], ptr %[[SLOT_AB:.*]], align 8
+// LLVM: %[[DIV_A_B:.*]] = load { float, float }, ptr %[[SLOT_AB]], align 4
// LLVM: %[[TMP_B:.*]] = load { half, half }, ptr %[[B_ADDR]], align 2
// LLVM: %[[B_REAL:.*]] = extractvalue { half, half } %[[TMP_B]], 0
// LLVM: %[[B_IMAG:.*]] = extractvalue { half, half } %[[TMP_B]], 1
@@ -770,7 +788,9 @@ void foo13() {
// LLVM: %[[B_COMPLEX_F32:.*]] = insertvalue { float, float } %[[TMP_B_COMPLEX_F32]], float %[[B_IMAG_F32]], 1
// LLVM: %[[DIV_AB_REAL:.*]] = extractvalue { float, float } %[[DIV_A_B]], 0
// LLVM: %[[DIV_AB_IMAG:.*]] = extractvalue { float, float } %[[DIV_A_B]], 1
-// LLVM: %[[RESULT:.*]] = call { float, float } @__divsc3(float %[[B_REAL_F32]], float %[[B_IMAG_F32]], float %[[DIV_AB_REAL]], float %[[DIV_AB_IMAG]])
+// LLVM: %[[RESULT_COERCED:.*]] = call <2 x float> @__divsc3(float %[[B_REAL_F32]], float %[[B_IMAG_F32]], float %[[DIV_AB_REAL]], float %[[DIV_AB_IMAG]])
+// LLVM: store <2 x float> %[[RESULT_COERCED]], ptr %[[SLOT_R:.*]], align 8
+// LLVM: %[[RESULT:.*]] = load { float, float }, ptr %[[SLOT_R]], align 4
// LLVM: %[[RESULT_REAL_F32:.*]] = extractvalue { float, float } %[[RESULT]], 0
// LLVM: %[[RESULT_IMAG_F32:.*]] = extractvalue { float, float } %[[RESULT]], 1
// LLVM: %[[RESULT_REAL_F16:.*]] = fptrunc float %[[RESULT_REAL_F32]] to half
diff --git a/clang/test/CIR/CodeGen/complex-libcall-abi-global-init.cpp b/clang/test/CIR/CodeGen/complex-libcall-abi-global-init.cpp
new file mode 100644
index 0000000000000..3ec35a24d3ced
--- /dev/null
+++ b/clang/test/CIR/CodeGen/complex-libcall-abi-global-init.cpp
@@ -0,0 +1,30 @@
+// RUN: %clang_cc1 -triple x86_64-unknown-linux-gnu -std=c++20 -fclangir -emit-cir %s -o %t.cir
+// RUN: FileCheck --check-prefix=CIR --input-file=%t.cir %s
+// RUN: %clang_cc1 -triple x86_64-unknown-linux-gnu -std=c++20 -fclangir -emit-llvm %s -o %t-cir.ll
+// RUN: FileCheck --check-prefix=LLVM --input-file=%t-cir.ll %s
+// RUN: %clang_cc1 -triple x86_64-unknown-linux-gnu -std=c++20 -emit-llvm %s -o %t.ll
+// RUN: FileCheck --check-prefix=OGCG --input-file=%t.ll %s
+
+extern float _Complex a;
+extern float _Complex b;
+
+// A dynamic initializer at namespace scope is still inside the global's
+// initializer region when the helper call is coerced, so the coercion slot has
+// no enclosing function to be placed in. It has to land in the region that
+// later becomes the initializer function.
+float _Complex g = a / b;
+
+// CIR-LABEL: cir.func {{.*}}@__cxx_global_var_init
+// CIR: %[[SLOT:.*]] = cir.alloca "coerce"{{.*}} : !cir.ptr<!cir.vector<2 x !cir.float>>
+// CIR: %[[COERCED:.*]] = cir.call @__divsc3({{.*}}) : (!cir.float, !cir.float, !cir.float, !cir.float) -> !cir.vector<2 x !cir.float>
+// CIR: cir.store %[[COERCED]], %[[SLOT]] : !cir.vector<2 x !cir.float>, !cir.ptr<!cir.vector<2 x !cir.float>>
+// CIR: cir.store{{.*}} %{{.+}}, %{{.+}} : !cir.complex<!cir.float>, !cir.ptr<!cir.complex<!cir.float>>
+
+// LLVM-LABEL: @__cxx_global_var_init(
+// LLVM: %[[SLOT:.+]] = alloca <2 x float>, align 8
+// LLVM: %[[COERCED:.*]] = call <2 x float> @__divsc3(float %{{.+}}, float %{{.+}}, float %{{.+}}, float %{{.+}})
+// LLVM: store <2 x float> %[[COERCED]], ptr %[[SLOT]], align 8
+// LLVM: store { float, float } %{{.+}}, ptr @g, align 4
+
+// OGCG-LABEL: @__cxx_global_var_init(
+// OGCG: call noundef <2 x float> @__divsc3(float noundef %{{.+}}, float noundef %{{.+}}, float noundef %{{.+}}, float noundef %{{.+}})
diff --git a/clang/test/CIR/CodeGen/complex-libcall-abi.c b/clang/test/CIR/CodeGen/complex-libcall-abi.c
new file mode 100644
index 0000000000000..efc459a30ad64
--- /dev/null
+++ b/clang/test/CIR/CodeGen/complex-libcall-abi.c
@@ -0,0 +1,76 @@
+// RUN: %clang_cc1 -triple x86_64-unknown-linux-gnu -fclangir -emit-cir %s -o %t.cir
+// RUN: FileCheck --check-prefix=CIR --input-file=%t.cir %s
+// RUN: %clang_cc1 -triple x86_64-unknown-linux-gnu -fclangir -emit-llvm %s -o %t-cir.ll
+// RUN: FileCheck --check-prefixes=LLVM,LLVMCIR --input-file=%t-cir.ll %s
+// RUN: %clang_cc1 -triple x86_64-unknown-linux-gnu -emit-llvm %s -o %t.ll
+// RUN: FileCheck --check-prefixes=LLVM,OGCG --input-file=%t.ll %s
+
+float _Complex divf(float _Complex a, float _Complex b) { return a / b; }
+
+// A float pair fits one eightbyte and returns in a single SSE register, so the
+// helper's return coerces to a vector and comes back through memory.
+
+// CIR-LABEL: cir.func {{.*}}@divf
+// CIR: %[[COERCED:.*]] = cir.call @__divsc3({{.*}}) : (!cir.float, !cir.float, !cir.float, !cir.float) -> !cir.vector<2 x !cir.float>
+// CIR: cir.store %[[COERCED]], %[[SLOT:.*]] : !cir.vector<2 x !cir.float>, !cir.ptr<!cir.vector<2 x !cir.float>>
+// CIR: %[[SLOT_PTR:.*]] = cir.cast bitcast %[[SLOT]] : !cir.ptr<!cir.vector<2 x !cir.float>> -> !cir.ptr<!cir.complex<!cir.float>>
+// CIR: cir.load %[[SLOT_PTR]] : !cir.ptr<!cir.complex<!cir.float>>, !cir.complex<!cir.float>
+
+// The caller's own signature is coerced the same way on both paths.
+// LLVM: define dso_local <2 x float> @divf(<2 x float> noundef %{{.+}}, <2 x float> noundef %{{.+}})
+
+// LLVMCIR: %[[COERCED:.*]] = call <2 x float> @__divsc3(float %{{.+}}, float %{{.+}}, float %{{.+}}, float %{{.+}})
+// LLVMCIR: store <2 x float> %[[COERCED]], ptr %[[SLOT:.+]], align 8
+// LLVMCIR: load { float, float }, ptr %[[SLOT]], align 4
+// OGCG: call <2 x float> @__divsc3(float noundef %{{.+}}, float noundef %{{.+}}, float noundef %{{.+}}, float noundef %{{.+}})
+
+float _Complex mulf(float _Complex a, float _Complex b) { return a * b; }
+
+// CIR-LABEL: cir.func {{.*}}@mulf
+// CIR: cir.call @__mulsc3({{.*}}) : (!cir.float, !cir.float, !cir.float, !cir.float) -> !cir.vector<2 x !cir.float>
+
+// LLVM: define dso_local <2 x float> @mulf(<2 x float> noundef %{{.+}}, <2 x float> noundef %{{.+}})
+// LLVMCIR: call <2 x float> @__mulsc3(float %{{.+}}, float %{{.+}}, float %{{.+}}, float %{{.+}})
+// OGCG: call <2 x float> @__mulsc3(float noundef %{{.+}}, float noundef %{{.+}}, float noundef %{{.+}}, float noundef %{{.+}})
+
+double _Complex divd(double _Complex a, double _Complex b) { return a / b; }
+
+// A double pair spans two eightbytes and returns in two registers, so it stays
+// a two-field record and its lowered call is unchanged by the coercion.
+
+// CIR-LABEL: cir.func {{.*}}@divd
+// CIR: cir.call @__divdc3({{.*}}) : (!cir.double, !cir.double, !cir.double, !cir.double) -> [[REC_D:!rec_anon_struct[0-9]*]]
+
+// LLVMCIR: call { double, double } @__divdc3(double %{{.+}}, double %{{.+}}, double %{{.+}}, double %{{.+}})
+// OGCG: call { double, double } @__divdc3(double noundef %{{.+}}, double noundef %{{.+}}, double noundef %{{.+}}, double noundef %{{.+}})
+
+double _Complex muld(double _Complex a, double _Complex b) { return a * b; }
+
+// CIR-LABEL: cir.func {{.*}}@muld
+// CIR: cir.call @__muldc3({{.*}}) : (!cir.double, !cir.double, !cir.double, !cir.double) -> [[REC_D]]
+
+// LLVMCIR: call { double, double } @__muldc3(double %{{.+}}, double %{{.+}}, double %{{.+}}, double %{{.+}})
+// OGCG: call { double, double } @__muldc3(double noundef %{{.+}}, double noundef %{{.+}}, double noundef %{{.+}}, double noundef %{{.+}})
+
+long double _Complex divld(long double _Complex a, long double _Complex b) {
+ return a / b;
+}
+
+// A long double pair is x87-classified and returned in memory, so it also
+// keeps a two-field record.
+
+// CIR-LABEL: cir.func {{.*}}@divld
+// CIR: cir.call @__divxc3({{.*}}) : (!cir.long_double<!cir.f80>, !cir.long_double<!cir.f80>, !cir.long_double<!cir.f80>, !cir.long_double<!cir.f80>) -> [[REC_LD:!rec_anon_struct[0-9]*]]
+
+// LLVMCIR: call { x86_fp80, x86_fp80 } @__divxc3(x86_fp80 %{{.+}}, x86_fp80 %{{.+}}, x86_fp80 %{{.+}}, x86_fp80 %{{.+}})
+// OGCG: call { x86_fp80, x86_fp80 } @__divxc3(x86_fp80 noundef %{{.+}}, x86_fp80 noundef %{{.+}}, x86_fp80 noundef %{{.+}}, x86_fp80 noundef %{{.+}})
+
+long double _Complex mulld(long double _Complex a, long double _Complex b) {
+ return a * b;
+}
+
+// CIR-LABEL: cir.func {{.*}}@mulld
+// CIR: cir.call @__mulxc3({{.*}}) : (!cir.long_double<!cir.f80>, !cir.long_double<!cir.f80>, !cir.long_double<!cir.f80>, !cir.long_double<!cir.f80>) -> [[REC_LD]]
+
+// LLVMCIR: call { x86_fp80, x86_fp80 } @__mulxc3(x86_fp80 %{{.+}}, x86_fp80 %{{.+}}, x86_fp80 %{{.+}}, x86_fp80 %{{.+}})
+// OGCG: call { x86_fp80, x86_fp80 } @__mulxc3(x86_fp80 noundef %{{.+}}, x86_fp80 noundef %{{.+}}, x86_fp80 noundef %{{.+}}, x86_fp80 noundef %{{.+}})
diff --git a/clang/test/CIR/CodeGen/complex-mul-div.cpp b/clang/test/CIR/CodeGen/complex-mul-div.cpp
index 50d77e5a6c2a5..13e3b39f7f10b 100644
--- a/clang/test/CIR/CodeGen/complex-mul-div.cpp
+++ b/clang/test/CIR/CodeGen/complex-mul-div.cpp
@@ -128,7 +128,10 @@ void foo() {
// CIR-AFTER-FULL: %[[CONST_FALSE:.*]] = cir.const #false
// CIR-AFTER-FULL: %[[SELECT_CONDITION:.*]] = cir.select if %[[IS_C_REAL_NAN]] then %[[IS_C_IMAG_NAN]] else %[[CONST_FALSE]] : (!cir.bool, !cir.bool, !cir.bool) -> !cir.bool
// CIR-AFTER-FULL: %[[RESULT:.*]] = cir.ternary(%[[SELECT_CONDITION]], true {
-// CIR-AFTER-FULL: %[[LIBC_COMPLEX:.*]] = cir.call @__mulsc3(%[[A_REAL]], %[[A_IMAG]], %[[B_REAL]], %[[B_IMAG]]) : (!cir.float, !cir.float, !cir.float, !cir.float) -> !cir.complex<!cir.float>
+// CIR-AFTER-FULL: %[[LIBC_COERCED:.*]] = cir.call @__mulsc3(%[[A_REAL]], %[[A_IMAG]], %[[B_REAL]], %[[B_IMAG]]) : (!cir.float, !cir.float, !cir.float, !cir.float) -> !cir.vector<2 x !cir.float>
+// CIR-AFTER-FULL: cir.store %[[LIBC_COERCED]], %[[COERCE_SLOT:.*]] : !cir.vector<2 x !cir.float>, !cir.ptr<!cir.vector<2 x !cir.float>>
+// CIR-AFTER-FULL: %[[COERCE_PTR:.*]] = cir.cast bitcast %[[COERCE_SLOT]] : !cir.ptr<!cir.vector<2 x !cir.float>> -> !cir.ptr<!cir.complex<!cir.float>>
+// CIR-AFTER-FULL: %[[LIBC_COMPLEX:.*]] = cir.load %[[COERCE_PTR]] : !cir.ptr<!cir.complex<!cir.float>>, !cir.complex<!cir.float>
// CIR-AFTER-FULL: cir.yield %[[LIBC_COMPLEX]] : !cir.complex<!cir.float>
// CIR-AFTER-FULL: }, false {
// CIR-AFTER-FULL: cir.yield %[[COMPLEX]] : !cir.complex<!cir.float>
@@ -157,7 +160,9 @@ void foo() {
// LLVM-FULL: %[[SELECT_CONDITION:.*]] = and i1 %[[IS_C_REAL_NAN]], %[[IS_C_IMAG_NAN]]
// LLVM-FULL: br i1 %[[SELECT_CONDITION]], label %[[THEN_LABEL:.*]], label %[[ELSE_LABEL:.*]]
// LLVM-FULL: [[THEN_LABEL]]:
-// LLVM-FULL: %[[LIBC_COMPLEX:.*]] = call { float, float } @__mulsc3(float %[[A_REAL]], float %[[A_IMAG]], float %[[B_REAL]], float %[[B_IMAG]])
+// LLVM-FULL: %[[LIBC_COERCED:.*]] = call <2 x float> @__mulsc3(float %[[A_REAL]], float %[[A_IMAG]], float %[[B_REAL]], float %[[B_IMAG]])
+// LLVM-FULL: store <2 x float> %[[LIBC_COERCED]], ptr %[[COERCE_SLOT:.*]], align 8
+// LLVM-FULL: %[[LIBC_COMPLEX:.*]] = load { float, float }, ptr %[[COERCE_SLOT]], align 4
// LLVM-FULL: br label %[[PHI_BRANCH:.*]]
// LLVM-FULL: [[ELSE_LABEL]]:
// LLVM-FULL: br label %[[PHI_BRANCH:]]
@@ -648,7 +653,10 @@ void foo3() {
// CIR-AFTER-FULL: %[[A_IMAG:.*]] = cir.complex.imag %[[TMP_A]] : !cir.complex<!cir.float> -> !cir.float
// CIR-AFTER-FULL: %[[B_REAL:.*]] = cir.complex.real %[[TMP_B]] : !cir.complex<!cir.float> -> !cir.float
// CIR-AFTER-FULL: %[[B_IMAG:.*]] = cir.complex.imag %[[TMP_B]] : !cir.complex<!cir.float> -> !cir.float
-// CIR-AFTER-FULL: %[[RESULT:.*]] = cir.call @__divsc3(%[[A_REAL]], %[[A_IMAG]], %[[B_REAL]], %[[B_IMAG]]) : (!cir.float, !cir.float, !cir.float, !cir.float) -> !cir.complex<!cir.float>
+// CIR-AFTER-FULL: %[[COERCED:.*]] = cir.call @__divsc3(%[[A_REAL]], %[[A_IMAG]], %[[B_REAL]], %[[B_IMAG]]) : (!cir.float, !cir.float, !cir.float, !cir.float) -> !cir.vector<2 x !cir.float>
+// CIR-AFTER-FULL: cir.store %[[COERCED]], %[[COERCE_SLOT:.*]] : !cir.vector<2 x !cir.float>, !cir.ptr<!cir.vector<2 x !cir.float>>
+// CIR-AFTER-FULL: %[[COERCE_PTR:.*]] = cir.cast bitcast %[[COERCE_SLOT]] : !cir.ptr<!cir.vector<2 x !cir.float>> -> !cir.ptr<!cir.complex<!cir.float>>
+// CIR-AFTER-FULL: %[[RESULT:.*]] = cir.load %[[COERCE_PTR]] : !cir.ptr<!cir.complex<!cir.float>>, !cir.complex<!cir.float>
// CIR-AFTER-FULL: cir.store{{.*}} %[[RESULT]], %[[C_ADDR]] : !cir.complex<!cir.float>, !cir.ptr<!cir.complex<!cir.float>>
// LLVM-FULL: %[[A_ADDR:.*]] = alloca { float, float }, align 4
@@ -660,7 +668,9 @@ void foo3() {
// LLVM-FULL: %[[A_IMAG:.*]] = extractvalue { float, float } %[[TMP_A]], 1
// LLVM-FULL: %[[B_REAL:.*]] = extractvalue { float, float } %[[TMP_B]], 0
// LLVM-FULL: %[[B_IMAG:.*]] = extractvalue { float, float } %[[TMP_B]], 1
-// LLVM-FULL: %[[RESULT:.*]] = call { float, float } @__divsc3(float %[[A_REAL]], float %[[A_IMAG]], float %[[B_REAL]], float %[[B_IMAG]])
+// LLVM-FULL: %[[COERCED:.*]] = call <2 x float> @__divsc3(float %[[A_REAL]], float %[[A_IMAG]], float %[[B_REAL]], float %[[B_IMAG]])
+// LLVM-FULL: store <2 x float> %[[COERCED]], ptr %[[COERCE_SLOT:.*]], align 8
+// LLVM-FULL: %[[RESULT:.*]] = load { float, float }, ptr %[[COERCE_SLOT]], align 4
// LLVM-FULL: store { float, float } %[[RESULT]], ptr %[[C_ADDR]], align 4
// OGCG-FULL: %[[A_ADDR:.*]] = alloca { float, float }, align 4
@@ -1140,7 +1150,10 @@ void foo6() {
// CIR-AFTER-FULL: %[[A_IMAG:.*]] = cir.complex.imag %[[COMPLEX_A]] : !cir.complex<!cir.float> -> !cir.float
// CIR-AFTER-FULL: %[[B_REAL:.*]] = cir.complex.real %[[TMP_B]] : !cir.complex<!cir.float> -> !cir.float
// CIR-AFTER-FULL: %[[B_IMAG:.*]] = cir.complex.imag %[[TMP_B]] : !cir.complex<!cir.float> -> !cir.float
-// CIR-AFTER-FULL: %[[RESULT:.*]] = cir.call @__divsc3(%[[A_REAL]], %[[A_IMAG]], %[[B_REAL]], %[[B_IMAG]]) : (!cir.float, !cir.float, !cir.float, !cir.float) -> !cir.complex<!cir.float>
+// CIR-AFTER-FULL: %[[COERCED:.*]] = cir.call @__divsc3(%[[A_REAL]], %[[A_IMAG]], %[[B_REAL]], %[[B_IMAG]]) : (!cir.float, !cir.float, !cir.float, !cir.float) -> !cir.vector<2 x !cir.float>
+// CIR-AFTER-FULL: cir.store %[[COERCED]], %[[COERCE_SLOT:.*]] : !cir.vector<2 x !cir.float>, !cir.ptr<!cir.vector<2 x !cir.float>>
+// CIR-AFTER-FULL: %[[COERCE_PTR:.*]] = cir.cast bitcast %[[COERCE_SLOT]] : !cir.ptr<!cir.vector<2 x !cir.float>> -> !cir.ptr<!cir.complex<!cir.float>>
+// CIR-AFTER-FULL: %[[RESULT:.*]] = cir.load %[[COERCE_PTR]] : !cir.ptr<!cir.complex<!cir.float>>, !cir.complex<!cir.float>
// CIR-AFTER-FULL: cir.store{{.*}} %[[RESULT]], %[[C_ADDR]] : !cir.complex<!cir.float>, !cir.ptr<!cir.complex<!cir.float>>
// LLVM-FULL: %[[A_ADDR:.*]] = alloca float, align 4
@@ -1149,10 +1162,12 @@ void foo6() {
// LLVM-FULL: %[[TMP_A:.*]] = load float, ptr %[[A_ADDR]], align 4
// LLVM-FULL: %[[TMP_B:.*]] = load { float, float }, ptr %[[B_ADDR]], align 4
// LLVM-FULL: %[[TMP_COMPLEX_A:.*]] = insertvalue { float, float } {{.*}}, float %[[TMP_A]], 0
-// LLVM-FULL: %[[COMPLEX_A:.*]] = insertvalue { float, float } %6, float 0.000000e+00, 1
+// LLVM-FULL: %[[COMPLEX_A:.*]] = insertvalue { float, float } %[[TMP_COMPLEX_A]], float 0.000000e+00, 1
// LLVM-FULL: %[[B_REAL:.*]] = extractvalue { float, float } %[[TMP_B]], 0
// LLVM-FULL: %[[B_IMAG:.*]] = extractvalue { float, float } %[[TMP_B]], 1
-// LLVM-FULL: %[[RESULT:.*]] = call { float, float } @__divsc3(float %[[TMP_A]], float 0.000000e+00, float %[[B_REAL]], float %[[B_IMAG]])
+// LLVM-FULL: %[[COERCED:.*]] = call <2 x float> @__divsc3(float %[[TMP_A]], float 0.000000e+00, float %[[B_REAL]], float %[[B_IMAG]])
+// LLVM-FULL: store <2 x float> %[[COERCED]], ptr %[[COERCE_SLOT:.*]], align 8
+// LLVM-FULL: %[[RESULT:.*]] = load { float, float }, ptr %[[COERCE_SLOT]], align 4
// LLVM-FULL: store { float, float } %[[RESULT]], ptr %[[C_ADDR]], align 4
// OGCG-FULL: %[[A_ADDR:.*]] = alloca float, align 4
>From 37d718110582b6917489af1f7273c5ba4a8f537e Mon Sep 17 00:00:00 2001
From: Adam Smith <adams at nvidia.com>
Date: Mon, 17 Aug 2026 11:39:27 -0700
Subject: [PATCH 2/3] [CIR] Move complex lowering out of LoweringPrepare
The new pass was declared and defined inside LoweringPrepare.cpp, which left it
entangled with the pass it had just been split out of. It moves to
ComplexLowering.cpp alongside the other CIR transforms, taking the mul and div
expansion with it, and LoweringPrepare keeps only the complex conj and cast
lowering it still owns.
`buildRuntimeFunction` was a file-static helper reached through a member
forwarder, which kept a generally useful routine private to one file. It moves
to CIRTransformUtils, where the shared transform helpers live, and the call
sites in LoweringPrepare now call it directly so the forwarder is gone.
Assisted-by: Cursor / claude-opus-5
---
clang/include/clang/CIR/Dialect/Passes.td | 5 +-
.../Dialect/Transforms/CIRTransformUtils.h | 11 +
.../Dialect/Transforms/CIRTransformUtils.cpp | 21 +
.../lib/CIR/Dialect/Transforms/CMakeLists.txt | 1 +
.../Dialect/Transforms/ComplexLowering.cpp | 494 +++++++++++++++
.../Dialect/Transforms/LoweringPrepare.cpp | 578 ++----------------
clang/lib/CIR/Lowering/CIRPasses.cpp | 5 +-
7 files changed, 585 insertions(+), 530 deletions(-)
create mode 100644 clang/lib/CIR/Dialect/Transforms/ComplexLowering.cpp
diff --git a/clang/include/clang/CIR/Dialect/Passes.td b/clang/include/clang/CIR/Dialect/Passes.td
index c0a79d6fde66c..2b98a46ca4f29 100644
--- a/clang/include/clang/CIR/Dialect/Passes.td
+++ b/clang/include/clang/CIR/Dialect/Passes.td
@@ -196,8 +196,9 @@ def ComplexLowering : Pass<"cir-complex-lowering", "mlir::ModuleOp"> {
run never gets its return type coerced, which places the real and
imaginary halves of the result in the wrong registers.
- Only these two operations need the earlier position, so `cir.complex.conj`
- and the complex casts are still expanded by `cir-lowering-prepare`.
+ Among the complex operations, only these two synthesize a call whose
+ return type has to be coerced, so `cir.complex.conj` and the complex casts
+ are still expanded by `cir-lowering-prepare`.
}];
let constructor = "mlir::createComplexLoweringPass()";
let dependentDialects = ["cir::CIRDialect"];
diff --git a/clang/include/clang/CIR/Dialect/Transforms/CIRTransformUtils.h b/clang/include/clang/CIR/Dialect/Transforms/CIRTransformUtils.h
index b708da61bf01f..6b184e1747a27 100644
--- a/clang/include/clang/CIR/Dialect/Transforms/CIRTransformUtils.h
+++ b/clang/include/clang/CIR/Dialect/Transforms/CIRTransformUtils.h
@@ -53,6 +53,17 @@ mlir::Block *replaceThrowWithTryThrow(cir::ThrowOp throwOp,
void collectUnreachable(mlir::Operation *parent,
llvm::SmallVectorImpl<mlir::Operation *> &ops);
+/// Return a declaration of the runtime function \p name in \p mlirModule,
+/// creating a private declaration of type \p type if the module does not
+/// already have one. When a declaration already exists, \p type and
+/// \p linkage are ignored and the existing type is not checked against
+/// \p type. Returns a null FuncOp if \p name resolves to a symbol that is
+/// not a cir::FuncOp.
+cir::FuncOp buildRuntimeFunction(
+ mlir::OpBuilder &builder, mlir::ModuleOp mlirModule, llvm::StringRef name,
+ mlir::Location loc, cir::FuncType type,
+ cir::GlobalLinkageKind linkage = cir::GlobalLinkageKind::ExternalLinkage);
+
} // namespace cir
#endif // LLVM_CLANG_CIR_DIALECT_TRANSFORMS_CIRTRANSFORMUTILS_H
diff --git a/clang/lib/CIR/Dialect/Transforms/CIRTransformUtils.cpp b/clang/lib/CIR/Dialect/Transforms/CIRTransformUtils.cpp
index 1653e673859dd..6476282e63301 100644
--- a/clang/lib/CIR/Dialect/Transforms/CIRTransformUtils.cpp
+++ b/clang/lib/CIR/Dialect/Transforms/CIRTransformUtils.cpp
@@ -9,9 +9,30 @@
#include "clang/CIR/Dialect/Transforms/CIRTransformUtils.h"
#include "clang/CIR/Dialect/IR/CIRTypes.h"
+#include "clang/CIR/MissingFeatures.h"
#include "llvm/ADT/DepthFirstIterator.h"
+cir::FuncOp cir::buildRuntimeFunction(mlir::OpBuilder &builder,
+ mlir::ModuleOp mlirModule,
+ llvm::StringRef name, mlir::Location loc,
+ cir::FuncType type,
+ cir::GlobalLinkageKind linkage) {
+ auto f = mlir::dyn_cast_or_null<cir::FuncOp>(
+ mlir::SymbolTable::lookupNearestSymbolFrom(
+ mlirModule, mlir::StringAttr::get(mlirModule->getContext(), name)));
+ if (!f) {
+ f = cir::FuncOp::create(builder, loc, name, type);
+ f.setLinkageAttr(
+ cir::GlobalLinkageKindAttr::get(builder.getContext(), linkage));
+ mlir::SymbolTable::setSymbolVisibility(
+ f, mlir::SymbolTable::Visibility::Private);
+
+ assert(!cir::MissingFeatures::opFuncExtraAttrs());
+ }
+ return f;
+}
+
void cir::collectUnreachable(mlir::Operation *parent,
llvm::SmallVectorImpl<mlir::Operation *> &ops) {
// For every region under `parent`, find the blocks unreachable from the
diff --git a/clang/lib/CIR/Dialect/Transforms/CMakeLists.txt b/clang/lib/CIR/Dialect/Transforms/CMakeLists.txt
index 82078bf2e8f73..ccd046b5d4b38 100644
--- a/clang/lib/CIR/Dialect/Transforms/CMakeLists.txt
+++ b/clang/lib/CIR/Dialect/Transforms/CMakeLists.txt
@@ -5,6 +5,7 @@ add_clang_library(MLIRCIRTransforms
CIRCanonicalize.cpp
CIRSimplify.cpp
CIRTransformUtils.cpp
+ ComplexLowering.cpp
CXXABILowering.cpp
EHABILowering.cpp
TargetLowering.cpp
diff --git a/clang/lib/CIR/Dialect/Transforms/ComplexLowering.cpp b/clang/lib/CIR/Dialect/Transforms/ComplexLowering.cpp
new file mode 100644
index 0000000000000..be14c38f1e559
--- /dev/null
+++ b/clang/lib/CIR/Dialect/Transforms/ComplexLowering.cpp
@@ -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);
+
+ if (mlir::isa<cir::DoubleType>(type))
+ return cir::LongDoubleType::get(&context, type);
+
+ return type;
+ };
+
+ auto getFloatTypeSemantics =
+ [&cc](mlir::Type type) -> const llvm::fltSemantics & {
+ const clang::TargetInfo &info = cc.getTargetInfo();
+ if (mlir::isa<cir::FP16Type>(type))
+ return info.getHalfFormat();
+
+ if (mlir::isa<cir::BF16Type>(type))
+ return info.getBFloat16Format();
+
+ if (mlir::isa<cir::SingleType>(type))
+ return info.getFloatFormat();
+
+ if (mlir::isa<cir::DoubleType>(type))
+ return info.getDoubleFormat();
+
+ if (mlir::isa<cir::LongDoubleType>(type)) {
+ if (cc.getLangOpts().OpenMP && cc.getLangOpts().OpenMPIsTargetDevice)
+ llvm_unreachable("NYI Float type semantics with OpenMP");
+ return info.getLongDoubleFormat();
+ }
+
+ if (mlir::isa<cir::FP128Type>(type)) {
+ if (cc.getLangOpts().OpenMP && cc.getLangOpts().OpenMPIsTargetDevice)
+ llvm_unreachable("NYI Float type semantics with OpenMP");
+ return info.getFloat128Format();
+ }
+
+ llvm_unreachable("Unsupported float type semantics");
+ };
+
+ const mlir::Type higherElementType = getHigherPrecisionFPType(elementType);
+ const llvm::fltSemantics &elementTypeSemantics =
+ getFloatTypeSemantics(elementType);
+ const llvm::fltSemantics &higherElementTypeSemantics =
+ getFloatTypeSemantics(higherElementType);
+
+ // Check that the promoted type can handle the intermediate values without
+ // overflowing. This can be interpreted as:
+ // (SmallerType.LargestFiniteVal * SmallerType.LargestFiniteVal) * 2 <=
+ // LargerType.LargestFiniteVal.
+ // In terms of exponent it gives this formula:
+ // (SmallerType.LargestFiniteVal * SmallerType.LargestFiniteVal
+ // doubles the exponent of SmallerType.LargestFiniteVal)
+ if (llvm::APFloat::semanticsMaxExponent(elementTypeSemantics) * 2 + 1 <=
+ llvm::APFloat::semanticsMaxExponent(higherElementTypeSemantics)) {
+ return higherElementType;
+ }
+
+ // The intermediate values can't be represented in the promoted type
+ // without overflowing.
+ return {};
+}
+
+static mlir::Value
+lowerComplexDiv(mlir::ModuleOp mlirModule, CIRBaseBuilderTy &builder,
+ mlir::Location loc, cir::ComplexDivOp op, mlir::Value lhsReal,
+ mlir::Value lhsImag, mlir::Value rhsReal, mlir::Value rhsImag,
+ mlir::MLIRContext &mlirCx, clang::ASTContext &cc) {
+ cir::ComplexType complexTy = op.getType();
+ if (mlir::isa<cir::FPTypeInterface>(complexTy.getElementType())) {
+ cir::ComplexRangeKind range = op.getRange();
+ if (range == cir::ComplexRangeKind::Improved)
+ return buildRangeReductionComplexDiv(builder, loc, lhsReal, lhsImag,
+ rhsReal, rhsImag);
+
+ if (range == cir::ComplexRangeKind::Full)
+ return buildComplexBinOpLibCall(mlirModule, builder,
+ &getComplexDivLibCallName, loc, complexTy,
+ lhsReal, lhsImag, rhsReal, rhsImag);
+
+ if (range == cir::ComplexRangeKind::Promoted) {
+ mlir::Type originalElementType = complexTy.getElementType();
+ mlir::Type higherPrecisionElementType =
+ higherPrecisionElementTypeForComplexArithmetic(mlirCx, cc, builder,
+ originalElementType);
+
+ if (!higherPrecisionElementType)
+ return buildRangeReductionComplexDiv(builder, loc, lhsReal, lhsImag,
+ rhsReal, rhsImag);
+
+ cir::CastKind floatingCastKind = cir::CastKind::floating;
+ lhsReal = builder.createCast(floatingCastKind, lhsReal,
+ higherPrecisionElementType);
+ lhsImag = builder.createCast(floatingCastKind, lhsImag,
+ higherPrecisionElementType);
+ rhsReal = builder.createCast(floatingCastKind, rhsReal,
+ higherPrecisionElementType);
+ rhsImag = builder.createCast(floatingCastKind, rhsImag,
+ higherPrecisionElementType);
+
+ mlir::Value algebraicResult = buildAlgebraicComplexDiv(
+ builder, loc, lhsReal, lhsImag, rhsReal, rhsImag);
+
+ mlir::Value resultReal = builder.createComplexReal(loc, algebraicResult);
+ mlir::Value resultImag = builder.createComplexImag(loc, algebraicResult);
+
+ mlir::Value finalReal =
+ builder.createCast(floatingCastKind, resultReal, originalElementType);
+ mlir::Value finalImag =
+ builder.createCast(floatingCastKind, resultImag, originalElementType);
+ return builder.createComplexCreate(loc, finalReal, finalImag);
+ }
+ }
+
+ return buildAlgebraicComplexDiv(builder, loc, lhsReal, lhsImag, rhsReal,
+ rhsImag);
+}
+
+void ComplexLoweringPass::lowerComplexDivOp(cir::ComplexDivOp op) {
+ cir::CIRBaseBuilderTy builder(getContext());
+ builder.setInsertionPointAfter(op);
+ mlir::Location loc = op.getLoc();
+ mlir::TypedValue<cir::ComplexType> lhs = op.getLhs();
+ mlir::TypedValue<cir::ComplexType> rhs = op.getRhs();
+ mlir::Value lhsReal = builder.createComplexReal(loc, lhs);
+ mlir::Value lhsImag = builder.createComplexImag(loc, lhs);
+ mlir::Value rhsReal = builder.createComplexReal(loc, rhs);
+ mlir::Value rhsImag = builder.createComplexImag(loc, rhs);
+
+ mlir::Value loweredResult =
+ lowerComplexDiv(mlirModule, builder, loc, op, lhsReal, lhsImag, rhsReal,
+ rhsImag, getContext(), *astCtx);
+ op.replaceAllUsesWith(loweredResult);
+ op.erase();
+}
+
+static llvm::StringRef
+getComplexMulLibCallName(llvm::APFloat::Semantics semantics) {
+ switch (semantics) {
+ case llvm::APFloat::S_IEEEhalf:
+ return "__mulhc3";
+ case llvm::APFloat::S_IEEEsingle:
+ return "__mulsc3";
+ case llvm::APFloat::S_IEEEdouble:
+ return "__muldc3";
+ case llvm::APFloat::S_PPCDoubleDouble:
+ return "__multc3";
+ case llvm::APFloat::S_x87DoubleExtended:
+ return "__mulxc3";
+ case llvm::APFloat::S_IEEEquad:
+ return "__multc3";
+ default:
+ llvm_unreachable("unsupported floating point type");
+ }
+}
+
+static mlir::Value lowerComplexMul(mlir::ModuleOp mlirModule,
+ CIRBaseBuilderTy &builder,
+ mlir::Location loc, cir::ComplexMulOp op,
+ mlir::Value lhsReal, mlir::Value lhsImag,
+ mlir::Value rhsReal, mlir::Value rhsImag) {
+ // (a+bi) * (c+di) = (ac-bd) + (ad+bc)i
+ bool isFP = cir::isFPOrVectorOfFPType(lhsReal.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);
+ };
+
+ mlir::Value resultRealLhs = mul(loc, lhsReal, rhsReal); // ac
+ mlir::Value resultRealRhs = mul(loc, lhsImag, rhsImag); // bd
+ mlir::Value resultImagLhs = mul(loc, lhsReal, rhsImag); // ad
+ mlir::Value resultImagRhs = mul(loc, lhsImag, rhsReal); // bc
+ mlir::Value resultReal = sub(loc, resultRealLhs, resultRealRhs);
+ mlir::Value resultImag = add(loc, resultImagLhs, resultImagRhs);
+ mlir::Value algebraicResult =
+ builder.createComplexCreate(loc, resultReal, resultImag);
+
+ cir::ComplexType complexTy = op.getType();
+ cir::ComplexRangeKind rangeKind = op.getRange();
+ if (mlir::isa<cir::IntType>(complexTy.getElementType()) ||
+ rangeKind == cir::ComplexRangeKind::Basic ||
+ rangeKind == cir::ComplexRangeKind::Improved ||
+ rangeKind == cir::ComplexRangeKind::Promoted)
+ return algebraicResult;
+
+ assert(!cir::MissingFeatures::fastMathFlags());
+
+ // Check whether the real part and the imaginary part of the result are both
+ // NaN. If so, emit a library call to compute the multiplication instead.
+ // We check a value against NaN by comparing the value against itself.
+ mlir::Value resultRealIsNaN = builder.createIsNaN(loc, resultReal);
+ mlir::Value resultImagIsNaN = builder.createIsNaN(loc, resultImag);
+ mlir::Value resultRealAndImagAreNaN =
+ builder.createLogicalAnd(loc, resultRealIsNaN, resultImagIsNaN);
+
+ return cir::TernaryOp::create(
+ builder, loc, resultRealAndImagAreNaN,
+ [&](mlir::OpBuilder &, mlir::Location) {
+ mlir::Value libCallResult = buildComplexBinOpLibCall(
+ mlirModule, builder, &getComplexMulLibCallName, loc,
+ complexTy, lhsReal, lhsImag, rhsReal, rhsImag);
+ builder.createYield(loc, libCallResult);
+ },
+ [&](mlir::OpBuilder &, mlir::Location) {
+ builder.createYield(loc, algebraicResult);
+ })
+ .getResult();
+}
+
+void ComplexLoweringPass::lowerComplexMulOp(cir::ComplexMulOp op) {
+ cir::CIRBaseBuilderTy builder(getContext());
+ builder.setInsertionPointAfter(op);
+ mlir::Location loc = op.getLoc();
+ mlir::TypedValue<cir::ComplexType> lhs = op.getLhs();
+ mlir::TypedValue<cir::ComplexType> rhs = op.getRhs();
+ mlir::Value lhsReal = builder.createComplexReal(loc, lhs);
+ mlir::Value lhsImag = builder.createComplexImag(loc, lhs);
+ mlir::Value rhsReal = builder.createComplexReal(loc, rhs);
+ mlir::Value rhsImag = builder.createComplexImag(loc, rhs);
+ mlir::Value loweredResult = lowerComplexMul(
+ mlirModule, builder, loc, op, lhsReal, lhsImag, rhsReal, rhsImag);
+ op.replaceAllUsesWith(loweredResult);
+ op.erase();
+}
+
+void ComplexLoweringPass::runOnOperation() {
+ assert(astCtx && "complex lowering requires an ASTContext");
+ mlirModule = mlir::cast<mlir::ModuleOp>(getOperation());
+
+ llvm::SmallVector<mlir::Operation *> opsToTransform;
+ mlirModule->walk([&](mlir::Operation *op) {
+ if (mlir::isa<cir::ComplexMulOp, cir::ComplexDivOp>(op))
+ opsToTransform.push_back(op);
+ });
+
+ for (mlir::Operation *o : opsToTransform) {
+ if (auto complexDiv = mlir::dyn_cast<cir::ComplexDivOp>(o))
+ lowerComplexDivOp(complexDiv);
+ else
+ lowerComplexMulOp(mlir::cast<cir::ComplexMulOp>(o));
+ }
+}
+
+std::unique_ptr<Pass> mlir::createComplexLoweringPass() {
+ return std::make_unique<ComplexLoweringPass>();
+}
+
+std::unique_ptr<Pass>
+mlir::createComplexLoweringPass(clang::ASTContext *astCtx) {
+ auto pass = std::make_unique<ComplexLoweringPass>();
+ pass->setASTContext(astCtx);
+ return std::move(pass);
+}
diff --git a/clang/lib/CIR/Dialect/Transforms/LoweringPrepare.cpp b/clang/lib/CIR/Dialect/Transforms/LoweringPrepare.cpp
index a71aba6c1460f..5bac3fc5703ad 100644
--- a/clang/lib/CIR/Dialect/Transforms/LoweringPrepare.cpp
+++ b/clang/lib/CIR/Dialect/Transforms/LoweringPrepare.cpp
@@ -27,6 +27,7 @@
#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/Interfaces/ASTAttrInterfaces.h"
#include "clang/CIR/MissingFeatures.h"
#include "llvm/ADT/StringRef.h"
@@ -44,7 +45,6 @@ using namespace mlir;
using namespace cir;
namespace mlir {
-#define GEN_PASS_DEF_COMPLEXLOWERING
#define GEN_PASS_DEF_LOWERINGPREPARE
#include "clang/CIR/Dialect/Passes.h.inc"
} // namespace mlir
@@ -149,11 +149,6 @@ struct LoweringPreparePass
/// Materialize global ctor/dtor list
void buildGlobalCtorDtorList();
- cir::FuncOp buildRuntimeFunction(
- mlir::OpBuilder &builder, llvm::StringRef name, mlir::Location loc,
- cir::FuncType type,
- cir::GlobalLinkageKind linkage = cir::GlobalLinkageKind::ExternalLinkage);
-
cir::GlobalOp getOrCreateRuntimeVariable(
mlir::OpBuilder &builder, llvm::StringRef name, mlir::Location loc,
mlir::Type type,
@@ -294,7 +289,8 @@ struct LoweringPreparePass
/// `symbolTables.getSymbolTable(mlirModule).insert(...)` (as
/// `getOrCreateConstAggregateGlobal` does), or (b) creates a symbol
/// that is never resolved through the cache later. Today
- /// `buildRuntimeFunction` and `getOrCreateRuntimeVariable` fall in the
+ /// `cir::buildRuntimeFunction` (in CIRTransformUtils.cpp) and
+ /// `getOrCreateRuntimeVariable` fall in the
/// (b) bucket: their callers either use a separate map
/// (`cudaKernelMap`, `staticLocalDeclGuardMap`, `dynamicInitializers`)
/// or the static `mlir::SymbolTable::lookupNearestSymbolFrom`, never
@@ -369,8 +365,8 @@ struct LoweringPreparePass
? llvm::StringLiteral("_tlv_atexit")
: llvm::StringLiteral("__cxa_thread_atexit");
- cir::FuncOp fnAtExit = buildRuntimeFunction(builder, nameAtExit,
- global.getLoc(), fnAtExitType);
+ cir::FuncOp fnAtExit = cir::buildRuntimeFunction(
+ builder, mlirModule, nameAtExit, global.getLoc(), fnAtExitType);
// Replace the dtor (or helper) call with a call to
// __cxa_atexit(&dtor, &var, &__dso_handle)
@@ -504,26 +500,6 @@ struct LoweringPreparePass
void setASTContext(clang::ASTContext *c) { astCtx = c; }
};
-/// Expand `cir.complex.mul` and `cir.complex.div`. See the pass description
-/// in Passes.td for why this cannot run with the rest of LoweringPrepare.
-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
cir::GlobalOp LoweringPreparePass::getOrCreateRuntimeVariable(
@@ -544,33 +520,6 @@ cir::GlobalOp LoweringPreparePass::getOrCreateRuntimeVariable(
return g;
}
-/// Declare `name` in `mlirModule` if it is not already declared there, and
-/// return the declaration. Free-standing so that ComplexLoweringPass can
-/// reach it without a LoweringPreparePass instance.
-static cir::FuncOp buildRuntimeFunction(
- mlir::OpBuilder &builder, mlir::ModuleOp mlirModule, llvm::StringRef name,
- mlir::Location loc, cir::FuncType type,
- cir::GlobalLinkageKind linkage = cir::GlobalLinkageKind::ExternalLinkage) {
- cir::FuncOp f = dyn_cast_or_null<FuncOp>(SymbolTable::lookupNearestSymbolFrom(
- mlirModule, StringAttr::get(mlirModule->getContext(), name)));
- if (!f) {
- f = cir::FuncOp::create(builder, loc, name, type);
- f.setLinkageAttr(
- cir::GlobalLinkageKindAttr::get(builder.getContext(), linkage));
- mlir::SymbolTable::setSymbolVisibility(
- f, mlir::SymbolTable::Visibility::Private);
-
- assert(!cir::MissingFeatures::opFuncExtraAttrs());
- }
- return f;
-}
-
-cir::FuncOp LoweringPreparePass::buildRuntimeFunction(
- mlir::OpBuilder &builder, llvm::StringRef name, mlir::Location loc,
- cir::FuncType type, cir::GlobalLinkageKind linkage) {
- return ::buildRuntimeFunction(builder, mlirModule, name, loc, type, linkage);
-}
-
static mlir::Value lowerScalarToComplexCast(mlir::MLIRContext &ctx,
cir::CastOp op) {
cir::CIRBaseBuilderTy builder(ctx);
@@ -656,426 +605,6 @@ void LoweringPreparePass::lowerCastOp(cir::CastOp op) {
}
}
-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 =
- 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);
-
- if (mlir::isa<cir::DoubleType>(type))
- return cir::LongDoubleType::get(&context, type);
-
- return type;
- };
-
- auto getFloatTypeSemantics =
- [&cc](mlir::Type type) -> const llvm::fltSemantics & {
- const clang::TargetInfo &info = cc.getTargetInfo();
- if (mlir::isa<cir::FP16Type>(type))
- return info.getHalfFormat();
-
- if (mlir::isa<cir::BF16Type>(type))
- return info.getBFloat16Format();
-
- if (mlir::isa<cir::SingleType>(type))
- return info.getFloatFormat();
-
- if (mlir::isa<cir::DoubleType>(type))
- return info.getDoubleFormat();
-
- if (mlir::isa<cir::LongDoubleType>(type)) {
- if (cc.getLangOpts().OpenMP && cc.getLangOpts().OpenMPIsTargetDevice)
- llvm_unreachable("NYI Float type semantics with OpenMP");
- return info.getLongDoubleFormat();
- }
-
- if (mlir::isa<cir::FP128Type>(type)) {
- if (cc.getLangOpts().OpenMP && cc.getLangOpts().OpenMPIsTargetDevice)
- llvm_unreachable("NYI Float type semantics with OpenMP");
- return info.getFloat128Format();
- }
-
- llvm_unreachable("Unsupported float type semantics");
- };
-
- const mlir::Type higherElementType = getHigherPrecisionFPType(elementType);
- const llvm::fltSemantics &elementTypeSemantics =
- getFloatTypeSemantics(elementType);
- const llvm::fltSemantics &higherElementTypeSemantics =
- getFloatTypeSemantics(higherElementType);
-
- // Check that the promoted type can handle the intermediate values without
- // overflowing. This can be interpreted as:
- // (SmallerType.LargestFiniteVal * SmallerType.LargestFiniteVal) * 2 <=
- // LargerType.LargestFiniteVal.
- // In terms of exponent it gives this formula:
- // (SmallerType.LargestFiniteVal * SmallerType.LargestFiniteVal
- // doubles the exponent of SmallerType.LargestFiniteVal)
- if (llvm::APFloat::semanticsMaxExponent(elementTypeSemantics) * 2 + 1 <=
- llvm::APFloat::semanticsMaxExponent(higherElementTypeSemantics)) {
- return higherElementType;
- }
-
- // The intermediate values can't be represented in the promoted type
- // without overflowing.
- return {};
-}
-
-static mlir::Value
-lowerComplexDiv(mlir::ModuleOp mlirModule, CIRBaseBuilderTy &builder,
- mlir::Location loc, cir::ComplexDivOp op, mlir::Value lhsReal,
- mlir::Value lhsImag, mlir::Value rhsReal, mlir::Value rhsImag,
- mlir::MLIRContext &mlirCx, clang::ASTContext &cc) {
- cir::ComplexType complexTy = op.getType();
- if (mlir::isa<cir::FPTypeInterface>(complexTy.getElementType())) {
- cir::ComplexRangeKind range = op.getRange();
- if (range == cir::ComplexRangeKind::Improved)
- return buildRangeReductionComplexDiv(builder, loc, lhsReal, lhsImag,
- rhsReal, rhsImag);
-
- if (range == cir::ComplexRangeKind::Full)
- return buildComplexBinOpLibCall(mlirModule, builder,
- &getComplexDivLibCallName, loc, complexTy,
- lhsReal, lhsImag, rhsReal, rhsImag);
-
- if (range == cir::ComplexRangeKind::Promoted) {
- mlir::Type originalElementType = complexTy.getElementType();
- mlir::Type higherPrecisionElementType =
- higherPrecisionElementTypeForComplexArithmetic(mlirCx, cc, builder,
- originalElementType);
-
- if (!higherPrecisionElementType)
- return buildRangeReductionComplexDiv(builder, loc, lhsReal, lhsImag,
- rhsReal, rhsImag);
-
- cir::CastKind floatingCastKind = cir::CastKind::floating;
- lhsReal = builder.createCast(floatingCastKind, lhsReal,
- higherPrecisionElementType);
- lhsImag = builder.createCast(floatingCastKind, lhsImag,
- higherPrecisionElementType);
- rhsReal = builder.createCast(floatingCastKind, rhsReal,
- higherPrecisionElementType);
- rhsImag = builder.createCast(floatingCastKind, rhsImag,
- higherPrecisionElementType);
-
- mlir::Value algebraicResult = buildAlgebraicComplexDiv(
- builder, loc, lhsReal, lhsImag, rhsReal, rhsImag);
-
- mlir::Value resultReal = builder.createComplexReal(loc, algebraicResult);
- mlir::Value resultImag = builder.createComplexImag(loc, algebraicResult);
-
- mlir::Value finalReal =
- builder.createCast(floatingCastKind, resultReal, originalElementType);
- mlir::Value finalImag =
- builder.createCast(floatingCastKind, resultImag, originalElementType);
- return builder.createComplexCreate(loc, finalReal, finalImag);
- }
- }
-
- return buildAlgebraicComplexDiv(builder, loc, lhsReal, lhsImag, rhsReal,
- rhsImag);
-}
-
-void ComplexLoweringPass::lowerComplexDivOp(cir::ComplexDivOp op) {
- cir::CIRBaseBuilderTy builder(getContext());
- builder.setInsertionPointAfter(op);
- mlir::Location loc = op.getLoc();
- mlir::TypedValue<cir::ComplexType> lhs = op.getLhs();
- mlir::TypedValue<cir::ComplexType> rhs = op.getRhs();
- mlir::Value lhsReal = builder.createComplexReal(loc, lhs);
- mlir::Value lhsImag = builder.createComplexImag(loc, lhs);
- mlir::Value rhsReal = builder.createComplexReal(loc, rhs);
- mlir::Value rhsImag = builder.createComplexImag(loc, rhs);
-
- mlir::Value loweredResult =
- lowerComplexDiv(mlirModule, builder, loc, op, lhsReal, lhsImag, rhsReal,
- rhsImag, getContext(), *astCtx);
- op.replaceAllUsesWith(loweredResult);
- op.erase();
-}
-
-static llvm::StringRef
-getComplexMulLibCallName(llvm::APFloat::Semantics semantics) {
- switch (semantics) {
- case llvm::APFloat::S_IEEEhalf:
- return "__mulhc3";
- case llvm::APFloat::S_IEEEsingle:
- return "__mulsc3";
- case llvm::APFloat::S_IEEEdouble:
- return "__muldc3";
- case llvm::APFloat::S_PPCDoubleDouble:
- return "__multc3";
- case llvm::APFloat::S_x87DoubleExtended:
- return "__mulxc3";
- case llvm::APFloat::S_IEEEquad:
- return "__multc3";
- default:
- llvm_unreachable("unsupported floating point type");
- }
-}
-
-static mlir::Value lowerComplexMul(mlir::ModuleOp mlirModule,
- CIRBaseBuilderTy &builder,
- mlir::Location loc, cir::ComplexMulOp op,
- mlir::Value lhsReal, mlir::Value lhsImag,
- mlir::Value rhsReal, mlir::Value rhsImag) {
- // (a+bi) * (c+di) = (ac-bd) + (ad+bc)i
- bool isFP = cir::isFPOrVectorOfFPType(lhsReal.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);
- };
-
- mlir::Value resultRealLhs = mul(loc, lhsReal, rhsReal); // ac
- mlir::Value resultRealRhs = mul(loc, lhsImag, rhsImag); // bd
- mlir::Value resultImagLhs = mul(loc, lhsReal, rhsImag); // ad
- mlir::Value resultImagRhs = mul(loc, lhsImag, rhsReal); // bc
- mlir::Value resultReal = sub(loc, resultRealLhs, resultRealRhs);
- mlir::Value resultImag = add(loc, resultImagLhs, resultImagRhs);
- mlir::Value algebraicResult =
- builder.createComplexCreate(loc, resultReal, resultImag);
-
- cir::ComplexType complexTy = op.getType();
- cir::ComplexRangeKind rangeKind = op.getRange();
- if (mlir::isa<cir::IntType>(complexTy.getElementType()) ||
- rangeKind == cir::ComplexRangeKind::Basic ||
- rangeKind == cir::ComplexRangeKind::Improved ||
- rangeKind == cir::ComplexRangeKind::Promoted)
- return algebraicResult;
-
- assert(!cir::MissingFeatures::fastMathFlags());
-
- // Check whether the real part and the imaginary part of the result are both
- // NaN. If so, emit a library call to compute the multiplication instead.
- // We check a value against NaN by comparing the value against itself.
- mlir::Value resultRealIsNaN = builder.createIsNaN(loc, resultReal);
- mlir::Value resultImagIsNaN = builder.createIsNaN(loc, resultImag);
- mlir::Value resultRealAndImagAreNaN =
- builder.createLogicalAnd(loc, resultRealIsNaN, resultImagIsNaN);
-
- return cir::TernaryOp::create(
- builder, loc, resultRealAndImagAreNaN,
- [&](mlir::OpBuilder &, mlir::Location) {
- mlir::Value libCallResult = buildComplexBinOpLibCall(
- mlirModule, builder, &getComplexMulLibCallName, loc,
- complexTy, lhsReal, lhsImag, rhsReal, rhsImag);
- builder.createYield(loc, libCallResult);
- },
- [&](mlir::OpBuilder &, mlir::Location) {
- builder.createYield(loc, algebraicResult);
- })
- .getResult();
-}
-
-void ComplexLoweringPass::lowerComplexMulOp(cir::ComplexMulOp op) {
- cir::CIRBaseBuilderTy builder(getContext());
- builder.setInsertionPointAfter(op);
- mlir::Location loc = op.getLoc();
- mlir::TypedValue<cir::ComplexType> lhs = op.getLhs();
- mlir::TypedValue<cir::ComplexType> rhs = op.getRhs();
- mlir::Value lhsReal = builder.createComplexReal(loc, lhs);
- mlir::Value lhsImag = builder.createComplexImag(loc, lhs);
- mlir::Value rhsReal = builder.createComplexReal(loc, rhs);
- mlir::Value rhsImag = builder.createComplexImag(loc, rhs);
- mlir::Value loweredResult = lowerComplexMul(
- mlirModule, builder, loc, op, lhsReal, lhsImag, rhsReal, rhsImag);
- op.replaceAllUsesWith(loweredResult);
- op.erase();
-}
-
-void ComplexLoweringPass::runOnOperation() {
- mlirModule = cast<mlir::ModuleOp>(getOperation());
-
- llvm::SmallVector<mlir::Operation *> opsToTransform;
- mlirModule->walk([&](mlir::Operation *op) {
- if (mlir::isa<cir::ComplexMulOp, cir::ComplexDivOp>(op))
- opsToTransform.push_back(op);
- });
-
- for (mlir::Operation *o : opsToTransform)
- if (auto complexDiv = mlir::dyn_cast<cir::ComplexDivOp>(o))
- lowerComplexDivOp(complexDiv);
- else
- lowerComplexMulOp(mlir::cast<cir::ComplexMulOp>(o));
-}
-
void LoweringPreparePass::lowerComplexConjOp(cir::ComplexConjOp op) {
mlir::Location loc = op.getLoc();
CIRBaseBuilderTy builder(getContext());
@@ -1144,9 +673,9 @@ cir::FuncOp LoweringPreparePass::getOrCreateDtorFunc(CIRBaseBuilderTy &builder,
// Create the helper function.
auto fnType = cir::FuncType::get({voidPtrTy}, voidTy);
- cir::FuncOp dtorFunc =
- buildRuntimeFunction(builder, fnName, op.getLoc(), fnType,
- cir::GlobalLinkageKind::InternalLinkage);
+ cir::FuncOp dtorFunc = cir::buildRuntimeFunction(
+ builder, mlirModule, fnName, op.getLoc(), fnType,
+ cir::GlobalLinkageKind::InternalLinkage);
SmallVector<mlir::NamedAttribute> paramAttrs;
paramAttrs.push_back(
@@ -1216,8 +745,9 @@ LoweringPreparePass::buildCXXGlobalVarDeclInitFunc(cir::GlobalOp op) {
builder.setInsertionPointAfter(op);
cir::VoidType voidTy = builder.getVoidTy();
auto fnType = cir::FuncType::get({}, voidTy);
- FuncOp f = buildRuntimeFunction(builder, fnName, op.getLoc(), fnType,
- cir::GlobalLinkageKind::InternalLinkage);
+ FuncOp f = cir::buildRuntimeFunction(builder, mlirModule, fnName, op.getLoc(),
+ fnType,
+ cir::GlobalLinkageKind::InternalLinkage);
// Forward the constrained floating-point marker recorded on the global by
// CodeGen onto the generated initializer function. The marker on the global
@@ -1301,7 +831,8 @@ LoweringPreparePass::getGuardAcquireFn(cir::PointerType guardPtrTy) {
mlir::Location loc = mlirModule.getLoc();
cir::IntType intTy = cir::IntType::get(&getContext(), 32, /*isSigned=*/true);
auto fnType = cir::FuncType::get({guardPtrTy}, intTy);
- return buildRuntimeFunction(builder, "__cxa_guard_acquire", loc, fnType);
+ return cir::buildRuntimeFunction(builder, mlirModule, "__cxa_guard_acquire",
+ loc, fnType);
}
cir::FuncOp
@@ -1313,7 +844,8 @@ LoweringPreparePass::getGuardReleaseFn(cir::PointerType guardPtrTy) {
mlir::Location loc = mlirModule.getLoc();
cir::VoidType voidTy = cir::VoidType::get(&getContext());
auto fnType = cir::FuncType::get({guardPtrTy}, voidTy);
- return buildRuntimeFunction(builder, "__cxa_guard_release", loc, fnType);
+ return cir::buildRuntimeFunction(builder, mlirModule, "__cxa_guard_release",
+ loc, fnType);
}
cir::FuncOp LoweringPreparePass::getTlsInitFn() {
@@ -1323,8 +855,9 @@ cir::FuncOp LoweringPreparePass::getTlsInitFn() {
builder.setInsertionPointToStart(mlirModule.getBody());
mlir::Location loc = mlirModule.getLoc();
auto fnType = builder.getVoidFnTy();
- return buildRuntimeFunction(builder, "__tls_init", loc, fnType,
- cir::GlobalLinkageKind::InternalLinkage);
+ return cir::buildRuntimeFunction(builder, mlirModule, "__tls_init", loc,
+ fnType,
+ cir::GlobalLinkageKind::InternalLinkage);
}
cir::GlobalOp LoweringPreparePass::createGuardGlobalOp(
@@ -1935,8 +1468,8 @@ void LoweringPreparePass::buildCXXGlobalInitFunc() {
CIRBaseBuilderTy builder(getContext());
builder.setInsertionPointToEnd(&mlirModule.getBodyRegion().back());
auto fnType = cir::FuncType::get({}, builder.getVoidTy());
- cir::FuncOp f = buildRuntimeFunction(builder, fnName, mlirModule.getLoc(),
- fnType, linkage);
+ cir::FuncOp f = cir::buildRuntimeFunction(
+ builder, mlirModule, fnName, mlirModule.getLoc(), fnType, linkage);
builder.setInsertionPointToStart(f.addEntryBlock());
for (cir::FuncOp &f : dynamicInitializers)
builder.createCallOp(f.getLoc(), f, {});
@@ -2527,12 +2060,12 @@ void LoweringPreparePass::buildCUDAModuleCtor() {
std::string regFuncName =
addUnderscoredPrefix(cudaPrefix, "RegisterFatBinary");
FuncType regFuncType = FuncType::get({voidPtrTy}, voidPtrPtrTy);
- cir::FuncOp regFunc =
- buildRuntimeFunction(builder, regFuncName, loc, regFuncType);
+ cir::FuncOp regFunc = cir::buildRuntimeFunction(
+ builder, mlirModule, regFuncName, loc, regFuncType);
std::string moduleCtorName = addUnderscoredPrefix(cudaPrefix, "_module_ctor");
- cir::FuncOp moduleCtor = buildRuntimeFunction(
- builder, moduleCtorName, loc, FuncType::get({}, voidTy),
+ cir::FuncOp moduleCtor = cir::buildRuntimeFunction(
+ builder, mlirModule, moduleCtorName, loc, FuncType::get({}, voidTy),
GlobalLinkageKind::InternalLinkage);
globalCtorList.emplace_back(moduleCtorName,
@@ -2588,8 +2121,8 @@ void LoweringPreparePass::buildCUDAModuleCtor() {
if (std::optional<FuncOp> dtor = buildHIPModuleDtor()) {
cir::CIRBaseBuilderTy globalBuilder(getContext());
globalBuilder.setInsertionPointToStart(mlirModule.getBody());
- FuncOp atexit = buildRuntimeFunction(
- globalBuilder, "atexit", loc,
+ FuncOp atexit = cir::buildRuntimeFunction(
+ globalBuilder, mlirModule, "atexit", loc,
FuncType::get(PointerType::get(dtor->getFunctionType()), intTy));
mlir::Value dtorFunc = GetGlobalOp::create(
builder, loc, PointerType::get(dtor->getFunctionType()),
@@ -2629,9 +2162,9 @@ void LoweringPreparePass::buildCUDAModuleCtor() {
clang::CudaFeature::CUDA_USES_FATBIN_REGISTER_END)) {
cir::CIRBaseBuilderTy globalBuilder(getContext());
globalBuilder.setInsertionPointToStart(mlirModule.getBody());
- FuncOp endFunc =
- buildRuntimeFunction(globalBuilder, "__cudaRegisterFatBinaryEnd", loc,
- FuncType::get({voidPtrPtrTy}, voidTy));
+ FuncOp endFunc = cir::buildRuntimeFunction(
+ globalBuilder, mlirModule, "__cudaRegisterFatBinaryEnd", loc,
+ FuncType::get({voidPtrPtrTy}, voidTy));
builder.createCallOp(loc, endFunc, gpuBinaryHandle);
}
} else
@@ -2645,8 +2178,8 @@ void LoweringPreparePass::buildCUDAModuleCtor() {
// extern "C" int atexit(void (*f)(void));
cir::CIRBaseBuilderTy globalBuilder(getContext());
globalBuilder.setInsertionPointToStart(mlirModule.getBody());
- FuncOp atexit = buildRuntimeFunction(
- globalBuilder, "atexit", loc,
+ FuncOp atexit = cir::buildRuntimeFunction(
+ globalBuilder, mlirModule, "atexit", loc,
FuncType::get(PointerType::get(dtor->getFunctionType()), intTy));
mlir::Value dtorFunc = GetGlobalOp::create(
builder, loc, PointerType::get(dtor->getFunctionType()),
@@ -2673,8 +2206,9 @@ std::optional<FuncOp> LoweringPreparePass::buildCUDAModuleDtor() {
// define: void __cudaUnregisterFatBinary(void ** handle);
std::string unregisterFuncName =
addUnderscoredPrefix(prefix, "UnregisterFatBinary");
- FuncOp unregisterFunc = buildRuntimeFunction(
- builder, unregisterFuncName, loc, FuncType::get({voidPtrPtrTy}, voidTy));
+ FuncOp unregisterFunc =
+ cir::buildRuntimeFunction(builder, mlirModule, unregisterFuncName, loc,
+ FuncType::get({voidPtrPtrTy}, voidTy));
// void __cuda_module_dtor();
// Despite the name, OG doesn't treat it as a destructor, so it shouldn't be
@@ -2682,9 +2216,9 @@ std::optional<FuncOp> LoweringPreparePass::buildCUDAModuleDtor() {
// double free above CUDA 9.2. The way to use it is to manually call
// atexit() at end of module ctor.
std::string dtorName = addUnderscoredPrefix(prefix, "_module_dtor");
- FuncOp dtor =
- buildRuntimeFunction(builder, dtorName, loc, FuncType::get({}, voidTy),
- GlobalLinkageKind::InternalLinkage);
+ FuncOp dtor = cir::buildRuntimeFunction(builder, mlirModule, dtorName, loc,
+ FuncType::get({}, voidTy),
+ GlobalLinkageKind::InternalLinkage);
builder.setInsertionPointToStart(dtor.addEntryBlock());
@@ -2730,13 +2264,14 @@ std::optional<FuncOp> LoweringPreparePass::buildHIPModuleDtor() {
// void __hipUnregisterFatBinary(void ** handle);
std::string unregisterFuncName =
addUnderscoredPrefix(prefix, "UnregisterFatBinary");
- FuncOp unregisterFunc = buildRuntimeFunction(
- builder, unregisterFuncName, loc, FuncType::get({voidPtrPtrTy}, voidTy));
+ FuncOp unregisterFunc =
+ cir::buildRuntimeFunction(builder, mlirModule, unregisterFuncName, loc,
+ FuncType::get({voidPtrPtrTy}, voidTy));
std::string dtorName = addUnderscoredPrefix(prefix, "_module_dtor");
- FuncOp dtor =
- buildRuntimeFunction(builder, dtorName, loc, FuncType::get({}, voidTy),
- GlobalLinkageKind::InternalLinkage);
+ FuncOp dtor = cir::buildRuntimeFunction(builder, mlirModule, dtorName, loc,
+ FuncType::get({}, voidTy),
+ GlobalLinkageKind::InternalLinkage);
std::string gpubinName = addUnderscoredPrefix(prefix, "_gpubin_handle");
GlobalOp gpuBinGlobal = cast<GlobalOp>(mlirModule.lookupSymbol(gpubinName));
@@ -2791,9 +2326,9 @@ std::optional<FuncOp> LoweringPreparePass::buildCUDARegisterGlobals() {
std::string regGlobalFuncName =
addUnderscoredPrefix(cudaPrefix, "_register_globals");
auto regGlobalFuncTy = FuncType::get({voidPtrPtrTy}, voidTy);
- FuncOp regGlobalFunc =
- buildRuntimeFunction(builder, regGlobalFuncName, loc, regGlobalFuncTy,
- /*linkage=*/GlobalLinkageKind::InternalLinkage);
+ FuncOp regGlobalFunc = cir::buildRuntimeFunction(
+ builder, mlirModule, regGlobalFuncName, loc, regGlobalFuncTy,
+ /*linkage=*/GlobalLinkageKind::InternalLinkage);
builder.setInsertionPointToStart(regGlobalFunc.addEntryBlock());
buildCUDARegisterGlobalFunctions(builder, regGlobalFunc);
@@ -2834,8 +2369,9 @@ void LoweringPreparePass::buildCUDARegisterGlobalFunctions(
// )
// OG doesn't care about the types at all. They're treated as void*.
- FuncOp cudaRegisterFunction = buildRuntimeFunction(
- globalBuilder, addUnderscoredPrefix(cudaPrefix, "RegisterFunction"), loc,
+ FuncOp cudaRegisterFunction = cir::buildRuntimeFunction(
+ globalBuilder, mlirModule,
+ addUnderscoredPrefix(cudaPrefix, "RegisterFunction"), loc,
FuncType::get({voidPtrPtrTy, voidPtrTy, voidPtrTy, voidPtrTy, intTy,
voidPtrTy, voidPtrTy, voidPtrTy, voidPtrTy, voidPtrTy},
intTy));
@@ -2916,8 +2452,9 @@ void LoweringPreparePass::buildCUDARegisterVars(cir::CIRBaseBuilderTy &builder,
// size_t size, int constant, int normalized);
// OG ignores parameter types, treating pointers as void*.
cir::VoidType voidTy = builder.getVoidTy();
- FuncOp cudaRegisterVar = buildRuntimeFunction(
- globalBuilder, addUnderscoredPrefix(cudaPrefix, "RegisterVar"), loc,
+ FuncOp cudaRegisterVar = cir::buildRuntimeFunction(
+ globalBuilder, mlirModule,
+ addUnderscoredPrefix(cudaPrefix, "RegisterVar"), loc,
FuncType::get({voidPtrPtrTy, voidPtrTy, voidPtrTy, voidPtrTy, intTy,
sizeTy, intTy, intTy},
voidTy));
@@ -3006,14 +2543,3 @@ mlir::createLoweringPreparePass(clang::ASTContext *astCtx) {
pass->setASTContext(astCtx);
return std::move(pass);
}
-
-std::unique_ptr<Pass> mlir::createComplexLoweringPass() {
- return std::make_unique<ComplexLoweringPass>();
-}
-
-std::unique_ptr<Pass>
-mlir::createComplexLoweringPass(clang::ASTContext *astCtx) {
- auto pass = std::make_unique<ComplexLoweringPass>();
- pass->setASTContext(astCtx);
- return std::move(pass);
-}
diff --git a/clang/lib/CIR/Lowering/CIRPasses.cpp b/clang/lib/CIR/Lowering/CIRPasses.cpp
index 1d1501e3719ee..dc1182fda6ae0 100644
--- a/clang/lib/CIR/Lowering/CIRPasses.cpp
+++ b/clang/lib/CIR/Lowering/CIRPasses.cpp
@@ -108,8 +108,9 @@ runCIRToCIRPasses(mlir::ModuleOp theModule, mlir::MLIRContext &mlirContext,
// Complex multiplication and division synthesize calls to runtime helpers
// such as __mulsc3 and __divsc3, so they must be expanded before
- // CallConvLowering classifies calls. The rest of the lowering-prepare work
- // stays after it.
+ // CallConvLowering classifies calls. A division in a global initializer is
+ // therefore coerced while still inside the global's initializer region,
+ // before LoweringPrepare turns that region into __cxx_global_var_init.
pm.addPass(mlir::createComplexLoweringPass(&astContext));
if (enableCallConvLowering) {
>From a23d8c438e143330c317a032e440e88689372f29 Mon Sep 17 00:00:00 2001
From: Adam Smith <adams at nvidia.com>
Date: Mon, 17 Aug 2026 11:45:48 -0700
Subject: [PATCH 3/3] [CIR] Tighten the complex libcall ABI tests
The global-initializer test checked that something was stored to the global
without tying it to the value loaded out of the coercion slot, so it would have
passed on a raw load of the dividend. Both layers now capture that chain end to
end. The test also covered only division, which is straight-line, so it gains
the multiply case, where the call sits inside the NaN-check region and the
coercion slot has to be hoisted to the initializer function's entry block.
The rest pins what was loose: `byval` on the long double arms, and the complex
range on the RUN lines, which the helper calls depend on and which was being
inherited from a default.
Assisted-by: Cursor / claude-opus-5
---
.../complex-libcall-abi-global-init.cpp | 54 ++++++++++++-------
clang/test/CIR/CodeGen/complex-libcall-abi.c | 22 ++++----
2 files changed, 46 insertions(+), 30 deletions(-)
diff --git a/clang/test/CIR/CodeGen/complex-libcall-abi-global-init.cpp b/clang/test/CIR/CodeGen/complex-libcall-abi-global-init.cpp
index 3ec35a24d3ced..5d0007981b9be 100644
--- a/clang/test/CIR/CodeGen/complex-libcall-abi-global-init.cpp
+++ b/clang/test/CIR/CodeGen/complex-libcall-abi-global-init.cpp
@@ -1,30 +1,48 @@
-// RUN: %clang_cc1 -triple x86_64-unknown-linux-gnu -std=c++20 -fclangir -emit-cir %s -o %t.cir
+// RUN: %clang_cc1 -triple x86_64-unknown-linux-gnu -std=c++20 -complex-range=full -fclangir -emit-cir %s -o %t.cir
// RUN: FileCheck --check-prefix=CIR --input-file=%t.cir %s
-// RUN: %clang_cc1 -triple x86_64-unknown-linux-gnu -std=c++20 -fclangir -emit-llvm %s -o %t-cir.ll
-// RUN: FileCheck --check-prefix=LLVM --input-file=%t-cir.ll %s
-// RUN: %clang_cc1 -triple x86_64-unknown-linux-gnu -std=c++20 -emit-llvm %s -o %t.ll
-// RUN: FileCheck --check-prefix=OGCG --input-file=%t.ll %s
+// RUN: %clang_cc1 -triple x86_64-unknown-linux-gnu -std=c++20 -complex-range=full -fclangir -emit-llvm %s -o %t-cir.ll
+// RUN: FileCheck --check-prefixes=LLVM,LLVMCIR --input-file=%t-cir.ll %s
+// RUN: %clang_cc1 -triple x86_64-unknown-linux-gnu -std=c++20 -complex-range=full -emit-llvm %s -o %t.ll
+// RUN: FileCheck --check-prefixes=LLVM,OGCG --input-file=%t.ll %s
extern float _Complex a;
extern float _Complex b;
-// A dynamic initializer at namespace scope is still inside the global's
-// initializer region when the helper call is coerced, so the coercion slot has
-// no enclosing function to be placed in. It has to land in the region that
-// later becomes the initializer function.
float _Complex g = a / b;
+// The coercion slot is allocated in the global initializer function.
+
// CIR-LABEL: cir.func {{.*}}@__cxx_global_var_init
-// CIR: %[[SLOT:.*]] = cir.alloca "coerce"{{.*}} : !cir.ptr<!cir.vector<2 x !cir.float>>
+// CIR: %[[SLOT:.*]] = cir.alloca "coerce" align(8) : !cir.ptr<!cir.vector<2 x !cir.float>>
+// CIR: %[[G:.*]] = cir.get_global @g : !cir.ptr<!cir.complex<!cir.float>>
// CIR: %[[COERCED:.*]] = cir.call @__divsc3({{.*}}) : (!cir.float, !cir.float, !cir.float, !cir.float) -> !cir.vector<2 x !cir.float>
// CIR: cir.store %[[COERCED]], %[[SLOT]] : !cir.vector<2 x !cir.float>, !cir.ptr<!cir.vector<2 x !cir.float>>
-// CIR: cir.store{{.*}} %{{.+}}, %{{.+}} : !cir.complex<!cir.float>, !cir.ptr<!cir.complex<!cir.float>>
+// CIR: %[[SLOT_PTR:.*]] = cir.cast bitcast %[[SLOT]] : !cir.ptr<!cir.vector<2 x !cir.float>> -> !cir.ptr<!cir.complex<!cir.float>>
+// CIR: %[[RESULT:.*]] = cir.load %[[SLOT_PTR]] : !cir.ptr<!cir.complex<!cir.float>>, !cir.complex<!cir.float>
+// CIR: cir.store{{.*}} %[[RESULT]], %[[G]] : !cir.complex<!cir.float>, !cir.ptr<!cir.complex<!cir.float>>
+
+// LLVM-LABEL: define internal void @__cxx_global_var_init()
+// LLVMCIR: %[[SLOT:.+]] = alloca <2 x float>, align 8
+// LLVMCIR: %[[COERCED:.*]] = call <2 x float> @__divsc3(float %{{.+}}, float %{{.+}}, float %{{.+}}, float %{{.+}})
+// LLVMCIR: store <2 x float> %[[COERCED]], ptr %[[SLOT]], align 8
+// LLVMCIR: %[[RESULT:.+]] = load { float, float }, ptr %[[SLOT]], align 4
+// LLVMCIR: store { float, float } %[[RESULT]], ptr @g, align 4
+
+// OGCG: %[[COERCED:.*]] = call noundef <2 x float> @__divsc3(float noundef %{{.+}}, float noundef %{{.+}}, float noundef %{{.+}}, float noundef %{{.+}})
+
+float _Complex h = a * b;
+
+// Multiply puts its call inside the NaN-check region, one level further in, so
+// the slot has to be hoisted to the initializer function's entry block.
-// LLVM-LABEL: @__cxx_global_var_init(
-// LLVM: %[[SLOT:.+]] = alloca <2 x float>, align 8
-// LLVM: %[[COERCED:.*]] = call <2 x float> @__divsc3(float %{{.+}}, float %{{.+}}, float %{{.+}}, float %{{.+}})
-// LLVM: store <2 x float> %[[COERCED]], ptr %[[SLOT]], align 8
-// LLVM: store { float, float } %{{.+}}, ptr @g, align 4
+// LLVM-LABEL: define internal void @__cxx_global_var_init.1()
+// LLVMCIR: %[[SLOT:.+]] = alloca <2 x float>, align 8
+// LLVMCIR: br i1 %{{.+}}, label %[[THEN:.+]], label %[[ELSE:.+]]
+// LLVMCIR: [[THEN]]:
+// LLVMCIR: %[[COERCED:.*]] = call <2 x float> @__mulsc3(float %{{.+}}, float %{{.+}}, float %{{.+}}, float %{{.+}})
+// LLVMCIR: store <2 x float> %[[COERCED]], ptr %[[SLOT]], align 8
+// LLVMCIR: %[[CALLED:.+]] = load { float, float }, ptr %[[SLOT]], align 4
+// LLVMCIR: %[[RESULT:.+]] = phi { float, float } [ %{{.+}}, %{{.+}} ], [ %[[CALLED]], %[[THEN]] ]
+// LLVMCIR: store { float, float } %[[RESULT]], ptr @h, align 4
-// OGCG-LABEL: @__cxx_global_var_init(
-// OGCG: call noundef <2 x float> @__divsc3(float noundef %{{.+}}, float noundef %{{.+}}, float noundef %{{.+}}, float noundef %{{.+}})
+// OGCG: call noundef <2 x float> @__mulsc3(float noundef %{{.+}}, float noundef %{{.+}}, float noundef %{{.+}}, float noundef %{{.+}})
diff --git a/clang/test/CIR/CodeGen/complex-libcall-abi.c b/clang/test/CIR/CodeGen/complex-libcall-abi.c
index efc459a30ad64..514837775f6db 100644
--- a/clang/test/CIR/CodeGen/complex-libcall-abi.c
+++ b/clang/test/CIR/CodeGen/complex-libcall-abi.c
@@ -1,15 +1,13 @@
-// RUN: %clang_cc1 -triple x86_64-unknown-linux-gnu -fclangir -emit-cir %s -o %t.cir
+// RUN: %clang_cc1 -triple x86_64-unknown-linux-gnu -complex-range=full -fclangir -emit-cir %s -o %t.cir
// RUN: FileCheck --check-prefix=CIR --input-file=%t.cir %s
-// RUN: %clang_cc1 -triple x86_64-unknown-linux-gnu -fclangir -emit-llvm %s -o %t-cir.ll
+// RUN: %clang_cc1 -triple x86_64-unknown-linux-gnu -complex-range=full -fclangir -emit-llvm %s -o %t-cir.ll
// RUN: FileCheck --check-prefixes=LLVM,LLVMCIR --input-file=%t-cir.ll %s
-// RUN: %clang_cc1 -triple x86_64-unknown-linux-gnu -emit-llvm %s -o %t.ll
+// RUN: %clang_cc1 -triple x86_64-unknown-linux-gnu -complex-range=full -emit-llvm %s -o %t.ll
// RUN: FileCheck --check-prefixes=LLVM,OGCG --input-file=%t.ll %s
float _Complex divf(float _Complex a, float _Complex b) { return a / b; }
-// A float pair fits one eightbyte and returns in a single SSE register, so the
-// helper's return coerces to a vector and comes back through memory.
-
+// A float pair is SSE-classified, so the helper return coerces to <2 x float>.
// CIR-LABEL: cir.func {{.*}}@divf
// CIR: %[[COERCED:.*]] = cir.call @__divsc3({{.*}}) : (!cir.float, !cir.float, !cir.float, !cir.float) -> !cir.vector<2 x !cir.float>
// CIR: cir.store %[[COERCED]], %[[SLOT:.*]] : !cir.vector<2 x !cir.float>, !cir.ptr<!cir.vector<2 x !cir.float>>
@@ -35,9 +33,7 @@ float _Complex mulf(float _Complex a, float _Complex b) { return a * b; }
double _Complex divd(double _Complex a, double _Complex b) { return a / b; }
-// A double pair spans two eightbytes and returns in two registers, so it stays
-// a two-field record and its lowered call is unchanged by the coercion.
-
+// A double pair spans two eightbytes, so the return stays a two-field record.
// CIR-LABEL: cir.func {{.*}}@divd
// CIR: cir.call @__divdc3({{.*}}) : (!cir.double, !cir.double, !cir.double, !cir.double) -> [[REC_D:!rec_anon_struct[0-9]*]]
@@ -56,12 +52,14 @@ long double _Complex divld(long double _Complex a, long double _Complex b) {
return a / b;
}
-// A long double pair is x87-classified and returned in memory, so it also
-// keeps a two-field record.
-
+// A long double pair is x87-classified, so the return stays a two-field record.
// CIR-LABEL: cir.func {{.*}}@divld
// CIR: cir.call @__divxc3({{.*}}) : (!cir.long_double<!cir.f80>, !cir.long_double<!cir.f80>, !cir.long_double<!cir.f80>, !cir.long_double<!cir.f80>) -> [[REC_LD:!rec_anon_struct[0-9]*]]
+// An x87 pair is passed indirectly, so the caller's own parameters are byval.
+// LLVMCIR: define dso_local { x86_fp80, x86_fp80 } @divld(ptr noalias noundef byval({ x86_fp80, x86_fp80 }) align 16 %{{.+}}, ptr noalias noundef byval({ x86_fp80, x86_fp80 }) align 16 %{{.+}})
+// OGCG: define dso_local { x86_fp80, x86_fp80 } @divld(ptr noundef byval({ x86_fp80, x86_fp80 }) align 16 %{{.+}}, ptr noundef byval({ x86_fp80, x86_fp80 }) align 16 %{{.+}})
+
// LLVMCIR: call { x86_fp80, x86_fp80 } @__divxc3(x86_fp80 %{{.+}}, x86_fp80 %{{.+}}, x86_fp80 %{{.+}}, x86_fp80 %{{.+}})
// OGCG: call { x86_fp80, x86_fp80 } @__divxc3(x86_fp80 noundef %{{.+}}, x86_fp80 noundef %{{.+}}, x86_fp80 noundef %{{.+}}, x86_fp80 noundef %{{.+}})
More information about the cfe-commits
mailing list