[clang] [CIR] Record -ffp-contract=fast as a per-op contract flag (PR #226334)
David Rivera via cfe-commits
cfe-commits at lists.llvm.org
Fri Sep 25 13:47:58 PDT 2026
https://github.com/RiverDave updated https://github.com/llvm/llvm-project/pull/226334
>From 6cd5e778ca421ab16e5a363131f20d26a2253971 Mon Sep 17 00:00:00 2001
From: David Rivera <davidriverg at gmail.com>
Date: Thu, 24 Sep 2026 20:54:26 -0400
Subject: [PATCH 1/5] [CIR] Record -ffp-contract=fast as a per-op contract flag
Classic CodeGen stamps contract on floating-point instructions so a later
Standard-fusion backend can still form an FMA. CIR only fused within a
statement via cir.fmuladd, which dropped FFMA on the CUDA device default.
Co-authored-by: Cursor <cursoragent at cursor.com>
---
.../CIR/Dialect/Builder/CIRBaseBuilder.h | 38 ++++-
.../clang/CIR/Dialect/IR/CIREnumAttr.td | 30 ++++
clang/include/clang/CIR/Dialect/IR/CIROps.td | 71 ++++++---
clang/lib/CIR/CodeGen/CIRGenBuiltin.cpp | 21 ++-
clang/lib/CIR/CodeGen/CIRGenExprScalar.cpp | 32 ++--
clang/lib/CIR/CodeGen/CIRGenFunction.cpp | 18 ++-
clang/lib/CIR/CodeGen/CIRGenFunction.h | 2 +
clang/lib/CIR/Dialect/IR/CIRDialect.cpp | 6 +
.../CIR/Lowering/DirectToLLVM/LowerToLLVM.cpp | 104 ++++++++++---
.../CIR/Lowering/DirectToLLVM/LowerToLLVM.h | 4 +
clang/test/CIR/CodeGen/fp-contract-fast.c | 138 ++++++++++++++++++
.../test/CIR/CodeGen/fp-math-precision-opts.c | 8 +-
clang/test/CIR/CodeGenCUDA/fp-contract.cu | 57 ++++++++
clang/test/CIR/Lowering/fastmath-contract.cir | 47 ++++++
clang/utils/TableGen/CIRLoweringEmitter.cpp | 17 ++-
15 files changed, 527 insertions(+), 66 deletions(-)
create mode 100644 clang/test/CIR/CodeGen/fp-contract-fast.c
create mode 100644 clang/test/CIR/CodeGenCUDA/fp-contract.cu
create mode 100644 clang/test/CIR/Lowering/fastmath-contract.cir
diff --git a/clang/include/clang/CIR/Dialect/Builder/CIRBaseBuilder.h b/clang/include/clang/CIR/Dialect/Builder/CIRBaseBuilder.h
index 29f1a64ad1d17f..e8ec07a66a77a1 100644
--- a/clang/include/clang/CIR/Dialect/Builder/CIRBaseBuilder.h
+++ b/clang/include/clang/CIR/Dialect/Builder/CIRBaseBuilder.h
@@ -74,6 +74,18 @@ class CIRBaseBuilderTy : public mlir::OpBuilder {
clang::LangOptions::FPE_Ignore;
llvm::RoundingMode defaultConstrainedRounding =
llvm::RoundingMode::NearestTiesToEven;
+ // Fast-math flags applied to floating-point ops created by this builder.
+ // CIRGen currently populates `contract` only.
+ cir::FastMathFlags fastMathFlags = cir::FastMathFlags::none;
+
+ void setFastMathFlags(cir::FastMathFlags flags) { fastMathFlags = flags; }
+ cir::FastMathFlags getFastMathFlags() const { return fastMathFlags; }
+
+ cir::FastMathFlagsAttr getFastMathFlagsAttr() {
+ if (fastMathFlags == cir::FastMathFlags::none)
+ return {};
+ return cir::FastMathFlagsAttr::get(getContext(), fastMathFlags);
+ }
mlir::Value getConstAPInt(mlir::Location loc, mlir::Type typ,
const llvm::APInt &val) {
@@ -848,32 +860,39 @@ class CIRBaseBuilderTy : public mlir::OpBuilder {
mlir::Value createFAdd(mlir::Location loc, mlir::Value lhs, mlir::Value rhs) {
assert(!cir::MissingFeatures::metaDataNode());
+ // `contract` is applied via getFastMathFlagsAttr(). The other fast-math
+ // bits are still unimplemented.
assert(!cir::MissingFeatures::fastMathFlags());
- return cir::FAddOp::create(*this, loc, lhs, rhs, getConstrainedFPAttr());
+ return cir::FAddOp::create(*this, loc, lhs, rhs, getConstrainedFPAttr(),
+ getFastMathFlagsAttr());
}
mlir::Value createFSub(mlir::Location loc, mlir::Value lhs, mlir::Value rhs) {
assert(!cir::MissingFeatures::metaDataNode());
assert(!cir::MissingFeatures::fastMathFlags());
- return cir::FSubOp::create(*this, loc, lhs, rhs, getConstrainedFPAttr());
+ return cir::FSubOp::create(*this, loc, lhs, rhs, getConstrainedFPAttr(),
+ getFastMathFlagsAttr());
}
mlir::Value createFMul(mlir::Location loc, mlir::Value lhs, mlir::Value rhs) {
assert(!cir::MissingFeatures::metaDataNode());
assert(!cir::MissingFeatures::fastMathFlags());
- return cir::FMulOp::create(*this, loc, lhs, rhs, getConstrainedFPAttr());
+ return cir::FMulOp::create(*this, loc, lhs, rhs, getConstrainedFPAttr(),
+ getFastMathFlagsAttr());
}
mlir::Value createFDiv(mlir::Location loc, mlir::Value lhs, mlir::Value rhs) {
assert(!cir::MissingFeatures::metaDataNode());
assert(!cir::MissingFeatures::fastMathFlags());
- return cir::FDivOp::create(*this, loc, lhs, rhs, getConstrainedFPAttr());
+ return cir::FDivOp::create(*this, loc, lhs, rhs, getConstrainedFPAttr(),
+ getFastMathFlagsAttr());
}
mlir::Value createFRem(mlir::Location loc, mlir::Value lhs, mlir::Value rhs) {
assert(!cir::MissingFeatures::metaDataNode());
assert(!cir::MissingFeatures::fastMathFlags());
- return cir::FRemOp::create(*this, loc, lhs, rhs, getConstrainedFPAttr());
+ return cir::FRemOp::create(*this, loc, lhs, rhs, getConstrainedFPAttr(),
+ getFastMathFlagsAttr());
}
mlir::Value createFNeg(mlir::Location loc, mlir::Value operand) {
@@ -883,7 +902,7 @@ class CIRBaseBuilderTy : public mlir::OpBuilder {
assert(!cir::MissingFeatures::fastMathFlags());
// fneg does not raise FP exceptions or depend on the rounding mode, so it
// never carries an fenv attribute.
- return cir::FNegOp::create(*this, loc, operand);
+ return cir::FNegOp::create(*this, loc, operand, getFastMathFlagsAttr());
}
mlir::Value createXor(mlir::Location loc, mlir::Value lhs, mlir::Value rhs) {
@@ -903,7 +922,10 @@ class CIRBaseBuilderTy : public mlir::OpBuilder {
cir::FenvAttr fenv;
if (cir::isAnyFloatingPointType(lhs.getType()))
fenv = getConstrainedFPAttr();
- return cir::CmpOp::create(*this, loc, kind, lhs, rhs, fenv);
+ return cir::CmpOp::create(*this, loc, kind, lhs, rhs, fenv,
+ cir::isAnyFloatingPointType(lhs.getType())
+ ? getFastMathFlagsAttr()
+ : cir::FastMathFlagsAttr{});
}
cir::VecCmpOp createVecCompare(mlir::Location loc, cir::CmpOpKind kind,
@@ -917,7 +939,7 @@ class CIRBaseBuilderTy : public mlir::OpBuilder {
if (cir::isFPOrVectorOfFPType(lhs.getType()))
fenv = getConstrainedFPAttr();
return cir::VecCmpOp::create(*this, loc, integralVecTy, kind, lhs, rhs,
- fenv);
+ fenv, getFastMathFlagsAttr());
}
mlir::Value createIsNaN(mlir::Location loc, mlir::Value operand) {
diff --git a/clang/include/clang/CIR/Dialect/IR/CIREnumAttr.td b/clang/include/clang/CIR/Dialect/IR/CIREnumAttr.td
index 3ee06412d8a90f..ecd981593ec047 100644
--- a/clang/include/clang/CIR/Dialect/IR/CIREnumAttr.td
+++ b/clang/include/clang/CIR/Dialect/IR/CIREnumAttr.td
@@ -47,6 +47,36 @@ class CIR_EnumAttr<EnumInfo info, string name = "", list<Trait> traits = []>
let assemblyFormat = "`<` $value `>`";
}
+// Bit positions match `mlir::LLVM::FastmathFlags`. Only `contract` is set by
+// CIRGen today; the other bits exist so the attribute can round-trip.
+def CIR_FastMathFlags : CIR_I32BitEnum<
+ "FastMathFlags", "fast-math flags", [
+ I32BitEnumCaseNone<"none">,
+ I32BitEnumCaseBit<"nnan", 0>,
+ I32BitEnumCaseBit<"ninf", 1>,
+ I32BitEnumCaseBit<"nsz", 2>,
+ I32BitEnumCaseBit<"arcp", 3>,
+ I32BitEnumCaseBit<"contract", 4>,
+ I32BitEnumCaseBit<"afn", 5>,
+ I32BitEnumCaseBit<"reassoc", 6>
+]> {
+ let description = [{
+ Per-operation fast-math flags. These are the LLVM fast-math bits, carried
+ on CIR floating-point operations and lowered onto the corresponding LLVM
+ dialect operation's `fastmathFlags`.
+
+ `contract` allows the backend to fuse this operation with another
+ floating-point operation, including across statements. It is what
+ `-ffp-contract=fast` records. `-ffp-contract=on` is represented separately
+ by `cir.fmuladd`.
+ }];
+ let separator = ", ";
+}
+
+def CIR_FastMathFlagsAttr : CIR_EnumAttr<CIR_FastMathFlags, "fastmath"> {
+ let summary = "Fast-math flags for a floating-point operation";
+}
+
def CIR_LangAddressSpace : CIR_I32Enum<
"LangAddressSpace", "language address space kind", [
I32EnumCase<"Default", 0, "default">,
diff --git a/clang/include/clang/CIR/Dialect/IR/CIROps.td b/clang/include/clang/CIR/Dialect/IR/CIROps.td
index 4996037ea5f560..9d222de1cd8a66 100644
--- a/clang/include/clang/CIR/Dialect/IR/CIROps.td
+++ b/clang/include/clang/CIR/Dialect/IR/CIROps.td
@@ -87,6 +87,10 @@ class LLVMLoweringInfo {
string llvmOp = "";
string constrainedLLVMIntrinsic = "";
bit constrainedLLVMIntrinsicHasRoundingMode = true;
+ // Copy an optional `fastmath` attribute onto the lowered LLVM operation.
+ // Floating-point ops that go through `lowerConstrainableFPOp` propagate the
+ // attribute there and do not need this bit.
+ bit propagateFastMathFlags = false;
}
class LoweringBuilders<dag p> {
@@ -2160,18 +2164,31 @@ def CIR_FNegOp : CIR_UnaryOp<"fneg", CIR_AnyFloatOrVecOfFloatType> {
The `cir.fneg` operation negates the operand. The operand and result must
have the same type.
+ The optional `fastmath` attribute carries LLVM fast-math flags for this
+ operation. `-ffp-contract=fast` sets `contract`.
+
Example:
```
%1 = cir.fneg %0 : !cir.float
- %3 = cir.fneg %2 : !cir.double
+ %3 = cir.fneg %2 : !cir.double {fastmath = #cir.fastmath<contract>}
%5 = cir.fneg %4 : !cir.vector<4 x !cir.float>
```
}];
+ let arguments = !con(commonArgs,
+ (ins OptionalAttr<CIR_FastMathFlagsAttr>:$fastmath));
+
+ let builders = [
+ OpBuilder<(ins "mlir::Value":$input), [{
+ build($_builder, $_state, input, cir::FastMathFlagsAttr{});
+ }]>
+ ];
+
let hasFolder = 1;
let llvmOp = "FNegOp";
+ let propagateFastMathFlags = true;
}
//===----------------------------------------------------------------------===//
@@ -2626,7 +2643,8 @@ def CIR_CmpOp : CIR_Op<"cmp",
CIR_CmpOpKindAttr:$kind,
CIR_ComparableType:$lhs,
CIR_ComparableType:$rhs,
- OptionalAttr<CIR_FenvAttr>:$fenv
+ OptionalAttr<CIR_FenvAttr>:$fenv,
+ OptionalAttr<CIR_FastMathFlagsAttr>:$fastmath
);
let results = (outs CIR_BoolType:$result);
@@ -2638,11 +2656,13 @@ def CIR_CmpOp : CIR_Op<"cmp",
let builders = [
OpBuilder<(ins "cir::CmpOpKind":$kind, "mlir::Value":$lhs,
"mlir::Value":$rhs), [{
- build($_builder, $_state, kind, lhs, rhs, cir::FenvAttr{});
+ build($_builder, $_state, kind, lhs, rhs, cir::FenvAttr{},
+ cir::FastMathFlagsAttr{});
}]>,
OpBuilder<(ins "mlir::Type":$result, "cir::CmpOpKind":$kind,
"mlir::Value":$lhs, "mlir::Value":$rhs), [{
- build($_builder, $_state, result, kind, lhs, rhs, cir::FenvAttr{});
+ build($_builder, $_state, result, kind, lhs, rhs, cir::FenvAttr{},
+ cir::FastMathFlagsAttr{});
}]>
];
@@ -2941,18 +2961,23 @@ def CIR_RemOp : CIR_BinaryOp<"rem", CIR_AnyIntOrVecOfIntType> {
// and result must all be the same floating-point scalar or vector type.
//
// The optional `fenv` attribute describes constraints on the floating-point
-// handling of the operation.
+// handling of the operation. The optional `fastmath` attribute carries LLVM
+// fast-math flags; `-ffp-contract=fast` sets `contract` here rather than
+// forming `cir.fmuladd`.
class CIR_FPBinaryOp<string mnemonic, list<Trait> traits = []>
: CIR_BinaryOp<mnemonic, CIR_AnyFloatOrVecOfFloatType,
!listconcat(CIR_FenvOpTraits, traits),
CIR_DynamicMemoryEffects> {
- let arguments = !con(commonArgs, (ins OptionalAttr<CIR_FenvAttr>:$fenv));
+ let arguments = !con(commonArgs, (ins
+ OptionalAttr<CIR_FenvAttr>:$fenv,
+ OptionalAttr<CIR_FastMathFlagsAttr>:$fastmath));
let constrainedLLVMIntrinsic = mnemonic;
let builders = [
OpBuilder<(ins "mlir::Value":$lhs, "mlir::Value":$rhs), [{
- build($_builder, $_state, lhs, rhs, cir::FenvAttr{});
+ build($_builder, $_state, lhs, rhs, cir::FenvAttr{},
+ cir::FastMathFlagsAttr{});
}]>
];
}
@@ -6094,7 +6119,8 @@ def CIR_VecCmpOp : CIR_Op<"vec.cmp",
CIR_CmpOpKindAttr:$kind,
CIR_VectorType:$lhs,
CIR_VectorType:$rhs,
- OptionalAttr<CIR_FenvAttr>:$fenv
+ OptionalAttr<CIR_FenvAttr>:$fenv,
+ OptionalAttr<CIR_FastMathFlagsAttr>:$fastmath
);
let results = (outs CIR_VectorType:$result);
@@ -6107,7 +6133,8 @@ def CIR_VecCmpOp : CIR_Op<"vec.cmp",
let builders = [
OpBuilder<(ins "mlir::Type":$result, "cir::CmpOpKind":$kind,
"mlir::Value":$lhs, "mlir::Value":$rhs), [{
- build($_builder, $_state, result, kind, lhs, rhs, cir::FenvAttr{});
+ build($_builder, $_state, result, kind, lhs, rhs, cir::FenvAttr{},
+ cir::FastMathFlagsAttr{});
}]>
];
@@ -7448,14 +7475,16 @@ class CIR_UnaryFPToFPBuiltinOp<string mnemonic, string llvmOpName>
!listconcat([SameOperandsAndResultType], CIR_FenvOpTraits)>
{
let arguments = (ins CIR_AnyFloatOrVecOfFloatType:$src,
- OptionalAttr<CIR_FenvAttr>:$fenv);
+ OptionalAttr<CIR_FenvAttr>:$fenv,
+ OptionalAttr<CIR_FastMathFlagsAttr>:$fastmath);
let results = (outs CIR_AnyFloatOrVecOfFloatType:$result);
let assemblyFormat = "$src `:` type($src) attr-dict";
let builders = [
OpBuilder<(ins "mlir::Value":$src), [{
- build($_builder, $_state, src, cir::FenvAttr{});
+ build($_builder, $_state, src, cir::FenvAttr{},
+ cir::FastMathFlagsAttr{});
}]>
];
@@ -7716,6 +7745,7 @@ def CIR_FAbsOp : CIR_UnaryFPToFPBuiltinOp<"fabs", "FAbsOp"> {
// fabs is exact and does not raise exceptions, so it is always lowered to
// the plain llvm.fabs intrinsic.
let constrainedLLVMIntrinsic = "";
+ let propagateFastMathFlags = true;
}
def CIR_AbsOp : CIR_Op<"abs", [Pure, SameOperandsAndResultType]> {
@@ -7770,7 +7800,8 @@ class CIR_UnaryFPToIntBuiltinOp<string mnemonic, string llvmOpName>
: CIR_Op<mnemonic, CIR_FenvOpTraits>
{
let arguments = (ins CIR_AnyFloatType:$src,
- OptionalAttr<CIR_FenvAttr>:$fenv);
+ OptionalAttr<CIR_FenvAttr>:$fenv,
+ OptionalAttr<CIR_FastMathFlagsAttr>:$fastmath);
let results = (outs CIR_IntType:$result);
let summary = [{
@@ -7784,7 +7815,8 @@ class CIR_UnaryFPToIntBuiltinOp<string mnemonic, string llvmOpName>
let builders = [
OpBuilder<(ins "mlir::Type":$result, "mlir::Value":$src), [{
- build($_builder, $_state, result, src, cir::FenvAttr{});
+ build($_builder, $_state, result, src, cir::FenvAttr{},
+ cir::FastMathFlagsAttr{});
}]>
];
@@ -7839,7 +7871,8 @@ class CIR_BinaryFPToFPBuiltinOp<string mnemonic, string llvmOpName>
let arguments = (ins
CIR_AnyFloatOrVecOfFloatType:$lhs,
CIR_AnyFloatOrVecOfFloatType:$rhs,
- OptionalAttr<CIR_FenvAttr>:$fenv
+ OptionalAttr<CIR_FenvAttr>:$fenv,
+ OptionalAttr<CIR_FastMathFlagsAttr>:$fastmath
);
let results = (outs CIR_AnyFloatOrVecOfFloatType:$result);
@@ -7851,7 +7884,8 @@ class CIR_BinaryFPToFPBuiltinOp<string mnemonic, string llvmOpName>
let builders = [
OpBuilder<(ins "mlir::Type":$result, "mlir::Value":$lhs,
"mlir::Value":$rhs), [{
- build($_builder, $_state, result, lhs, rhs, cir::FenvAttr{});
+ build($_builder, $_state, result, lhs, rhs, cir::FenvAttr{},
+ cir::FastMathFlagsAttr{});
}]>
];
@@ -7867,6 +7901,7 @@ def CIR_CopysignOp : CIR_BinaryFPToFPBuiltinOp<"copysign", "CopySignOp"> {
// copysign is exact and does not raise exceptions, so it is always lowered
// to the plain llvm.copysign intrinsic.
+ let propagateFastMathFlags = true;
}
def CIR_FMaxNumOp : CIR_BinaryFPToFPBuiltinOp<"fmaxnum", "MaxNumOp"> {
@@ -7991,7 +8026,8 @@ class CIR_TernaryFPToFPBuiltinOp<string mnemonic, string llvmOpName>
CIR_AnyFloatOrVecOfFloatType:$a,
CIR_AnyFloatOrVecOfFloatType:$b,
CIR_AnyFloatOrVecOfFloatType:$c,
- OptionalAttr<CIR_FenvAttr>:$fenv
+ OptionalAttr<CIR_FenvAttr>:$fenv,
+ OptionalAttr<CIR_FastMathFlagsAttr>:$fastmath
);
let results = (outs CIR_AnyFloatOrVecOfFloatType:$result);
@@ -8001,7 +8037,8 @@ class CIR_TernaryFPToFPBuiltinOp<string mnemonic, string llvmOpName>
let builders = [
OpBuilder<(ins "mlir::Type":$result, "mlir::Value":$a, "mlir::Value":$b,
"mlir::Value":$c), [{
- build($_builder, $_state, result, a, b, c, cir::FenvAttr{});
+ build($_builder, $_state, result, a, b, c, cir::FenvAttr{},
+ cir::FastMathFlagsAttr{});
}]>
];
diff --git a/clang/lib/CIR/CodeGen/CIRGenBuiltin.cpp b/clang/lib/CIR/CodeGen/CIRGenBuiltin.cpp
index 245708691b7d99..41c09cb43fcb85 100644
--- a/clang/lib/CIR/CodeGen/CIRGenBuiltin.cpp
+++ b/clang/lib/CIR/CodeGen/CIRGenBuiltin.cpp
@@ -558,14 +558,18 @@ static RValue emitUnaryMaybeConstrainedFPBuiltin(CIRGenFunction &cgf,
CIRGenFunction::CIRGenFPOptionsRAII FPOptsRAII(cgf, &e);
auto call = Operation::create(cgf.getBuilder(), arg.getLoc(), arg.getType(),
- arg, cgf.getBuilder().getConstrainedFPAttr());
+ arg, cgf.getBuilder().getConstrainedFPAttr(),
+ cgf.getBuilder().getFastMathFlagsAttr());
return RValue::get(call->getResult(0));
}
template <class Operation>
static RValue emitUnaryFPBuiltin(CIRGenFunction &cgf, const CallExpr &e) {
mlir::Value arg = cgf.emitScalarExpr(e.getArg(0));
- auto call = Operation::create(cgf.getBuilder(), arg.getLoc(), arg);
+ CIRGenFunction::CIRGenFPOptionsRAII FPOptsRAII(cgf, &e);
+ auto call = Operation::create(cgf.getBuilder(), arg.getLoc(), arg.getType(),
+ arg, cir::FenvAttr{},
+ cgf.getBuilder().getFastMathFlagsAttr());
return RValue::get(call->getResult(0));
}
@@ -578,7 +582,8 @@ static RValue emitUnaryMaybeConstrainedFPToIntBuiltin(CIRGenFunction &cgf,
CIRGenFunction::CIRGenFPOptionsRAII FPOptsRAII(cgf, &e);
auto call = Op::create(cgf.getBuilder(), src.getLoc(), resultType, src,
- cgf.getBuilder().getConstrainedFPAttr());
+ cgf.getBuilder().getConstrainedFPAttr(),
+ cgf.getBuilder().getFastMathFlagsAttr());
return RValue::get(call->getResult(0));
}
@@ -587,9 +592,11 @@ static RValue emitBinaryFPBuiltin(CIRGenFunction &cgf, const CallExpr &e) {
mlir::Value arg0 = cgf.emitScalarExpr(e.getArg(0));
mlir::Value arg1 = cgf.emitScalarExpr(e.getArg(1));
+ CIRGenFunction::CIRGenFPOptionsRAII FPOptsRAII(cgf, &e);
mlir::Location loc = cgf.getLoc(e.getExprLoc());
mlir::Type ty = cgf.convertType(e.getType());
- auto call = Op::create(cgf.getBuilder(), loc, ty, arg0, arg1);
+ auto call = Op::create(cgf.getBuilder(), loc, ty, arg0, arg1, cir::FenvAttr{},
+ cgf.getBuilder().getFastMathFlagsAttr());
return RValue::get(call->getResult(0));
}
@@ -621,7 +628,8 @@ static RValue emitTernaryMaybeConstrainedFPBuiltin(CIRGenFunction &cgf,
mlir::Type ty = cgf.convertType(e.getType());
auto call = Op::create(cgf.getBuilder(), loc, ty, arg0, arg1, arg2,
- cgf.getBuilder().getConstrainedFPAttr());
+ cgf.getBuilder().getConstrainedFPAttr(),
+ cgf.getBuilder().getFastMathFlagsAttr());
return RValue::get(call->getResult(0));
}
@@ -637,7 +645,8 @@ static mlir::Value emitBinaryMaybeConstrainedFPBuiltin(CIRGenFunction &cgf,
mlir::Type ty = cgf.convertType(e.getType());
auto call = Op::create(cgf.getBuilder(), loc, ty, arg0, arg1,
- cgf.getBuilder().getConstrainedFPAttr());
+ cgf.getBuilder().getConstrainedFPAttr(),
+ cgf.getBuilder().getFastMathFlagsAttr());
return call->getResult(0);
}
diff --git a/clang/lib/CIR/CodeGen/CIRGenExprScalar.cpp b/clang/lib/CIR/CodeGen/CIRGenExprScalar.cpp
index 80e8b36dfda6bf..cc29109883b77b 100644
--- a/clang/lib/CIR/CodeGen/CIRGenExprScalar.cpp
+++ b/clang/lib/CIR/CodeGen/CIRGenExprScalar.cpp
@@ -899,8 +899,11 @@ class ScalarExprEmitter : public StmtVisitor<ScalarExprEmitter, mlir::Value> {
mlir::Location loc = cgf.getLoc(e->getSourceRange().getBegin());
- if (cir::isFPOrVectorOfFPType(operand.getType()))
- return builder.createOrFold<cir::FNegOp>(loc, operand);
+ if (cir::isFPOrVectorOfFPType(operand.getType())) {
+ CIRGenFunction::CIRGenFPOptionsRAII FPOptsRAII(cgf, e);
+ return builder.createOrFold<cir::FNegOp>(loc, operand,
+ builder.getFastMathFlagsAttr());
+ }
// TODO(cir): We might have to change this to support overflow trapping.
// Classic codegen routes unary minus through emitSub to ensure
@@ -1281,6 +1284,9 @@ class ScalarExprEmitter : public StmtVisitor<ScalarExprEmitter, mlir::Value> {
BinOpInfo boInfo = emitBinOps(e);
mlir::Value lhs = boInfo.lhs;
mlir::Value rhs = boInfo.rhs;
+ std::optional<CIRGenFunction::CIRGenFPOptionsRAII> fpOpts;
+ if (cir::isFPOrVectorOfFPType(lhs.getType()))
+ fpOpts.emplace(cgf, boInfo.fpFeatures);
if (lhsTy->isVectorType()) {
if (!e->getType()->isVectorType()) {
@@ -1290,9 +1296,15 @@ class ScalarExprEmitter : public StmtVisitor<ScalarExprEmitter, mlir::Value> {
} else {
// Other kinds of vectors. Element-wise comparison returning
// a vector.
- result = cir::VecCmpOp::create(builder, cgf.getLoc(boInfo.loc),
- cgf.convertType(boInfo.fullType), kind,
- boInfo.lhs, boInfo.rhs);
+ result = cir::VecCmpOp::create(
+ builder, cgf.getLoc(boInfo.loc), cgf.convertType(boInfo.fullType),
+ kind, boInfo.lhs, boInfo.rhs,
+ cir::isFPOrVectorOfFPType(boInfo.lhs.getType())
+ ? builder.getConstrainedFPAttr()
+ : cir::FenvAttr{},
+ cir::isFPOrVectorOfFPType(boInfo.lhs.getType())
+ ? builder.getFastMathFlagsAttr()
+ : cir::FastMathFlagsAttr{});
}
} else if (boInfo.isFixedPointOp()) {
result = emitFixedPointBinOp(boInfo);
@@ -2069,9 +2081,9 @@ static mlir::Value buildFMulAdd(mlir::Location addLoc, cir::FMulOp mulOp,
// Carry the mul's fenv attribute so a constrained fmul yields a constrained
// fmuladd; the builder is under the add's FP options, not the mul's.
- mlir::Value fmuladd =
- cir::FMulAddOp::create(builder, loc, addend.getType(), mulOp0, mulOp1,
- addend, mulOp.getFenvAttr());
+ mlir::Value fmuladd = cir::FMulAddOp::create(
+ builder, loc, addend.getType(), mulOp0, mulOp1, addend,
+ mulOp.getFenvAttr(), cir::FastMathFlagsAttr{});
mulOp.erase();
return fmuladd;
}
@@ -2090,8 +2102,8 @@ static mlir::Value tryEmitFMulAdd(mlir::Location loc, const BinOpInfo &op,
"Only fadd/fsub can be the root of an fmuladd.");
// Check whether this op is fusable, i.e. -ffp-contract=on. -ffp-contract=fast
- // needs fast-math flags on the fmul/fadd, which CIR does not model yet, so it
- // fuses nowhere for now.
+ // is not a cir.fmuladd: the builder stamps `contract` on the fmul and fadd,
+ // which is what the backend fuses when it is no longer in Fast mode.
assert(!cir::MissingFeatures::fastMathFlags());
if (!op.fpFeatures.allowFPContractWithinStatement())
return nullptr;
diff --git a/clang/lib/CIR/CodeGen/CIRGenFunction.cpp b/clang/lib/CIR/CodeGen/CIRGenFunction.cpp
index 333610196c7965..bb9ccdedbe5afb 100644
--- a/clang/lib/CIR/CodeGen/CIRGenFunction.cpp
+++ b/clang/lib/CIR/CodeGen/CIRGenFunction.cpp
@@ -44,6 +44,16 @@ static bool functionMightHaveBypass(const Stmt *s) {
return false;
}
+static cir::FastMathFlags fastMathFlagsFromFPOptions(clang::FPOptions fpFeatures) {
+ // Other fast-math bits (nnan, ninf, reassoc, ...) are not modeled yet.
+ // `contract` is the bit `-ffp-contract=fast` needs so a later backend in
+ // Standard fusion mode can still form an FMA.
+ cir::FastMathFlags flags = cir::FastMathFlags::none;
+ if (fpFeatures.allowFPContractAcrossStatement())
+ flags = flags | cir::FastMathFlags::contract;
+ return flags;
+}
+
CIRGenFunction::CIRGenFunction(CIRGenModule &cgm, CIRGenBuilderTy &builder,
bool suppressNewContext)
: CIRGenTypeCache(cgm), cgm{cgm}, builder(builder),
@@ -51,6 +61,7 @@ CIRGenFunction::CIRGenFunction(CIRGenModule &cgm, CIRGenBuilderTy &builder,
ehStack.setCGF(this);
shouldEmitLifetimeMarkers = CodeGenUtils::shouldEmitLifetimeMarkers(
cgm.getCodeGenOpts(), getContext().getLangOpts());
+ builder.setFastMathFlags(fastMathFlagsFromFPOptions(curFPFeatures));
}
CIRGenFunction::~CIRGenFunction() {}
@@ -1436,6 +1447,7 @@ void CIRGenFunction::CIRGenFPOptionsRAII::ConstructorHelper(
oldExcept = cgf.builder.getDefaultConstrainedExcept();
oldRounding = cgf.builder.getDefaultConstrainedRounding();
+ oldFastMathFlags = cgf.builder.getFastMathFlags();
if (oldFPFeatures == fpFeatures)
return;
@@ -1449,8 +1461,10 @@ void CIRGenFunction::CIRGenFPOptionsRAII::ConstructorHelper(
cgf.builder.setDefaultConstrainedRounding(newRoundingMode);
cgf.builder.setDefaultConstrainedExcept(newExceptionBehavior);
+ cgf.builder.setFastMathFlags(fastMathFlagsFromFPOptions(fpFeatures));
+ restoredFastMathFlags = true;
- // TODO(cir): override FP flags once FM configs are guarded.
+ // nnan/ninf/reassoc/arcp/afn are still missing. `contract` is applied above.
assert(!cir::MissingFeatures::fastMathFlags());
assert((cgf.curFuncDecl == nullptr || cgf.builder.getIsFPConstrained() ||
@@ -1468,6 +1482,8 @@ CIRGenFunction::CIRGenFPOptionsRAII::~CIRGenFPOptionsRAII() {
cgf.curFPFeatures = oldFPFeatures;
cgf.builder.setDefaultConstrainedExcept(oldExcept);
cgf.builder.setDefaultConstrainedRounding(oldRounding);
+ if (restoredFastMathFlags)
+ cgf.builder.setFastMathFlags(oldFastMathFlags);
}
// TODO(cir): should be shared with LLVM codegen.
diff --git a/clang/lib/CIR/CodeGen/CIRGenFunction.h b/clang/lib/CIR/CodeGen/CIRGenFunction.h
index b232abf5d9299d..b1230c79c8b611 100644
--- a/clang/lib/CIR/CodeGen/CIRGenFunction.h
+++ b/clang/lib/CIR/CodeGen/CIRGenFunction.h
@@ -337,6 +337,8 @@ class CIRGenFunction : public CIRGenTypeCache {
clang::FPOptions oldFPFeatures;
LangOptions::FPExceptionModeKind oldExcept;
llvm::RoundingMode oldRounding;
+ cir::FastMathFlags oldFastMathFlags = cir::FastMathFlags::none;
+ bool restoredFastMathFlags = false;
};
clang::FPOptions curFPFeatures;
diff --git a/clang/lib/CIR/Dialect/IR/CIRDialect.cpp b/clang/lib/CIR/Dialect/IR/CIRDialect.cpp
index ce406707f5942f..d948ed9608d352 100644
--- a/clang/lib/CIR/Dialect/IR/CIRDialect.cpp
+++ b/clang/lib/CIR/Dialect/IR/CIRDialect.cpp
@@ -3857,6 +3857,9 @@ LogicalResult cir::CmpOp::verify() {
if (getFenvAttr() && !cir::isAnyFloatingPointType(getLhs().getType()))
return emitOpError()
<< "'fenv' is only valid for floating-point comparisons";
+ if (getFastmathAttr() && !cir::isAnyFloatingPointType(getLhs().getType()))
+ return emitOpError()
+ << "'fastmath' is only valid for floating-point comparisons";
return success();
}
@@ -3868,6 +3871,9 @@ LogicalResult cir::VecCmpOp::verify() {
if (getFenvAttr() && !cir::isFPOrVectorOfFPType(getLhs().getType()))
return emitOpError()
<< "'fenv' is only valid for floating-point comparisons";
+ if (getFastmathAttr() && !cir::isFPOrVectorOfFPType(getLhs().getType()))
+ return emitOpError()
+ << "'fastmath' is only valid for floating-point comparisons";
return success();
}
diff --git a/clang/lib/CIR/Lowering/DirectToLLVM/LowerToLLVM.cpp b/clang/lib/CIR/Lowering/DirectToLLVM/LowerToLLVM.cpp
index 4bdbe38df24e8e..7db03d4e56fc4a 100644
--- a/clang/lib/CIR/Lowering/DirectToLLVM/LowerToLLVM.cpp
+++ b/clang/lib/CIR/Lowering/DirectToLLVM/LowerToLLVM.cpp
@@ -550,6 +550,58 @@ static llvm::StringRef getConstrainedExceptMetadata(cir::FenvAttr fenv) {
return strictExcept.getValue() ? "fpexcept.strict" : "fpexcept.maytrap";
}
+// CIR fast-math bits use the same positions as `mlir::LLVM::FastmathFlags`.
+static mlir::LLVM::FastmathFlags
+toLLVMFastMathFlags(cir::FastMathFlags flags) {
+ using CirFlags = cir::FastMathFlags;
+ using LLVMFlags = mlir::LLVM::FastmathFlags;
+ static_assert(static_cast<uint32_t>(CirFlags::nnan) ==
+ static_cast<uint32_t>(LLVMFlags::nnan),
+ "CIR nnan bit must match LLVM fastmath");
+ static_assert(static_cast<uint32_t>(CirFlags::ninf) ==
+ static_cast<uint32_t>(LLVMFlags::ninf),
+ "CIR ninf bit must match LLVM fastmath");
+ static_assert(static_cast<uint32_t>(CirFlags::nsz) ==
+ static_cast<uint32_t>(LLVMFlags::nsz),
+ "CIR nsz bit must match LLVM fastmath");
+ static_assert(static_cast<uint32_t>(CirFlags::arcp) ==
+ static_cast<uint32_t>(LLVMFlags::arcp),
+ "CIR arcp bit must match LLVM fastmath");
+ static_assert(static_cast<uint32_t>(CirFlags::contract) ==
+ static_cast<uint32_t>(LLVMFlags::contract),
+ "CIR contract bit must match LLVM fastmath");
+ static_assert(static_cast<uint32_t>(CirFlags::afn) ==
+ static_cast<uint32_t>(LLVMFlags::afn),
+ "CIR afn bit must match LLVM fastmath");
+ static_assert(static_cast<uint32_t>(CirFlags::reassoc) ==
+ static_cast<uint32_t>(LLVMFlags::reassoc),
+ "CIR reassoc bit must match LLVM fastmath");
+ return static_cast<LLVMFlags>(static_cast<uint32_t>(flags));
+}
+
+// Not inlined into lowerConstrainableFPOp. In that template, a local null
+// check on the attribute is dropped and getValue() crashes when `fastmath`
+// is absent (cir.fmuladd has no such property).
+__attribute__((noinline)) static mlir::LLVM::FastmathFlags
+readFastMathFlags(mlir::Operation *cirOp) {
+ auto cirFlags = cirOp->getAttrOfType<cir::FastMathFlagsAttr>("fastmath");
+ if (!cirFlags.getAsOpaquePointer() ||
+ cirFlags.getValue() == cir::FastMathFlags::none)
+ return mlir::LLVM::FastmathFlags::none;
+ return toLLVMFastMathFlags(cirFlags.getValue());
+}
+
+void propagateFastMathFlags(mlir::Operation *cirOp, mlir::Operation *llvmOp) {
+ mlir::LLVM::FastmathFlags flags = readFastMathFlags(cirOp);
+ if (flags == mlir::LLVM::FastmathFlags::none)
+ return;
+ auto fmfOp = dyn_cast<mlir::LLVM::FastmathFlagsInterface>(llvmOp);
+ if (!fmfOp)
+ return;
+ fmfOp.setFastmathAttr(
+ mlir::LLVM::FastmathFlagsAttr::get(llvmOp->getContext(), flags));
+}
+
static mlir::Value
createFenvMetadataValue(mlir::ConversionPatternRewriter &rewriter,
mlir::Location loc, llvm::StringRef str) {
@@ -588,14 +640,17 @@ mlir::LogicalResult lowerConstrainableFPOp(
return op->emitError("expected LLVM result type for floating-point op");
if (!fenv) {
- rewriter.replaceOpWithNewOp<LLVMOp>(
- op, mlir::TypeRange{llvmResTy}, operands,
+ LLVMOp llvmOp = LLVMOp::create(
+ rewriter, op->getLoc(), mlir::TypeRange{llvmResTy}, operands,
cir::getDefaultProperties<LLVMOp>(op->getContext()));
+ propagateFastMathFlags(op, llvmOp);
+ rewriter.replaceOp(op, llvmOp.getResult());
return mlir::success();
}
return lowerToConstrainedFPIntrinsic(op, operands, fenv, llvmResTy, rewriter,
- constrainedMnemonic, hasRoundingMode);
+ constrainedMnemonic, hasRoundingMode,
+ readFastMathFlags(op));
}
static mlir::LLVM::FastmathFlags
@@ -2091,13 +2146,15 @@ mlir::LogicalResult CIRToLLVMFMaxNumOpLowering::matchAndRewrite(
cir::FMaxNumOp op, OpAdaptor adaptor,
mlir::ConversionPatternRewriter &rewriter) const {
mlir::Type resTy = typeConverter->convertType(op.getType());
+ mlir::LLVM::FastmathFlags flags = static_cast<mlir::LLVM::FastmathFlags>(
+ static_cast<uint32_t>(mlir::LLVM::FastmathFlags::nsz) |
+ static_cast<uint32_t>(readFastMathFlags(op)));
if (cir::FenvAttr fenv = op.getFenvAttr())
return lowerToConstrainedFPIntrinsic(
op, adaptor.getOperands(), fenv, resTy, rewriter, "maxnum",
- /*hasRoundingMode=*/false, mlir::LLVM::FastmathFlags::nsz);
+ /*hasRoundingMode=*/false, flags);
rewriter.replaceOpWithNewOp<mlir::LLVM::MaxNumOp>(
- op, resTy, adaptor.getLhs(), adaptor.getRhs(),
- mlir::LLVM::FastmathFlags::nsz);
+ op, resTy, adaptor.getLhs(), adaptor.getRhs(), flags);
return mlir::success();
}
@@ -2105,13 +2162,15 @@ mlir::LogicalResult CIRToLLVMFMinNumOpLowering::matchAndRewrite(
cir::FMinNumOp op, OpAdaptor adaptor,
mlir::ConversionPatternRewriter &rewriter) const {
mlir::Type resTy = typeConverter->convertType(op.getType());
+ mlir::LLVM::FastmathFlags flags = static_cast<mlir::LLVM::FastmathFlags>(
+ static_cast<uint32_t>(mlir::LLVM::FastmathFlags::nsz) |
+ static_cast<uint32_t>(readFastMathFlags(op)));
if (cir::FenvAttr fenv = op.getFenvAttr())
return lowerToConstrainedFPIntrinsic(
op, adaptor.getOperands(), fenv, resTy, rewriter, "minnum",
- /*hasRoundingMode=*/false, mlir::LLVM::FastmathFlags::nsz);
+ /*hasRoundingMode=*/false, flags);
rewriter.replaceOpWithNewOp<mlir::LLVM::MinNumOp>(
- op, resTy, adaptor.getLhs(), adaptor.getRhs(),
- mlir::LLVM::FastmathFlags::nsz);
+ op, resTy, adaptor.getLhs(), adaptor.getRhs(), flags);
return mlir::success();
}
@@ -3613,7 +3672,8 @@ static mlir::LLVM::CallIntrinsicOp
createConstrainedFCmpCall(mlir::ConversionPatternRewriter &rewriter,
mlir::Location loc, mlir::Value lhs, mlir::Value rhs,
cir::CmpOpKind kind, cir::FenvAttr fenv,
- mlir::Type llvmResTy) {
+ mlir::Type llvmResTy,
+ mlir::LLVM::FastmathFlags fastmathFlags = {}) {
llvm::SmallVector<mlir::Value, 4> callOperands = {
lhs, rhs,
createFenvMetadataValue(rewriter, loc,
@@ -3624,7 +3684,7 @@ createConstrainedFCmpCall(mlir::ConversionPatternRewriter &rewriter,
? "llvm.experimental.constrained.fcmps"
: "llvm.experimental.constrained.fcmp";
return createCallLLVMIntrinsicOp(rewriter, loc, intrinsicName, llvmResTy,
- callOperands);
+ callOperands, fastmathFlags);
}
mlir::LogicalResult CIRToLLVMCmpOpLowering::matchAndRewrite(
@@ -3659,14 +3719,16 @@ mlir::LogicalResult CIRToLLVMCmpOpLowering::matchAndRewrite(
if (cir::FenvAttr fenv = cmpOp.getFenvAttr()) {
mlir::LLVM::CallIntrinsicOp call = createConstrainedFCmpCall(
rewriter, cmpOp.getLoc(), adaptor.getLhs(), adaptor.getRhs(),
- cmpOp.getKind(), fenv, llvmResTy);
+ cmpOp.getKind(), fenv, llvmResTy, readFastMathFlags(cmpOp));
rewriter.replaceOp(cmpOp, call.getResult(0));
return mlir::success();
}
mlir::LLVM::FCmpPredicate kind =
convertCmpKindToFCmpPredicate(cmpOp.getKind());
- rewriter.replaceOpWithNewOp<mlir::LLVM::FCmpOp>(
- cmpOp, kind, adaptor.getLhs(), adaptor.getRhs());
+ auto fcmp = mlir::LLVM::FCmpOp::create(
+ rewriter, cmpOp.getLoc(), kind, adaptor.getLhs(), adaptor.getRhs());
+ propagateFastMathFlags(cmpOp, fcmp);
+ rewriter.replaceOp(cmpOp, fcmp.getResult());
return mlir::success();
}
@@ -3702,6 +3764,8 @@ mlir::LogicalResult CIRToLLVMCmpOpLowering::matchAndRewrite(
rewriter, loc, mlir::LLVM::FCmpPredicate::oeq, lhsReal, rhsReal);
auto imagCmp = mlir::LLVM::FCmpOp::create(
rewriter, loc, mlir::LLVM::FCmpPredicate::oeq, lhsImag, rhsImag);
+ propagateFastMathFlags(cmpOp, realCmp);
+ propagateFastMathFlags(cmpOp, imagCmp);
rewriter.replaceOpWithNewOp<mlir::LLVM::AndOp>(cmpOp, realCmp, imagCmp);
return mlir::success();
}
@@ -3720,6 +3784,8 @@ mlir::LogicalResult CIRToLLVMCmpOpLowering::matchAndRewrite(
rewriter, loc, mlir::LLVM::FCmpPredicate::une, lhsReal, rhsReal);
auto imagCmp = mlir::LLVM::FCmpOp::create(
rewriter, loc, mlir::LLVM::FCmpPredicate::une, lhsImag, rhsImag);
+ propagateFastMathFlags(cmpOp, realCmp);
+ propagateFastMathFlags(cmpOp, imagCmp);
rewriter.replaceOpWithNewOp<mlir::LLVM::OrOp>(cmpOp, realCmp, imagCmp);
return mlir::success();
}
@@ -5054,14 +5120,16 @@ mlir::LogicalResult CIRToLLVMVecCmpOpLowering::matchAndRewrite(
if (cir::FenvAttr fenv = op.getFenvAttr()) {
auto i1VecTy = mlir::VectorType::get(op.getLhs().getType().getSize(),
rewriter.getI1Type());
- bitResult = createConstrainedFCmpCall(rewriter, op.getLoc(),
- adaptor.getLhs(), adaptor.getRhs(),
- op.getKind(), fenv, i1VecTy)
+ bitResult = createConstrainedFCmpCall(
+ rewriter, op.getLoc(), adaptor.getLhs(), adaptor.getRhs(),
+ op.getKind(), fenv, i1VecTy, readFastMathFlags(op))
.getResult(0);
} else {
- bitResult = mlir::LLVM::FCmpOp::create(
+ auto fcmp = mlir::LLVM::FCmpOp::create(
rewriter, op.getLoc(), convertCmpKindToFCmpPredicate(op.getKind()),
adaptor.getLhs(), adaptor.getRhs());
+ propagateFastMathFlags(op, fcmp);
+ bitResult = fcmp.getResult();
}
} else {
return op.emitError() << "unsupported type for VecCmpOp: " << elementType;
diff --git a/clang/lib/CIR/Lowering/DirectToLLVM/LowerToLLVM.h b/clang/lib/CIR/Lowering/DirectToLLVM/LowerToLLVM.h
index 56941b2aa51e7a..65f52cec486c63 100644
--- a/clang/lib/CIR/Lowering/DirectToLLVM/LowerToLLVM.h
+++ b/clang/lib/CIR/Lowering/DirectToLLVM/LowerToLLVM.h
@@ -77,6 +77,10 @@ struct LLVMBlockAddressInfo {
int32_t blockTagOpIndex;
};
+/// Copy a CIR `fastmath` attribute onto an LLVM dialect operation that
+/// implements `FastmathFlagsInterface`. No-op when the CIR op has no flags.
+void propagateFastMathFlags(mlir::Operation *cirOp, mlir::Operation *llvmOp);
+
mlir::LogicalResult lowerToConstrainedFPIntrinsic(
mlir::Operation *op, mlir::ValueRange operands, cir::FenvAttr fenv,
mlir::Type llvmResTy, mlir::ConversionPatternRewriter &rewriter,
diff --git a/clang/test/CIR/CodeGen/fp-contract-fast.c b/clang/test/CIR/CodeGen/fp-contract-fast.c
new file mode 100644
index 00000000000000..9c1b57e6a12354
--- /dev/null
+++ b/clang/test/CIR/CodeGen/fp-contract-fast.c
@@ -0,0 +1,138 @@
+// -ffp-contract=fast does not form cir.fmuladd. It stamps `contract` on the
+// individual floating-point ops so a backend in Standard fusion mode can
+// still contract them, including across statements.
+// -ffp-contract=on still forms cir.fmuladd and does not set `contract`.
+// -ffp-contract=off does neither.
+
+// RUN: %clang_cc1 -triple x86_64-unknown-linux-gnu -Wno-unused-value -fclangir -ffp-contract=fast -emit-cir %s -o %t-fast.cir
+// RUN: FileCheck --input-file=%t-fast.cir %s -check-prefix=CIR-FAST
+// RUN: %clang_cc1 -triple x86_64-unknown-linux-gnu -Wno-unused-value -fclangir -ffp-contract=fast -emit-llvm %s -o %t-fast.ll
+// RUN: FileCheck --input-file=%t-fast.ll %s -check-prefix=LLVM-FAST
+
+// RUN: %clang_cc1 -triple x86_64-unknown-linux-gnu -Wno-unused-value -fclangir -ffp-contract=on -emit-cir %s -o %t-on.cir
+// RUN: FileCheck --input-file=%t-on.cir %s -check-prefix=CIR-ON
+// RUN: %clang_cc1 -triple x86_64-unknown-linux-gnu -Wno-unused-value -fclangir -ffp-contract=off -emit-cir %s -o %t-off.cir
+// RUN: FileCheck --input-file=%t-off.cir %s -check-prefix=CIR-OFF
+
+// RUN: %clang_cc1 -triple x86_64-unknown-linux-gnu -Wno-unused-value -ffp-contract=fast -emit-llvm %s -o %t-og.ll
+// RUN: FileCheck --input-file=%t-og.ll %s -check-prefix=LLVM-FAST
+
+float same_stmt(float a, float b, float c) { return a * b + c; }
+// CIR-FAST-LABEL: cir.func {{.*}}@same_stmt
+// CIR-FAST: cir.fmul {{.*}} : !cir.float {fastmath = #cir.fastmath<contract>}
+// CIR-FAST: cir.fadd {{.*}} : !cir.float {fastmath = #cir.fastmath<contract>}
+// CIR-FAST-NOT: cir.fmuladd
+
+// CIR-ON-LABEL: cir.func {{.*}}@same_stmt
+// CIR-ON: cir.fmuladd
+// CIR-ON-NOT: #cir.fastmath
+
+// CIR-OFF-LABEL: cir.func {{.*}}@same_stmt
+// CIR-OFF: cir.fmul {{.*}} : !cir.float
+// CIR-OFF: cir.fadd {{.*}} : !cir.float
+// CIR-OFF-NOT: #cir.fastmath
+// CIR-OFF-NOT: cir.fmuladd
+
+// LLVM-FAST-LABEL: @same_stmt
+// LLVM-FAST: fmul contract float
+// LLVM-FAST: fadd contract float
+// LLVM-FAST-NOT: @llvm.fmuladd
+
+float across_stmt(float a, float b, float c) {
+ float t = a * b;
+ return t + c;
+}
+// CIR-FAST-LABEL: cir.func {{.*}}@across_stmt
+// CIR-FAST: cir.fmul {{.*}} : !cir.float {fastmath = #cir.fastmath<contract>}
+// CIR-FAST: cir.fadd {{.*}} : !cir.float {fastmath = #cir.fastmath<contract>}
+// CIR-FAST-NOT: cir.fmuladd
+
+// CIR-ON-LABEL: cir.func {{.*}}@across_stmt
+// CIR-ON: cir.fmul {{.*}} : !cir.float
+// CIR-ON: cir.fadd {{.*}} : !cir.float
+// CIR-ON-NOT: cir.fmuladd
+// CIR-ON-NOT: #cir.fastmath
+
+// LLVM-FAST-LABEL: @across_stmt
+// LLVM-FAST: fmul contract float
+// LLVM-FAST: fadd contract float
+
+float sub_stmt(float a, float b, float c) { return a * b - c; }
+// CIR-FAST-LABEL: cir.func {{.*}}@sub_stmt
+// CIR-FAST: cir.fmul {{.*}} : !cir.float {fastmath = #cir.fastmath<contract>}
+// CIR-FAST: cir.fsub {{.*}} : !cir.float {fastmath = #cir.fastmath<contract>}
+// CIR-FAST-NOT: cir.fmuladd
+
+// CIR-ON-LABEL: cir.func {{.*}}@sub_stmt
+// CIR-ON: cir.fmuladd
+// CIR-ON-NOT: #cir.fastmath
+
+// CIR-OFF-LABEL: cir.func {{.*}}@sub_stmt
+// CIR-OFF: cir.fmul {{.*}} : !cir.float
+// CIR-OFF: cir.fsub {{.*}} : !cir.float
+// CIR-OFF-NOT: #cir.fastmath
+// CIR-OFF-NOT: cir.fmuladd
+
+// LLVM-FAST-LABEL: @sub_stmt
+// LLVM-FAST: fmul contract float
+// LLVM-FAST: fsub contract float
+
+float neg(float a) { return -a; }
+// CIR-FAST-LABEL: cir.func {{.*}}@neg
+// CIR-FAST: cir.fneg {{.*}} : !cir.float {fastmath = #cir.fastmath<contract>}
+// CIR-ON-LABEL: cir.func {{.*}}@neg
+// CIR-ON: cir.fneg {{.*}} : !cir.float
+// CIR-ON-NOT: #cir.fastmath
+// CIR-OFF-LABEL: cir.func {{.*}}@neg
+// CIR-OFF: cir.fneg {{.*}} : !cir.float
+// CIR-OFF-NOT: #cir.fastmath
+// LLVM-FAST-LABEL: @neg
+// LLVM-FAST: fneg contract float
+
+int cmp(float a, float b) { return a < b; }
+// CIR-FAST-LABEL: cir.func {{.*}}@cmp
+// CIR-FAST: cir.cmp lt {{.*}} : !cir.float {fastmath = #cir.fastmath<contract>}
+// CIR-ON-LABEL: cir.func {{.*}}@cmp
+// CIR-ON: cir.cmp lt {{.*}} : !cir.float
+// CIR-ON-NOT: #cir.fastmath
+// CIR-OFF-LABEL: cir.func {{.*}}@cmp
+// CIR-OFF: cir.cmp lt {{.*}} : !cir.float
+// CIR-OFF-NOT: #cir.fastmath
+// LLVM-FAST-LABEL: @cmp
+// LLVM-FAST: fcmp contract olt float
+
+float rem(float a, float b) { return __builtin_fmodf(a, b); }
+// CIR-FAST-LABEL: cir.func {{.*}}@rem
+// CIR-FAST: cir.fmod {{.*}} : !cir.float {fastmath = #cir.fastmath<contract>}
+// CIR-ON-LABEL: cir.func {{.*}}@rem
+// CIR-ON: cir.fmod {{.*}} : !cir.float
+// CIR-ON-NOT: #cir.fastmath
+// CIR-OFF-LABEL: cir.func {{.*}}@rem
+// CIR-OFF: cir.fmod {{.*}} : !cir.float
+// CIR-OFF-NOT: #cir.fastmath
+// LLVM-FAST-LABEL: @rem
+// LLVM-FAST: frem contract float
+
+float sq(float a) { return __builtin_sqrtf(a); }
+// CIR-FAST-LABEL: cir.func {{.*}}@sq
+// CIR-FAST: cir.sqrt {{.*}} : !cir.float {fastmath = #cir.fastmath<contract>}
+// CIR-ON-LABEL: cir.func {{.*}}@sq
+// CIR-ON: cir.sqrt {{.*}} : !cir.float
+// CIR-ON-NOT: #cir.fastmath
+// CIR-OFF-LABEL: cir.func {{.*}}@sq
+// CIR-OFF: cir.sqrt {{.*}} : !cir.float
+// CIR-OFF-NOT: #cir.fastmath
+// LLVM-FAST-LABEL: @sq
+// LLVM-FAST: call contract float @llvm.sqrt.f32
+
+float absf(float a) { return __builtin_fabsf(a); }
+// CIR-FAST-LABEL: cir.func {{.*}}@absf
+// CIR-FAST: cir.fabs {{.*}} : !cir.float {fastmath = #cir.fastmath<contract>}
+// CIR-ON-LABEL: cir.func {{.*}}@absf
+// CIR-ON: cir.fabs {{.*}} : !cir.float
+// CIR-ON-NOT: #cir.fastmath
+// CIR-OFF-LABEL: cir.func {{.*}}@absf
+// CIR-OFF: cir.fabs {{.*}} : !cir.float
+// CIR-OFF-NOT: #cir.fastmath
+// LLVM-FAST-LABEL: @absf
+// LLVM-FAST: call contract float @llvm.fabs.f32
diff --git a/clang/test/CIR/CodeGen/fp-math-precision-opts.c b/clang/test/CIR/CodeGen/fp-math-precision-opts.c
index 4c04eb14c6309b..e365d594f33c31 100644
--- a/clang/test/CIR/CodeGen/fp-math-precision-opts.c
+++ b/clang/test/CIR/CodeGen/fp-math-precision-opts.c
@@ -57,10 +57,10 @@ float test_fast(float f) {
// Should produce an intrinsic at -O1
return __builtin_cosf(f);
// ALL: test_fast
- // CIR-ERRNO-O1: cir.cos
- // CIR-NO-ERRNO-O1: cir.cos
- // LLVM-ERRNO-O1: call float @llvm.cos.f32
- // LLVM-NO-ERRNO-O1: call float @llvm.cos.f32
+ // CIR-ERRNO-O1: cir.cos {{.*}} {fastmath = #cir.fastmath<contract>}
+ // CIR-NO-ERRNO-O1: cir.cos {{.*}} {fastmath = #cir.fastmath<contract>}
+ // LLVM-ERRNO-O1: call contract float @llvm.cos.f32
+ // LLVM-NO-ERRNO-O1: call contract float @llvm.cos.f32
// OGCG-ERRNO-O1: call {{.*}} float @llvm.cos.f32
// OGCG-NO-ERRNO-O1: call {{.*}} float @llvm.cos.f32
}
diff --git a/clang/test/CIR/CodeGenCUDA/fp-contract.cu b/clang/test/CIR/CodeGenCUDA/fp-contract.cu
new file mode 100644
index 00000000000000..5b993e358f086f
--- /dev/null
+++ b/clang/test/CIR/CodeGenCUDA/fp-contract.cu
@@ -0,0 +1,57 @@
+// CUDA's default contract mode is -ffp-contract=fast. CIR records that as
+// `contract` on fmul/fadd, not as cir.fmuladd. The lowered LLVM IR must carry
+// the same flag so a later Standard-fusion backend still emits FMA.
+
+// RUN: %clang_cc1 -triple nvptx64-nvidia-cuda -x cuda -fcuda-is-device \
+// RUN: -fclangir -emit-cir %s -o %t.cir
+// RUN: FileCheck --input-file=%t.cir %s -check-prefix=CIR-FAST
+// RUN: %clang_cc1 -triple nvptx64-nvidia-cuda -x cuda -fcuda-is-device \
+// RUN: -fclangir -emit-llvm %s -o %t.ll
+// RUN: FileCheck --input-file=%t.ll %s -check-prefix=LLVM-FAST
+
+// RUN: %clang_cc1 -triple nvptx64-nvidia-cuda -x cuda -fcuda-is-device \
+// RUN: -ffp-contract=on -fclangir -emit-cir %s -o %t-on.cir
+// RUN: FileCheck --input-file=%t-on.cir %s -check-prefix=CIR-ON
+
+// RUN: %clang_cc1 -triple nvptx64-nvidia-cuda -x cuda -fcuda-is-device \
+// RUN: -emit-llvm %s -o %t-og.ll
+// RUN: FileCheck --input-file=%t-og.ll %s -check-prefix=LLVM-FAST
+
+#include "Inputs/cuda.h"
+
+__host__ __device__ float same_stmt(float a, float b, float c) {
+ return a * b + c;
+}
+// CIR-FAST-LABEL: cir.func {{.*}}@_Z9same_stmtfff
+// CIR-FAST: cir.fmul {{.*}} : !cir.float {fastmath = #cir.fastmath<contract>}
+// CIR-FAST: cir.fadd {{.*}} : !cir.float {fastmath = #cir.fastmath<contract>}
+// CIR-FAST-NOT: cir.fmuladd
+
+// CIR-ON-LABEL: cir.func {{.*}}@_Z9same_stmtfff
+// CIR-ON: cir.fmuladd
+// CIR-ON-NOT: #cir.fastmath
+
+// LLVM-FAST-LABEL: @_Z9same_stmtfff
+// LLVM-FAST: fmul contract float
+// LLVM-FAST: fadd contract float
+// LLVM-FAST-NOT: @llvm.fmuladd
+
+__host__ __device__ float across_stmt(float a, float b, float c) {
+ float t = a * b;
+ return t + c;
+}
+// CIR-FAST-LABEL: cir.func {{.*}}@_Z11across_stmtfff
+// CIR-FAST: cir.fmul {{.*}} : !cir.float {fastmath = #cir.fastmath<contract>}
+// CIR-FAST: cir.fadd {{.*}} : !cir.float {fastmath = #cir.fastmath<contract>}
+// CIR-FAST-NOT: cir.fmuladd
+
+// CIR-ON-LABEL: cir.func {{.*}}@_Z11across_stmtfff
+// CIR-ON: cir.fmul {{.*}} : !cir.float
+// CIR-ON: cir.fadd {{.*}} : !cir.float
+// CIR-ON-NOT: cir.fmuladd
+// CIR-ON-NOT: #cir.fastmath
+// CIR-ON: cir.return
+
+// LLVM-FAST-LABEL: @_Z11across_stmtfff
+// LLVM-FAST: fmul contract float
+// LLVM-FAST: fadd contract float
diff --git a/clang/test/CIR/Lowering/fastmath-contract.cir b/clang/test/CIR/Lowering/fastmath-contract.cir
new file mode 100644
index 00000000000000..0ece658a12f8d1
--- /dev/null
+++ b/clang/test/CIR/Lowering/fastmath-contract.cir
@@ -0,0 +1,47 @@
+// RUN: cir-opt %s --verify-roundtrip | FileCheck %s -check-prefix=CIR
+// RUN: cir-opt %s -cir-to-llvm -o - | FileCheck %s -check-prefix=MLIR
+// RUN: cir-translate %s -cir-to-llvmir --target nvptx64-nvidia-cuda --disable-cc-lowering | FileCheck %s -check-prefix=LLVM
+
+!f32 = !cir.float
+
+module {
+ cir.func @contract(%a: !f32, %b: !f32, %c: !f32) -> !f32 {
+ %m = cir.fmul %a, %b : !f32 {fastmath = #cir.fastmath<contract>}
+ %n = cir.fneg %c : !f32 {fastmath = #cir.fastmath<contract>}
+ %s = cir.fadd %m, %n : !f32 {fastmath = #cir.fastmath<contract>}
+ cir.return %s : !f32
+ }
+
+ cir.func @plain(%a: !f32, %b: !f32, %c: !f32) -> !f32 {
+ %m = cir.fmul %a, %b : !f32
+ %s = cir.fadd %m, %c : !f32
+ cir.return %s : !f32
+ }
+}
+
+// CIR-LABEL: cir.func @contract
+// CIR: cir.fmul {{.*}} {fastmath = #cir.fastmath<contract>}
+// CIR: cir.fneg {{.*}} {fastmath = #cir.fastmath<contract>}
+// CIR: cir.fadd {{.*}} {fastmath = #cir.fastmath<contract>}
+// CIR-LABEL: cir.func @plain
+// CIR: cir.fmul {{.*}} : !cir.float
+// CIR-NOT: #cir.fastmath
+// CIR: cir.fadd {{.*}} : !cir.float
+
+// MLIR-LABEL: llvm.func @contract
+// MLIR: llvm.fmul {{.*}} {fastmathFlags = #llvm.fastmath<contract>}
+// MLIR: llvm.fneg {{.*}} {fastmathFlags = #llvm.fastmath<contract>}
+// MLIR: llvm.fadd {{.*}} {fastmathFlags = #llvm.fastmath<contract>}
+// MLIR-LABEL: llvm.func @plain
+// MLIR: llvm.fmul {{.*}} : f32
+// MLIR-NOT: fastmathFlags
+// MLIR: llvm.fadd {{.*}} : f32
+
+// LLVM-LABEL: @contract
+// LLVM: fmul contract float
+// LLVM: fneg contract float
+// LLVM: fadd contract float
+// LLVM-LABEL: @plain
+// LLVM: fmul float
+// LLVM-NOT: contract
+// LLVM: fadd float
diff --git a/clang/utils/TableGen/CIRLoweringEmitter.cpp b/clang/utils/TableGen/CIRLoweringEmitter.cpp
index 756b6265b90f44..978f0484627b58 100644
--- a/clang/utils/TableGen/CIRLoweringEmitter.cpp
+++ b/clang/utils/TableGen/CIRLoweringEmitter.cpp
@@ -151,7 +151,8 @@ void GenerateLLVMLoweringPattern(
llvm::StringRef OpName, llvm::StringRef PatternName, bool IsRecursive,
llvm::StringRef ExtraDecl, const Record *CustomCtorRec,
llvm::StringRef LLVMOp, llvm::StringRef ConstrainedLLVMIntrinsic,
- bool ConstrainedHasRoundingMode, bool HasZeroResult) {
+ bool ConstrainedHasRoundingMode, bool PropagateFastMathFlags,
+ bool HasZeroResult) {
std::optional<CustomLoweringCtor> CustomCtor =
parseCustomLoweringCtor(CustomCtorRec);
std::string CodeBuffer;
@@ -229,6 +230,15 @@ void GenerateLLVMLoweringPattern(
Code << " rewriter.replaceOpWithNewOp<mlir::LLVM::" << LLVMOp
<< ">(op, mlir::TypeRange{}, adaptor.getOperands(), " << Properties
<< ");\n";
+ } else if (PropagateFastMathFlags) {
+ Code << " mlir::Type resTy = "
+ "typeConverter->convertType(op.getType());\n";
+ Code << " auto lowered = mlir::LLVM::" << LLVMOp
+ << "::create(rewriter, op.getLoc(), mlir::TypeRange{resTy}, "
+ "adaptor.getOperands(), "
+ << Properties << ");\n";
+ Code << " propagateFastMathFlags(op, lowered);\n";
+ Code << " rewriter.replaceOp(op, lowered.getResult());\n";
} else {
Code << " mlir::Type resTy = "
"typeConverter->convertType(op.getType());\n";
@@ -277,6 +287,8 @@ void Generate(const Record *OpRecord) {
OpRecord->getValueAsString("constrainedLLVMIntrinsic");
bool ConstrainedHasRoundingMode =
OpRecord->getValueAsBit("constrainedLLVMIntrinsicHasRoundingMode");
+ bool PropagateFastMathFlags =
+ OpRecord->getValueAsBit("propagateFastMathFlags");
if (!LLVMOp.empty() && CustomCtor)
PrintFatalError(OpRecord->getLoc(),
@@ -294,7 +306,8 @@ void Generate(const Record *OpRecord) {
bool IsZeroResult = ResultsDag->getNumArgs() == 0;
GenerateLLVMLoweringPattern(OpName, PatternName, IsRecursive, ExtraDecl,
CustomCtor, LLVMOp, ConstrainedLLVMIntrinsic,
- ConstrainedHasRoundingMode, IsZeroResult);
+ ConstrainedHasRoundingMode,
+ PropagateFastMathFlags, IsZeroResult);
// Only automatically register patterns that use the default constructor.
// Patterns with a custom constructor must be manually registered by the
// lowering pass.
>From e72a7a8cb885fb2007cdd734e35af96364052198 Mon Sep 17 00:00:00 2001
From: David Rivera <davidriverg at gmail.com>
Date: Thu, 24 Sep 2026 21:10:48 -0400
Subject: [PATCH 2/5] [CIR] clang-format the fp-contract changes
Co-authored-by: Cursor <cursoragent at cursor.com>
---
clang/lib/CIR/CodeGen/CIRGenFunction.cpp | 3 +-
.../CIR/Lowering/DirectToLLVM/LowerToLLVM.cpp | 37 +++++++++----------
2 files changed, 19 insertions(+), 21 deletions(-)
diff --git a/clang/lib/CIR/CodeGen/CIRGenFunction.cpp b/clang/lib/CIR/CodeGen/CIRGenFunction.cpp
index bb9ccdedbe5afb..c70b00eae1ec67 100644
--- a/clang/lib/CIR/CodeGen/CIRGenFunction.cpp
+++ b/clang/lib/CIR/CodeGen/CIRGenFunction.cpp
@@ -44,7 +44,8 @@ static bool functionMightHaveBypass(const Stmt *s) {
return false;
}
-static cir::FastMathFlags fastMathFlagsFromFPOptions(clang::FPOptions fpFeatures) {
+static cir::FastMathFlags
+fastMathFlagsFromFPOptions(clang::FPOptions fpFeatures) {
// Other fast-math bits (nnan, ninf, reassoc, ...) are not modeled yet.
// `contract` is the bit `-ffp-contract=fast` needs so a later backend in
// Standard fusion mode can still form an FMA.
diff --git a/clang/lib/CIR/Lowering/DirectToLLVM/LowerToLLVM.cpp b/clang/lib/CIR/Lowering/DirectToLLVM/LowerToLLVM.cpp
index 7db03d4e56fc4a..1bf89658d2402c 100644
--- a/clang/lib/CIR/Lowering/DirectToLLVM/LowerToLLVM.cpp
+++ b/clang/lib/CIR/Lowering/DirectToLLVM/LowerToLLVM.cpp
@@ -551,8 +551,7 @@ static llvm::StringRef getConstrainedExceptMetadata(cir::FenvAttr fenv) {
}
// CIR fast-math bits use the same positions as `mlir::LLVM::FastmathFlags`.
-static mlir::LLVM::FastmathFlags
-toLLVMFastMathFlags(cir::FastMathFlags flags) {
+static mlir::LLVM::FastmathFlags toLLVMFastMathFlags(cir::FastMathFlags flags) {
using CirFlags = cir::FastMathFlags;
using LLVMFlags = mlir::LLVM::FastmathFlags;
static_assert(static_cast<uint32_t>(CirFlags::nnan) ==
@@ -2150,11 +2149,11 @@ mlir::LogicalResult CIRToLLVMFMaxNumOpLowering::matchAndRewrite(
static_cast<uint32_t>(mlir::LLVM::FastmathFlags::nsz) |
static_cast<uint32_t>(readFastMathFlags(op)));
if (cir::FenvAttr fenv = op.getFenvAttr())
- return lowerToConstrainedFPIntrinsic(
- op, adaptor.getOperands(), fenv, resTy, rewriter, "maxnum",
- /*hasRoundingMode=*/false, flags);
- rewriter.replaceOpWithNewOp<mlir::LLVM::MaxNumOp>(
- op, resTy, adaptor.getLhs(), adaptor.getRhs(), flags);
+ return lowerToConstrainedFPIntrinsic(op, adaptor.getOperands(), fenv, resTy,
+ rewriter, "maxnum",
+ /*hasRoundingMode=*/false, flags);
+ rewriter.replaceOpWithNewOp<mlir::LLVM::MaxNumOp>(op, resTy, adaptor.getLhs(),
+ adaptor.getRhs(), flags);
return mlir::success();
}
@@ -2166,11 +2165,11 @@ mlir::LogicalResult CIRToLLVMFMinNumOpLowering::matchAndRewrite(
static_cast<uint32_t>(mlir::LLVM::FastmathFlags::nsz) |
static_cast<uint32_t>(readFastMathFlags(op)));
if (cir::FenvAttr fenv = op.getFenvAttr())
- return lowerToConstrainedFPIntrinsic(
- op, adaptor.getOperands(), fenv, resTy, rewriter, "minnum",
- /*hasRoundingMode=*/false, flags);
- rewriter.replaceOpWithNewOp<mlir::LLVM::MinNumOp>(
- op, resTy, adaptor.getLhs(), adaptor.getRhs(), flags);
+ return lowerToConstrainedFPIntrinsic(op, adaptor.getOperands(), fenv, resTy,
+ rewriter, "minnum",
+ /*hasRoundingMode=*/false, flags);
+ rewriter.replaceOpWithNewOp<mlir::LLVM::MinNumOp>(op, resTy, adaptor.getLhs(),
+ adaptor.getRhs(), flags);
return mlir::success();
}
@@ -3668,12 +3667,10 @@ static bool isSignalingConstrainedFCmp(cir::CmpOpKind kind) {
llvm_unreachable("Unknown CmpOpKind");
}
-static mlir::LLVM::CallIntrinsicOp
-createConstrainedFCmpCall(mlir::ConversionPatternRewriter &rewriter,
- mlir::Location loc, mlir::Value lhs, mlir::Value rhs,
- cir::CmpOpKind kind, cir::FenvAttr fenv,
- mlir::Type llvmResTy,
- mlir::LLVM::FastmathFlags fastmathFlags = {}) {
+static mlir::LLVM::CallIntrinsicOp createConstrainedFCmpCall(
+ mlir::ConversionPatternRewriter &rewriter, mlir::Location loc,
+ mlir::Value lhs, mlir::Value rhs, cir::CmpOpKind kind, cir::FenvAttr fenv,
+ mlir::Type llvmResTy, mlir::LLVM::FastmathFlags fastmathFlags = {}) {
llvm::SmallVector<mlir::Value, 4> callOperands = {
lhs, rhs,
createFenvMetadataValue(rewriter, loc,
@@ -3725,8 +3722,8 @@ mlir::LogicalResult CIRToLLVMCmpOpLowering::matchAndRewrite(
}
mlir::LLVM::FCmpPredicate kind =
convertCmpKindToFCmpPredicate(cmpOp.getKind());
- auto fcmp = mlir::LLVM::FCmpOp::create(
- rewriter, cmpOp.getLoc(), kind, adaptor.getLhs(), adaptor.getRhs());
+ auto fcmp = mlir::LLVM::FCmpOp::create(rewriter, cmpOp.getLoc(), kind,
+ adaptor.getLhs(), adaptor.getRhs());
propagateFastMathFlags(cmpOp, fcmp);
rewriter.replaceOp(cmpOp, fcmp.getResult());
return mlir::success();
>From 6c75084761cf6561685f28cc29f4bc241b7cd966 Mon Sep 17 00:00:00 2001
From: David Rivera <davidriverg at gmail.com>
Date: Thu, 24 Sep 2026 21:53:22 -0400
Subject: [PATCH 3/5] [CIR] Update the fastmath contract lowering test
cir-to-llvm now requires a module triple, and the LLVM dialect prints
fastmath flags as fastmath<contract> rather than an attribute dictionary.
Co-authored-by: Cursor <cursoragent at cursor.com>
---
clang/test/CIR/Lowering/fastmath-contract.cir | 10 +++++-----
1 file changed, 5 insertions(+), 5 deletions(-)
diff --git a/clang/test/CIR/Lowering/fastmath-contract.cir b/clang/test/CIR/Lowering/fastmath-contract.cir
index 0ece658a12f8d1..a75a9391936dbe 100644
--- a/clang/test/CIR/Lowering/fastmath-contract.cir
+++ b/clang/test/CIR/Lowering/fastmath-contract.cir
@@ -4,7 +4,7 @@
!f32 = !cir.float
-module {
+module attributes {cir.triple = "nvptx64-nvidia-cuda"} {
cir.func @contract(%a: !f32, %b: !f32, %c: !f32) -> !f32 {
%m = cir.fmul %a, %b : !f32 {fastmath = #cir.fastmath<contract>}
%n = cir.fneg %c : !f32 {fastmath = #cir.fastmath<contract>}
@@ -29,12 +29,12 @@ module {
// CIR: cir.fadd {{.*}} : !cir.float
// MLIR-LABEL: llvm.func @contract
-// MLIR: llvm.fmul {{.*}} {fastmathFlags = #llvm.fastmath<contract>}
-// MLIR: llvm.fneg {{.*}} {fastmathFlags = #llvm.fastmath<contract>}
-// MLIR: llvm.fadd {{.*}} {fastmathFlags = #llvm.fastmath<contract>}
+// MLIR: llvm.fmul {{.*}} fastmath<contract> : f32
+// MLIR: llvm.fneg {{.*}} fastmath<contract> : f32
+// MLIR: llvm.fadd {{.*}} fastmath<contract> : f32
// MLIR-LABEL: llvm.func @plain
// MLIR: llvm.fmul {{.*}} : f32
-// MLIR-NOT: fastmathFlags
+// MLIR-NOT: fastmath<
// MLIR: llvm.fadd {{.*}} : f32
// LLVM-LABEL: @contract
>From 34fe18494a509601e3d5aa8b2d7e6bd2cef2b569 Mon Sep 17 00:00:00 2001
From: David Rivera <davidriverg at gmail.com>
Date: Fri, 25 Sep 2026 15:42:04 -0400
Subject: [PATCH 4/5] [CIR] Reuse the existing fast-math flags attribute
#224899 already defined cir::FastMathFlags. Drop the duplicate enum and
let FP operation builders read the active fenv and fast-math flags from
the CIR builder.
Co-authored-by: Cursor <cursoragent at cursor.com>
---
.../CIR/Dialect/Builder/CIRBaseBuilder.h | 75 +++++++++++++----
.../include/clang/CIR/Dialect/IR/CIRAttrs.td | 4 +
.../include/clang/CIR/Dialect/IR/CIRDialect.h | 14 ++++
.../clang/CIR/Dialect/IR/CIREnumAttr.td | 30 -------
clang/include/clang/CIR/Dialect/IR/CIROps.td | 55 ++++++------
clang/lib/CIR/CodeGen/CIRGenBuiltin.cpp | 23 ++---
clang/lib/CIR/CodeGen/CIRGenExprScalar.cpp | 3 +-
clang/lib/CIR/Dialect/IR/CIRDialect.cpp | 49 +++++++++++
.../CIR/Lowering/DirectToLLVM/LowerToLLVM.cpp | 83 ++++++-------------
9 files changed, 186 insertions(+), 150 deletions(-)
diff --git a/clang/include/clang/CIR/Dialect/Builder/CIRBaseBuilder.h b/clang/include/clang/CIR/Dialect/Builder/CIRBaseBuilder.h
index e8ec07a66a77a1..fb9bb6eea3e56d 100644
--- a/clang/include/clang/CIR/Dialect/Builder/CIRBaseBuilder.h
+++ b/clang/include/clang/CIR/Dialect/Builder/CIRBaseBuilder.h
@@ -66,8 +66,15 @@ class CIRBaseBuilderTy : public mlir::OpBuilder {
public:
CIRBaseBuilderTy(mlir::MLIRContext &mlirContext)
- : mlir::OpBuilder(&mlirContext) {}
- CIRBaseBuilderTy(mlir::OpBuilder &builder) : mlir::OpBuilder(builder) {}
+ : mlir::OpBuilder(&mlirContext) {
+ registerFPDefaults();
+ }
+ CIRBaseBuilderTy(mlir::OpBuilder &builder) : mlir::OpBuilder(builder) {
+ registerFPDefaults();
+ }
+ CIRBaseBuilderTy(const CIRBaseBuilderTy &other);
+ CIRBaseBuilderTy &operator=(const CIRBaseBuilderTy &other);
+ ~CIRBaseBuilderTy() { cir::unregisterCIRBuilderFPDefaults(this); }
bool isFPConstrained = false;
clang::LangOptions::FPExceptionModeKind defaultConstrainedExcept =
@@ -75,7 +82,8 @@ class CIRBaseBuilderTy : public mlir::OpBuilder {
llvm::RoundingMode defaultConstrainedRounding =
llvm::RoundingMode::NearestTiesToEven;
// Fast-math flags applied to floating-point ops created by this builder.
- // CIRGen currently populates `contract` only.
+ // CIRGen currently populates `contract` only. FP op builders read these via
+ // fenvForBuilder / fastMathForBuilder, so create sites omit the attributes.
cir::FastMathFlags fastMathFlags = cir::FastMathFlags::none;
void setFastMathFlags(cir::FastMathFlags flags) { fastMathFlags = flags; }
@@ -87,6 +95,19 @@ class CIRBaseBuilderTy : public mlir::OpBuilder {
return cir::FastMathFlagsAttr::get(getContext(), fastMathFlags);
}
+private:
+ static cir::FenvAttr fetchFenv(void *self) {
+ return static_cast<CIRBaseBuilderTy *>(self)->getConstrainedFPAttr();
+ }
+ static cir::FastMathFlagsAttr fetchFastMath(void *self) {
+ return static_cast<CIRBaseBuilderTy *>(self)->getFastMathFlagsAttr();
+ }
+ void registerFPDefaults() {
+ cir::registerCIRBuilderFPDefaults(this, &CIRBaseBuilderTy::fetchFenv,
+ &CIRBaseBuilderTy::fetchFastMath, this);
+ }
+
+public:
mlir::Value getConstAPInt(mlir::Location loc, mlir::Type typ,
const llvm::APInt &val) {
return cir::ConstantOp::create(*this, loc, cir::IntAttr::get(typ, val));
@@ -860,39 +881,34 @@ class CIRBaseBuilderTy : public mlir::OpBuilder {
mlir::Value createFAdd(mlir::Location loc, mlir::Value lhs, mlir::Value rhs) {
assert(!cir::MissingFeatures::metaDataNode());
- // `contract` is applied via getFastMathFlagsAttr(). The other fast-math
- // bits are still unimplemented.
+ // `contract` is applied by FAddOp's builder. The other fast-math bits are
+ // still unimplemented.
assert(!cir::MissingFeatures::fastMathFlags());
- return cir::FAddOp::create(*this, loc, lhs, rhs, getConstrainedFPAttr(),
- getFastMathFlagsAttr());
+ return cir::FAddOp::create(*this, loc, lhs, rhs);
}
mlir::Value createFSub(mlir::Location loc, mlir::Value lhs, mlir::Value rhs) {
assert(!cir::MissingFeatures::metaDataNode());
assert(!cir::MissingFeatures::fastMathFlags());
- return cir::FSubOp::create(*this, loc, lhs, rhs, getConstrainedFPAttr(),
- getFastMathFlagsAttr());
+ return cir::FSubOp::create(*this, loc, lhs, rhs);
}
mlir::Value createFMul(mlir::Location loc, mlir::Value lhs, mlir::Value rhs) {
assert(!cir::MissingFeatures::metaDataNode());
assert(!cir::MissingFeatures::fastMathFlags());
- return cir::FMulOp::create(*this, loc, lhs, rhs, getConstrainedFPAttr(),
- getFastMathFlagsAttr());
+ return cir::FMulOp::create(*this, loc, lhs, rhs);
}
mlir::Value createFDiv(mlir::Location loc, mlir::Value lhs, mlir::Value rhs) {
assert(!cir::MissingFeatures::metaDataNode());
assert(!cir::MissingFeatures::fastMathFlags());
- return cir::FDivOp::create(*this, loc, lhs, rhs, getConstrainedFPAttr(),
- getFastMathFlagsAttr());
+ return cir::FDivOp::create(*this, loc, lhs, rhs);
}
mlir::Value createFRem(mlir::Location loc, mlir::Value lhs, mlir::Value rhs) {
assert(!cir::MissingFeatures::metaDataNode());
assert(!cir::MissingFeatures::fastMathFlags());
- return cir::FRemOp::create(*this, loc, lhs, rhs, getConstrainedFPAttr(),
- getFastMathFlagsAttr());
+ return cir::FRemOp::create(*this, loc, lhs, rhs);
}
mlir::Value createFNeg(mlir::Location loc, mlir::Value operand) {
@@ -902,7 +918,7 @@ class CIRBaseBuilderTy : public mlir::OpBuilder {
assert(!cir::MissingFeatures::fastMathFlags());
// fneg does not raise FP exceptions or depend on the rounding mode, so it
// never carries an fenv attribute.
- return cir::FNegOp::create(*this, loc, operand, getFastMathFlagsAttr());
+ return cir::FNegOp::create(*this, loc, operand);
}
mlir::Value createXor(mlir::Location loc, mlir::Value lhs, mlir::Value rhs) {
@@ -936,10 +952,13 @@ class CIRBaseBuilderTy : public mlir::OpBuilder {
VectorType integralVecTy =
cir::VectorType::get(integralTy, vecCast.getSize());
cir::FenvAttr fenv;
- if (cir::isFPOrVectorOfFPType(lhs.getType()))
+ cir::FastMathFlagsAttr fastmath;
+ if (cir::isFPOrVectorOfFPType(lhs.getType())) {
fenv = getConstrainedFPAttr();
+ fastmath = getFastMathFlagsAttr();
+ }
return cir::VecCmpOp::create(*this, loc, integralVecTy, kind, lhs, rhs,
- fenv, getFastMathFlagsAttr());
+ fenv, fastmath);
}
mlir::Value createIsNaN(mlir::Location loc, mlir::Value operand) {
@@ -1092,6 +1111,26 @@ class CIRBaseBuilderTy : public mlir::OpBuilder {
}
};
+inline CIRBaseBuilderTy::CIRBaseBuilderTy(const CIRBaseBuilderTy &other)
+ : mlir::OpBuilder(other), isFPConstrained(other.isFPConstrained),
+ defaultConstrainedExcept(other.defaultConstrainedExcept),
+ defaultConstrainedRounding(other.defaultConstrainedRounding),
+ fastMathFlags(other.fastMathFlags) {
+ registerFPDefaults();
+}
+
+inline CIRBaseBuilderTy &
+CIRBaseBuilderTy::operator=(const CIRBaseBuilderTy &other) {
+ if (this == &other)
+ return *this;
+ mlir::OpBuilder::operator=(other);
+ isFPConstrained = other.isFPConstrained;
+ defaultConstrainedExcept = other.defaultConstrainedExcept;
+ defaultConstrainedRounding = other.defaultConstrainedRounding;
+ fastMathFlags = other.fastMathFlags;
+ return *this;
+}
+
} // namespace cir
#endif
diff --git a/clang/include/clang/CIR/Dialect/IR/CIRAttrs.td b/clang/include/clang/CIR/Dialect/IR/CIRAttrs.td
index 04dc8bb178cb87..f2ddef865c8b72 100644
--- a/clang/include/clang/CIR/Dialect/IR/CIRAttrs.td
+++ b/clang/include/clang/CIR/Dialect/IR/CIRAttrs.td
@@ -973,6 +973,10 @@ def CIR_FastMathFlags : CIR_I32BitEnum<
Describes fast-math flags for CIR operations. This attribute is shared by
operations with floating-point semantics and is not specific to LLVM intrinsic
calls.
+
+ `contract` allows a backend in Standard fusion mode to form an FMA across
+ statements. It is what `-ffp-contract=fast` records. `-ffp-contract=on` is
+ represented separately by `cir.fmuladd`.
}];
let separator = ", ";
diff --git a/clang/include/clang/CIR/Dialect/IR/CIRDialect.h b/clang/include/clang/CIR/Dialect/IR/CIRDialect.h
index d9018e93628b69..238dcded0a0287 100644
--- a/clang/include/clang/CIR/Dialect/IR/CIRDialect.h
+++ b/clang/include/clang/CIR/Dialect/IR/CIRDialect.h
@@ -99,6 +99,20 @@ RecordLayoutAttr getRecordLayout(mlir::ModuleOp mod, mlir::StringAttr name);
/// Same lookup as getRecordLayout, but returns a null attribute instead of
/// asserting when the record has no layout entry.
RecordLayoutAttr tryGetRecordLayout(mlir::ModuleOp mod, mlir::StringAttr name);
+
+/// Active constrained-FP and fast-math state for a `CIRBaseBuilderTy`.
+/// Both are null when `builder` did not register any.
+///
+/// ODS builders only receive `mlir::OpBuilder &`, so floating-point ops call
+/// these instead of taking the attributes at every create site.
+FenvAttr fenvForBuilder(mlir::OpBuilder &builder);
+FastMathFlagsAttr fastMathForBuilder(mlir::OpBuilder &builder);
+
+void registerCIRBuilderFPDefaults(mlir::OpBuilder *builder,
+ FenvAttr (*fenv)(void *),
+ FastMathFlagsAttr (*fastMath)(void *),
+ void *self);
+void unregisterCIRBuilderFPDefaults(mlir::OpBuilder *builder);
} // namespace cir
// TableGen'erated files for MLIR dialects require that a macro be defined when
diff --git a/clang/include/clang/CIR/Dialect/IR/CIREnumAttr.td b/clang/include/clang/CIR/Dialect/IR/CIREnumAttr.td
index ecd981593ec047..3ee06412d8a90f 100644
--- a/clang/include/clang/CIR/Dialect/IR/CIREnumAttr.td
+++ b/clang/include/clang/CIR/Dialect/IR/CIREnumAttr.td
@@ -47,36 +47,6 @@ class CIR_EnumAttr<EnumInfo info, string name = "", list<Trait> traits = []>
let assemblyFormat = "`<` $value `>`";
}
-// Bit positions match `mlir::LLVM::FastmathFlags`. Only `contract` is set by
-// CIRGen today; the other bits exist so the attribute can round-trip.
-def CIR_FastMathFlags : CIR_I32BitEnum<
- "FastMathFlags", "fast-math flags", [
- I32BitEnumCaseNone<"none">,
- I32BitEnumCaseBit<"nnan", 0>,
- I32BitEnumCaseBit<"ninf", 1>,
- I32BitEnumCaseBit<"nsz", 2>,
- I32BitEnumCaseBit<"arcp", 3>,
- I32BitEnumCaseBit<"contract", 4>,
- I32BitEnumCaseBit<"afn", 5>,
- I32BitEnumCaseBit<"reassoc", 6>
-]> {
- let description = [{
- Per-operation fast-math flags. These are the LLVM fast-math bits, carried
- on CIR floating-point operations and lowered onto the corresponding LLVM
- dialect operation's `fastmathFlags`.
-
- `contract` allows the backend to fuse this operation with another
- floating-point operation, including across statements. It is what
- `-ffp-contract=fast` records. `-ffp-contract=on` is represented separately
- by `cir.fmuladd`.
- }];
- let separator = ", ";
-}
-
-def CIR_FastMathFlagsAttr : CIR_EnumAttr<CIR_FastMathFlags, "fastmath"> {
- let summary = "Fast-math flags for a floating-point operation";
-}
-
def CIR_LangAddressSpace : CIR_I32Enum<
"LangAddressSpace", "language address space kind", [
I32EnumCase<"Default", 0, "default">,
diff --git a/clang/include/clang/CIR/Dialect/IR/CIROps.td b/clang/include/clang/CIR/Dialect/IR/CIROps.td
index 9d222de1cd8a66..14ee7957549c50 100644
--- a/clang/include/clang/CIR/Dialect/IR/CIROps.td
+++ b/clang/include/clang/CIR/Dialect/IR/CIROps.td
@@ -97,6 +97,24 @@ class LoweringBuilders<dag p> {
dag dagParams = p;
}
+// Shorter builder for floating-point ops. Omitted `fenv` and `fastmath`
+// attributes come from the active CIR builder (null for any other builder).
+class CIR_DefaultFPAttrsBuilder<dag params, string forwarded>
+ : OpBuilder<params> {
+ let body = !strconcat(
+ "build($_builder, $_state, ", forwarded,
+ ", ::cir::fenvForBuilder($_builder), ::cir::fastMathForBuilder($_builder));");
+}
+
+// Same as CIR_DefaultFPAttrsBuilder for ops that carry fast-math flags but
+// not an fenv attribute.
+class CIR_DefaultFastMathBuilder<dag params, string forwarded>
+ : OpBuilder<params> {
+ let body = !strconcat(
+ "build($_builder, $_state, ", forwarded,
+ ", ::cir::fastMathForBuilder($_builder));");
+}
+
class CIR_Op<string mnemonic, list<Trait> traits = []> :
Op<CIR_Dialect, mnemonic, traits>, LLVMLoweringInfo {
// Should we generate an ABI lowering pattern for this op?
@@ -2180,9 +2198,7 @@ def CIR_FNegOp : CIR_UnaryOp<"fneg", CIR_AnyFloatOrVecOfFloatType> {
(ins OptionalAttr<CIR_FastMathFlagsAttr>:$fastmath));
let builders = [
- OpBuilder<(ins "mlir::Value":$input), [{
- build($_builder, $_state, input, cir::FastMathFlagsAttr{});
- }]>
+ CIR_DefaultFastMathBuilder<(ins "mlir::Value":$input), "input">
];
let hasFolder = 1;
@@ -2975,10 +2991,8 @@ class CIR_FPBinaryOp<string mnemonic, list<Trait> traits = []>
let constrainedLLVMIntrinsic = mnemonic;
let builders = [
- OpBuilder<(ins "mlir::Value":$lhs, "mlir::Value":$rhs), [{
- build($_builder, $_state, lhs, rhs, cir::FenvAttr{},
- cir::FastMathFlagsAttr{});
- }]>
+ CIR_DefaultFPAttrsBuilder<(ins "mlir::Value":$lhs, "mlir::Value":$rhs),
+ "lhs, rhs">
];
}
@@ -7482,10 +7496,7 @@ class CIR_UnaryFPToFPBuiltinOp<string mnemonic, string llvmOpName>
let assemblyFormat = "$src `:` type($src) attr-dict";
let builders = [
- OpBuilder<(ins "mlir::Value":$src), [{
- build($_builder, $_state, src, cir::FenvAttr{},
- cir::FastMathFlagsAttr{});
- }]>
+ CIR_DefaultFPAttrsBuilder<(ins "mlir::Value":$src), "src">
];
let llvmOp = llvmOpName;
@@ -7814,10 +7825,8 @@ class CIR_UnaryFPToIntBuiltinOp<string mnemonic, string llvmOpName>
}];
let builders = [
- OpBuilder<(ins "mlir::Type":$result, "mlir::Value":$src), [{
- build($_builder, $_state, result, src, cir::FenvAttr{},
- cir::FastMathFlagsAttr{});
- }]>
+ CIR_DefaultFPAttrsBuilder<(ins "mlir::Type":$result, "mlir::Value":$src),
+ "result, src">
];
let llvmOp = llvmOpName;
@@ -7882,11 +7891,9 @@ class CIR_BinaryFPToFPBuiltinOp<string mnemonic, string llvmOpName>
}];
let builders = [
- OpBuilder<(ins "mlir::Type":$result, "mlir::Value":$lhs,
- "mlir::Value":$rhs), [{
- build($_builder, $_state, result, lhs, rhs, cir::FenvAttr{},
- cir::FastMathFlagsAttr{});
- }]>
+ CIR_DefaultFPAttrsBuilder<(ins "mlir::Type":$result, "mlir::Value":$lhs,
+ "mlir::Value":$rhs),
+ "result, lhs, rhs">
];
let llvmOp = llvmOpName;
@@ -8035,11 +8042,9 @@ class CIR_TernaryFPToFPBuiltinOp<string mnemonic, string llvmOpName>
let assemblyFormat = "$a `,` $b `,` $c `:` type($a) attr-dict";
let builders = [
- OpBuilder<(ins "mlir::Type":$result, "mlir::Value":$a, "mlir::Value":$b,
- "mlir::Value":$c), [{
- build($_builder, $_state, result, a, b, c, cir::FenvAttr{},
- cir::FastMathFlagsAttr{});
- }]>
+ CIR_DefaultFPAttrsBuilder<(ins "mlir::Type":$result, "mlir::Value":$a,
+ "mlir::Value":$b, "mlir::Value":$c),
+ "result, a, b, c">
];
let llvmOp = llvmOpName;
diff --git a/clang/lib/CIR/CodeGen/CIRGenBuiltin.cpp b/clang/lib/CIR/CodeGen/CIRGenBuiltin.cpp
index 41c09cb43fcb85..fd36a2f8ed2693 100644
--- a/clang/lib/CIR/CodeGen/CIRGenBuiltin.cpp
+++ b/clang/lib/CIR/CodeGen/CIRGenBuiltin.cpp
@@ -557,9 +557,7 @@ static RValue emitUnaryMaybeConstrainedFPBuiltin(CIRGenFunction &cgf,
CIRGenFunction::CIRGenFPOptionsRAII FPOptsRAII(cgf, &e);
- auto call = Operation::create(cgf.getBuilder(), arg.getLoc(), arg.getType(),
- arg, cgf.getBuilder().getConstrainedFPAttr(),
- cgf.getBuilder().getFastMathFlagsAttr());
+ auto call = Operation::create(cgf.getBuilder(), arg.getLoc(), arg);
return RValue::get(call->getResult(0));
}
@@ -567,9 +565,7 @@ template <class Operation>
static RValue emitUnaryFPBuiltin(CIRGenFunction &cgf, const CallExpr &e) {
mlir::Value arg = cgf.emitScalarExpr(e.getArg(0));
CIRGenFunction::CIRGenFPOptionsRAII FPOptsRAII(cgf, &e);
- auto call = Operation::create(cgf.getBuilder(), arg.getLoc(), arg.getType(),
- arg, cir::FenvAttr{},
- cgf.getBuilder().getFastMathFlagsAttr());
+ auto call = Operation::create(cgf.getBuilder(), arg.getLoc(), arg);
return RValue::get(call->getResult(0));
}
@@ -581,9 +577,7 @@ static RValue emitUnaryMaybeConstrainedFPToIntBuiltin(CIRGenFunction &cgf,
CIRGenFunction::CIRGenFPOptionsRAII FPOptsRAII(cgf, &e);
- auto call = Op::create(cgf.getBuilder(), src.getLoc(), resultType, src,
- cgf.getBuilder().getConstrainedFPAttr(),
- cgf.getBuilder().getFastMathFlagsAttr());
+ auto call = Op::create(cgf.getBuilder(), src.getLoc(), resultType, src);
return RValue::get(call->getResult(0));
}
@@ -595,8 +589,7 @@ static RValue emitBinaryFPBuiltin(CIRGenFunction &cgf, const CallExpr &e) {
CIRGenFunction::CIRGenFPOptionsRAII FPOptsRAII(cgf, &e);
mlir::Location loc = cgf.getLoc(e.getExprLoc());
mlir::Type ty = cgf.convertType(e.getType());
- auto call = Op::create(cgf.getBuilder(), loc, ty, arg0, arg1, cir::FenvAttr{},
- cgf.getBuilder().getFastMathFlagsAttr());
+ auto call = Op::create(cgf.getBuilder(), loc, ty, arg0, arg1);
return RValue::get(call->getResult(0));
}
@@ -627,9 +620,7 @@ static RValue emitTernaryMaybeConstrainedFPBuiltin(CIRGenFunction &cgf,
mlir::Location loc = cgf.getLoc(e.getExprLoc());
mlir::Type ty = cgf.convertType(e.getType());
- auto call = Op::create(cgf.getBuilder(), loc, ty, arg0, arg1, arg2,
- cgf.getBuilder().getConstrainedFPAttr(),
- cgf.getBuilder().getFastMathFlagsAttr());
+ auto call = Op::create(cgf.getBuilder(), loc, ty, arg0, arg1, arg2);
return RValue::get(call->getResult(0));
}
@@ -644,9 +635,7 @@ static mlir::Value emitBinaryMaybeConstrainedFPBuiltin(CIRGenFunction &cgf,
mlir::Location loc = cgf.getLoc(e.getExprLoc());
mlir::Type ty = cgf.convertType(e.getType());
- auto call = Op::create(cgf.getBuilder(), loc, ty, arg0, arg1,
- cgf.getBuilder().getConstrainedFPAttr(),
- cgf.getBuilder().getFastMathFlagsAttr());
+ auto call = Op::create(cgf.getBuilder(), loc, ty, arg0, arg1);
return call->getResult(0);
}
diff --git a/clang/lib/CIR/CodeGen/CIRGenExprScalar.cpp b/clang/lib/CIR/CodeGen/CIRGenExprScalar.cpp
index cc29109883b77b..539fd36ca3fced 100644
--- a/clang/lib/CIR/CodeGen/CIRGenExprScalar.cpp
+++ b/clang/lib/CIR/CodeGen/CIRGenExprScalar.cpp
@@ -901,8 +901,7 @@ class ScalarExprEmitter : public StmtVisitor<ScalarExprEmitter, mlir::Value> {
if (cir::isFPOrVectorOfFPType(operand.getType())) {
CIRGenFunction::CIRGenFPOptionsRAII FPOptsRAII(cgf, e);
- return builder.createOrFold<cir::FNegOp>(loc, operand,
- builder.getFastMathFlagsAttr());
+ return builder.createOrFold<cir::FNegOp>(loc, operand);
}
// TODO(cir): We might have to change this to support overflow trapping.
diff --git a/clang/lib/CIR/Dialect/IR/CIRDialect.cpp b/clang/lib/CIR/Dialect/IR/CIRDialect.cpp
index d948ed9608d352..9739c8de5d4708 100644
--- a/clang/lib/CIR/Dialect/IR/CIRDialect.cpp
+++ b/clang/lib/CIR/Dialect/IR/CIRDialect.cpp
@@ -29,14 +29,63 @@
#include "clang/CIR/Dialect/IR/CIROpsDialect.cpp.inc"
#include "clang/CIR/Dialect/IR/CIROpsEnums.cpp.inc"
#include "clang/CIR/MissingFeatures.h"
+#include "llvm/ADT/DenseMap.h"
#include "llvm/ADT/SetOperations.h"
#include "llvm/ADT/SmallSet.h"
#include "llvm/ADT/TypeSwitch.h"
#include "llvm/Support/LogicalResult.h"
+#include "llvm/Support/Mutex.h"
using namespace mlir;
using namespace cir;
+namespace {
+struct CIRBuilderFPDefaults {
+ FenvAttr (*fenv)(void *);
+ FastMathFlagsAttr (*fastMath)(void *);
+ void *self;
+};
+
+llvm::sys::SmartMutex<true> &cirBuilderFPMutex() {
+ static llvm::sys::SmartMutex<true> mutex;
+ return mutex;
+}
+
+llvm::DenseMap<mlir::OpBuilder *, CIRBuilderFPDefaults> &cirBuilderFPMap() {
+ static llvm::DenseMap<mlir::OpBuilder *, CIRBuilderFPDefaults> map;
+ return map;
+}
+} // namespace
+
+void cir::registerCIRBuilderFPDefaults(mlir::OpBuilder *builder,
+ FenvAttr (*fenv)(void *),
+ FastMathFlagsAttr (*fastMath)(void *),
+ void *self) {
+ llvm::sys::SmartScopedLock<true> lock(cirBuilderFPMutex());
+ cirBuilderFPMap()[builder] = {fenv, fastMath, self};
+}
+
+void cir::unregisterCIRBuilderFPDefaults(mlir::OpBuilder *builder) {
+ llvm::sys::SmartScopedLock<true> lock(cirBuilderFPMutex());
+ cirBuilderFPMap().erase(builder);
+}
+
+FenvAttr cir::fenvForBuilder(mlir::OpBuilder &builder) {
+ llvm::sys::SmartScopedLock<true> lock(cirBuilderFPMutex());
+ auto it = cirBuilderFPMap().find(&builder);
+ if (it == cirBuilderFPMap().end())
+ return {};
+ return it->second.fenv(it->second.self);
+}
+
+FastMathFlagsAttr cir::fastMathForBuilder(mlir::OpBuilder &builder) {
+ llvm::sys::SmartScopedLock<true> lock(cirBuilderFPMutex());
+ auto it = cirBuilderFPMap().find(&builder);
+ if (it == cirBuilderFPMap().end())
+ return {};
+ return it->second.fastMath(it->second.self);
+}
+
//===----------------------------------------------------------------------===//
// CIR Dialect
//===----------------------------------------------------------------------===//
diff --git a/clang/lib/CIR/Lowering/DirectToLLVM/LowerToLLVM.cpp b/clang/lib/CIR/Lowering/DirectToLLVM/LowerToLLVM.cpp
index 1bf89658d2402c..a5383045b14465 100644
--- a/clang/lib/CIR/Lowering/DirectToLLVM/LowerToLLVM.cpp
+++ b/clang/lib/CIR/Lowering/DirectToLLVM/LowerToLLVM.cpp
@@ -550,44 +550,32 @@ static llvm::StringRef getConstrainedExceptMetadata(cir::FenvAttr fenv) {
return strictExcept.getValue() ? "fpexcept.strict" : "fpexcept.maytrap";
}
-// CIR fast-math bits use the same positions as `mlir::LLVM::FastmathFlags`.
-static mlir::LLVM::FastmathFlags toLLVMFastMathFlags(cir::FastMathFlags flags) {
- using CirFlags = cir::FastMathFlags;
- using LLVMFlags = mlir::LLVM::FastmathFlags;
- static_assert(static_cast<uint32_t>(CirFlags::nnan) ==
- static_cast<uint32_t>(LLVMFlags::nnan),
- "CIR nnan bit must match LLVM fastmath");
- static_assert(static_cast<uint32_t>(CirFlags::ninf) ==
- static_cast<uint32_t>(LLVMFlags::ninf),
- "CIR ninf bit must match LLVM fastmath");
- static_assert(static_cast<uint32_t>(CirFlags::nsz) ==
- static_cast<uint32_t>(LLVMFlags::nsz),
- "CIR nsz bit must match LLVM fastmath");
- static_assert(static_cast<uint32_t>(CirFlags::arcp) ==
- static_cast<uint32_t>(LLVMFlags::arcp),
- "CIR arcp bit must match LLVM fastmath");
- static_assert(static_cast<uint32_t>(CirFlags::contract) ==
- static_cast<uint32_t>(LLVMFlags::contract),
- "CIR contract bit must match LLVM fastmath");
- static_assert(static_cast<uint32_t>(CirFlags::afn) ==
- static_cast<uint32_t>(LLVMFlags::afn),
- "CIR afn bit must match LLVM fastmath");
- static_assert(static_cast<uint32_t>(CirFlags::reassoc) ==
- static_cast<uint32_t>(LLVMFlags::reassoc),
- "CIR reassoc bit must match LLVM fastmath");
- return static_cast<LLVMFlags>(static_cast<uint32_t>(flags));
-}
-
-// Not inlined into lowerConstrainableFPOp. In that template, a local null
-// check on the attribute is dropped and getValue() crashes when `fastmath`
-// is absent (cir.fmuladd has no such property).
-__attribute__((noinline)) static mlir::LLVM::FastmathFlags
-readFastMathFlags(mlir::Operation *cirOp) {
+static mlir::LLVM::FastmathFlags
+convertFastMathFlags(cir::FastMathFlags cirFlags) {
+ mlir::LLVM::FastmathFlags llvmFlags{};
+ const std::pair<cir::FastMathFlags, mlir::LLVM::FastmathFlags> flags[] = {
+ {cir::FastMathFlags::nnan, mlir::LLVM::FastmathFlags::nnan},
+ {cir::FastMathFlags::ninf, mlir::LLVM::FastmathFlags::ninf},
+ {cir::FastMathFlags::nsz, mlir::LLVM::FastmathFlags::nsz},
+ {cir::FastMathFlags::arcp, mlir::LLVM::FastmathFlags::arcp},
+ {cir::FastMathFlags::contract, mlir::LLVM::FastmathFlags::contract},
+ {cir::FastMathFlags::afn, mlir::LLVM::FastmathFlags::afn},
+ {cir::FastMathFlags::reassoc, mlir::LLVM::FastmathFlags::reassoc},
+ };
+
+ for (auto [cirFlag, llvmFlag] : flags) {
+ if (bitEnumContainsAny(cirFlags, cirFlag))
+ llvmFlags = llvmFlags | llvmFlag;
+ }
+
+ return llvmFlags;
+}
+
+static mlir::LLVM::FastmathFlags readFastMathFlags(mlir::Operation *cirOp) {
auto cirFlags = cirOp->getAttrOfType<cir::FastMathFlagsAttr>("fastmath");
- if (!cirFlags.getAsOpaquePointer() ||
- cirFlags.getValue() == cir::FastMathFlags::none)
- return mlir::LLVM::FastmathFlags::none;
- return toLLVMFastMathFlags(cirFlags.getValue());
+ if (!cirFlags)
+ return {};
+ return convertFastMathFlags(cirFlags.getValue());
}
void propagateFastMathFlags(mlir::Operation *cirOp, mlir::Operation *llvmOp) {
@@ -652,27 +640,6 @@ mlir::LogicalResult lowerConstrainableFPOp(
readFastMathFlags(op));
}
-static mlir::LLVM::FastmathFlags
-convertFastMathFlags(cir::FastMathFlags cirFlags) {
- mlir::LLVM::FastmathFlags llvmFlags{};
- const std::pair<cir::FastMathFlags, mlir::LLVM::FastmathFlags> flags[] = {
- {cir::FastMathFlags::nnan, mlir::LLVM::FastmathFlags::nnan},
- {cir::FastMathFlags::ninf, mlir::LLVM::FastmathFlags::ninf},
- {cir::FastMathFlags::nsz, mlir::LLVM::FastmathFlags::nsz},
- {cir::FastMathFlags::arcp, mlir::LLVM::FastmathFlags::arcp},
- {cir::FastMathFlags::contract, mlir::LLVM::FastmathFlags::contract},
- {cir::FastMathFlags::afn, mlir::LLVM::FastmathFlags::afn},
- {cir::FastMathFlags::reassoc, mlir::LLVM::FastmathFlags::reassoc},
- };
-
- for (auto [cirFlag, llvmFlag] : flags) {
- if (bitEnumContainsAny(cirFlags, cirFlag))
- llvmFlags = llvmFlags | llvmFlag;
- }
-
- return llvmFlags;
-}
-
mlir::LogicalResult CIRToLLVMLLVMIntrinsicCallOpLowering::matchAndRewrite(
cir::LLVMIntrinsicCallOp op, OpAdaptor adaptor,
mlir::ConversionPatternRewriter &rewriter) const {
>From 3bb794c0888e49a4d7722d9d653a5d80304067df Mon Sep 17 00:00:00 2001
From: David Rivera <davidriverg at gmail.com>
Date: Fri, 25 Sep 2026 16:43:43 -0400
Subject: [PATCH 5/5] [CIR] Keep fast-math flags on the CIR builder and address
review
Drop the process-wide OpBuilder registry. CIRBaseBuilderTy now holds the
fast-math flags next to the constrained-FP state, the same way IRBuilder
does in classic CodeGen:
- CIRGenFunction::setFastMathFlags derives the flags from FPOptions (only
`contract` for now) and is called from the constructor and from
CIRGenFPOptionsRAII, which restores them like classic CodeGen's FMFGuard.
ConstrainedFPRAII restores them for nested CIRGenFunctions.
- createFAdd/FSub/FMul/FDiv/FRem go through one createFPBinOp helper that
attaches both fenv and fastmath_flags.
- Scope is limited to the FP binary ops. fneg, cmp/vec.cmp and the FP
builtins are left for follow-ups.
- The generated lowering passes op.getFastmathFlagsAttr() instead of
looking the attribute up by name.
Tests cover -ffp-contract=fast-honor-pragmas, nested pragmas with state
restored after each scope, strict FP combined with contract, and a vector
op.
---
.../CIR/Dialect/Builder/CIRBaseBuilder.h | 111 ++++----------
.../include/clang/CIR/Dialect/IR/CIRAttrs.td | 4 -
.../include/clang/CIR/Dialect/IR/CIRDialect.h | 14 --
clang/include/clang/CIR/Dialect/IR/CIROps.td | 100 ++++--------
clang/include/clang/CIR/MissingFeatures.h | 1 -
clang/lib/CIR/CodeGen/CIRGenBuiltin.cpp | 14 +-
clang/lib/CIR/CodeGen/CIRGenExprScalar.cpp | 29 +---
clang/lib/CIR/CodeGen/CIRGenFunction.cpp | 35 ++---
clang/lib/CIR/CodeGen/CIRGenFunction.h | 10 +-
clang/lib/CIR/Dialect/IR/CIRDialect.cpp | 55 -------
.../CIR/Lowering/DirectToLLVM/LowerToLLVM.cpp | 142 ++++++++----------
.../CIR/Lowering/DirectToLLVM/LowerToLLVM.h | 6 +-
clang/test/CIR/CodeGen/fp-contract-fast.c | 138 -----------------
clang/test/CIR/CodeGen/fp-contract.c | 123 ++++++++++++++-
.../test/CIR/CodeGen/fp-math-precision-opts.c | 8 +-
clang/test/CIR/CodeGenCUDA/fp-contract.cu | 60 ++------
clang/test/CIR/Lowering/fastmath-contract.cir | 47 ------
clang/test/CIR/Lowering/fenv.cir | 9 ++
clang/utils/TableGen/CIRLoweringEmitter.cpp | 38 ++---
19 files changed, 327 insertions(+), 617 deletions(-)
delete mode 100644 clang/test/CIR/CodeGen/fp-contract-fast.c
delete mode 100644 clang/test/CIR/Lowering/fastmath-contract.cir
diff --git a/clang/include/clang/CIR/Dialect/Builder/CIRBaseBuilder.h b/clang/include/clang/CIR/Dialect/Builder/CIRBaseBuilder.h
index fb9bb6eea3e56d..d28c1f68eb7bd3 100644
--- a/clang/include/clang/CIR/Dialect/Builder/CIRBaseBuilder.h
+++ b/clang/include/clang/CIR/Dialect/Builder/CIRBaseBuilder.h
@@ -66,48 +66,16 @@ class CIRBaseBuilderTy : public mlir::OpBuilder {
public:
CIRBaseBuilderTy(mlir::MLIRContext &mlirContext)
- : mlir::OpBuilder(&mlirContext) {
- registerFPDefaults();
- }
- CIRBaseBuilderTy(mlir::OpBuilder &builder) : mlir::OpBuilder(builder) {
- registerFPDefaults();
- }
- CIRBaseBuilderTy(const CIRBaseBuilderTy &other);
- CIRBaseBuilderTy &operator=(const CIRBaseBuilderTy &other);
- ~CIRBaseBuilderTy() { cir::unregisterCIRBuilderFPDefaults(this); }
+ : mlir::OpBuilder(&mlirContext) {}
+ CIRBaseBuilderTy(mlir::OpBuilder &builder) : mlir::OpBuilder(builder) {}
bool isFPConstrained = false;
clang::LangOptions::FPExceptionModeKind defaultConstrainedExcept =
clang::LangOptions::FPE_Ignore;
llvm::RoundingMode defaultConstrainedRounding =
llvm::RoundingMode::NearestTiesToEven;
- // Fast-math flags applied to floating-point ops created by this builder.
- // CIRGen currently populates `contract` only. FP op builders read these via
- // fenvForBuilder / fastMathForBuilder, so create sites omit the attributes.
cir::FastMathFlags fastMathFlags = cir::FastMathFlags::none;
- void setFastMathFlags(cir::FastMathFlags flags) { fastMathFlags = flags; }
- cir::FastMathFlags getFastMathFlags() const { return fastMathFlags; }
-
- cir::FastMathFlagsAttr getFastMathFlagsAttr() {
- if (fastMathFlags == cir::FastMathFlags::none)
- return {};
- return cir::FastMathFlagsAttr::get(getContext(), fastMathFlags);
- }
-
-private:
- static cir::FenvAttr fetchFenv(void *self) {
- return static_cast<CIRBaseBuilderTy *>(self)->getConstrainedFPAttr();
- }
- static cir::FastMathFlagsAttr fetchFastMath(void *self) {
- return static_cast<CIRBaseBuilderTy *>(self)->getFastMathFlagsAttr();
- }
- void registerFPDefaults() {
- cir::registerCIRBuilderFPDefaults(this, &CIRBaseBuilderTy::fetchFenv,
- &CIRBaseBuilderTy::fetchFastMath, this);
- }
-
-public:
mlir::Value getConstAPInt(mlir::Location loc, mlir::Type typ,
const llvm::APInt &val) {
return cir::ConstantOp::create(*this, loc, cir::IntAttr::get(typ, val));
@@ -284,6 +252,19 @@ class CIRBaseBuilderTy : public mlir::OpBuilder {
return defaultConstrainedRounding;
}
+ /// Set the fast-math flags attached to floating-point operations.
+ void setFastMathFlags(cir::FastMathFlags flags) { fastMathFlags = flags; }
+
+ /// Get the fast-math flags attached to floating-point operations.
+ cir::FastMathFlags getFastMathFlags() const { return fastMathFlags; }
+
+ /// Returns a null attribute when no fast-math flags are set.
+ cir::FastMathFlagsAttr getFastMathFlagsAttr() {
+ if (fastMathFlags == cir::FastMathFlags::none)
+ return {};
+ return cir::FastMathFlagsAttr::get(getContext(), fastMathFlags);
+ }
+
/// Build the `#cir.fenv` attribute describing the constrained floating-point
/// environment currently in effect. This is attached to floating-point
/// operations that support it to capture the rounding and exception
@@ -879,36 +860,32 @@ class CIRBaseBuilderTy : public mlir::OpBuilder {
return cir::RemOp::create(*this, loc, lhs, rhs);
}
- mlir::Value createFAdd(mlir::Location loc, mlir::Value lhs, mlir::Value rhs) {
+ template <typename OpTy>
+ mlir::Value createFPBinOp(mlir::Location loc, mlir::Value lhs,
+ mlir::Value rhs) {
assert(!cir::MissingFeatures::metaDataNode());
- // `contract` is applied by FAddOp's builder. The other fast-math bits are
- // still unimplemented.
- assert(!cir::MissingFeatures::fastMathFlags());
- return cir::FAddOp::create(*this, loc, lhs, rhs);
+ return OpTy::create(*this, loc, lhs, rhs, getConstrainedFPAttr(),
+ getFastMathFlagsAttr());
+ }
+
+ mlir::Value createFAdd(mlir::Location loc, mlir::Value lhs, mlir::Value rhs) {
+ return createFPBinOp<cir::FAddOp>(loc, lhs, rhs);
}
mlir::Value createFSub(mlir::Location loc, mlir::Value lhs, mlir::Value rhs) {
- assert(!cir::MissingFeatures::metaDataNode());
- assert(!cir::MissingFeatures::fastMathFlags());
- return cir::FSubOp::create(*this, loc, lhs, rhs);
+ return createFPBinOp<cir::FSubOp>(loc, lhs, rhs);
}
mlir::Value createFMul(mlir::Location loc, mlir::Value lhs, mlir::Value rhs) {
- assert(!cir::MissingFeatures::metaDataNode());
- assert(!cir::MissingFeatures::fastMathFlags());
- return cir::FMulOp::create(*this, loc, lhs, rhs);
+ return createFPBinOp<cir::FMulOp>(loc, lhs, rhs);
}
mlir::Value createFDiv(mlir::Location loc, mlir::Value lhs, mlir::Value rhs) {
- assert(!cir::MissingFeatures::metaDataNode());
- assert(!cir::MissingFeatures::fastMathFlags());
- return cir::FDivOp::create(*this, loc, lhs, rhs);
+ return createFPBinOp<cir::FDivOp>(loc, lhs, rhs);
}
mlir::Value createFRem(mlir::Location loc, mlir::Value lhs, mlir::Value rhs) {
- assert(!cir::MissingFeatures::metaDataNode());
- assert(!cir::MissingFeatures::fastMathFlags());
- return cir::FRemOp::create(*this, loc, lhs, rhs);
+ return createFPBinOp<cir::FRemOp>(loc, lhs, rhs);
}
mlir::Value createFNeg(mlir::Location loc, mlir::Value operand) {
@@ -938,10 +915,7 @@ class CIRBaseBuilderTy : public mlir::OpBuilder {
cir::FenvAttr fenv;
if (cir::isAnyFloatingPointType(lhs.getType()))
fenv = getConstrainedFPAttr();
- return cir::CmpOp::create(*this, loc, kind, lhs, rhs, fenv,
- cir::isAnyFloatingPointType(lhs.getType())
- ? getFastMathFlagsAttr()
- : cir::FastMathFlagsAttr{});
+ return cir::CmpOp::create(*this, loc, kind, lhs, rhs, fenv);
}
cir::VecCmpOp createVecCompare(mlir::Location loc, cir::CmpOpKind kind,
@@ -952,13 +926,10 @@ class CIRBaseBuilderTy : public mlir::OpBuilder {
VectorType integralVecTy =
cir::VectorType::get(integralTy, vecCast.getSize());
cir::FenvAttr fenv;
- cir::FastMathFlagsAttr fastmath;
- if (cir::isFPOrVectorOfFPType(lhs.getType())) {
+ if (cir::isFPOrVectorOfFPType(lhs.getType()))
fenv = getConstrainedFPAttr();
- fastmath = getFastMathFlagsAttr();
- }
return cir::VecCmpOp::create(*this, loc, integralVecTy, kind, lhs, rhs,
- fenv, fastmath);
+ fenv);
}
mlir::Value createIsNaN(mlir::Location loc, mlir::Value operand) {
@@ -1111,26 +1082,6 @@ class CIRBaseBuilderTy : public mlir::OpBuilder {
}
};
-inline CIRBaseBuilderTy::CIRBaseBuilderTy(const CIRBaseBuilderTy &other)
- : mlir::OpBuilder(other), isFPConstrained(other.isFPConstrained),
- defaultConstrainedExcept(other.defaultConstrainedExcept),
- defaultConstrainedRounding(other.defaultConstrainedRounding),
- fastMathFlags(other.fastMathFlags) {
- registerFPDefaults();
-}
-
-inline CIRBaseBuilderTy &
-CIRBaseBuilderTy::operator=(const CIRBaseBuilderTy &other) {
- if (this == &other)
- return *this;
- mlir::OpBuilder::operator=(other);
- isFPConstrained = other.isFPConstrained;
- defaultConstrainedExcept = other.defaultConstrainedExcept;
- defaultConstrainedRounding = other.defaultConstrainedRounding;
- fastMathFlags = other.fastMathFlags;
- return *this;
-}
-
} // namespace cir
#endif
diff --git a/clang/include/clang/CIR/Dialect/IR/CIRAttrs.td b/clang/include/clang/CIR/Dialect/IR/CIRAttrs.td
index f2ddef865c8b72..04dc8bb178cb87 100644
--- a/clang/include/clang/CIR/Dialect/IR/CIRAttrs.td
+++ b/clang/include/clang/CIR/Dialect/IR/CIRAttrs.td
@@ -973,10 +973,6 @@ def CIR_FastMathFlags : CIR_I32BitEnum<
Describes fast-math flags for CIR operations. This attribute is shared by
operations with floating-point semantics and is not specific to LLVM intrinsic
calls.
-
- `contract` allows a backend in Standard fusion mode to form an FMA across
- statements. It is what `-ffp-contract=fast` records. `-ffp-contract=on` is
- represented separately by `cir.fmuladd`.
}];
let separator = ", ";
diff --git a/clang/include/clang/CIR/Dialect/IR/CIRDialect.h b/clang/include/clang/CIR/Dialect/IR/CIRDialect.h
index 238dcded0a0287..d9018e93628b69 100644
--- a/clang/include/clang/CIR/Dialect/IR/CIRDialect.h
+++ b/clang/include/clang/CIR/Dialect/IR/CIRDialect.h
@@ -99,20 +99,6 @@ RecordLayoutAttr getRecordLayout(mlir::ModuleOp mod, mlir::StringAttr name);
/// Same lookup as getRecordLayout, but returns a null attribute instead of
/// asserting when the record has no layout entry.
RecordLayoutAttr tryGetRecordLayout(mlir::ModuleOp mod, mlir::StringAttr name);
-
-/// Active constrained-FP and fast-math state for a `CIRBaseBuilderTy`.
-/// Both are null when `builder` did not register any.
-///
-/// ODS builders only receive `mlir::OpBuilder &`, so floating-point ops call
-/// these instead of taking the attributes at every create site.
-FenvAttr fenvForBuilder(mlir::OpBuilder &builder);
-FastMathFlagsAttr fastMathForBuilder(mlir::OpBuilder &builder);
-
-void registerCIRBuilderFPDefaults(mlir::OpBuilder *builder,
- FenvAttr (*fenv)(void *),
- FastMathFlagsAttr (*fastMath)(void *),
- void *self);
-void unregisterCIRBuilderFPDefaults(mlir::OpBuilder *builder);
} // namespace cir
// TableGen'erated files for MLIR dialects require that a macro be defined when
diff --git a/clang/include/clang/CIR/Dialect/IR/CIROps.td b/clang/include/clang/CIR/Dialect/IR/CIROps.td
index 14ee7957549c50..be3cfe44692a6c 100644
--- a/clang/include/clang/CIR/Dialect/IR/CIROps.td
+++ b/clang/include/clang/CIR/Dialect/IR/CIROps.td
@@ -87,34 +87,12 @@ class LLVMLoweringInfo {
string llvmOp = "";
string constrainedLLVMIntrinsic = "";
bit constrainedLLVMIntrinsicHasRoundingMode = true;
- // Copy an optional `fastmath` attribute onto the lowered LLVM operation.
- // Floating-point ops that go through `lowerConstrainableFPOp` propagate the
- // attribute there and do not need this bit.
- bit propagateFastMathFlags = false;
}
class LoweringBuilders<dag p> {
dag dagParams = p;
}
-// Shorter builder for floating-point ops. Omitted `fenv` and `fastmath`
-// attributes come from the active CIR builder (null for any other builder).
-class CIR_DefaultFPAttrsBuilder<dag params, string forwarded>
- : OpBuilder<params> {
- let body = !strconcat(
- "build($_builder, $_state, ", forwarded,
- ", ::cir::fenvForBuilder($_builder), ::cir::fastMathForBuilder($_builder));");
-}
-
-// Same as CIR_DefaultFPAttrsBuilder for ops that carry fast-math flags but
-// not an fenv attribute.
-class CIR_DefaultFastMathBuilder<dag params, string forwarded>
- : OpBuilder<params> {
- let body = !strconcat(
- "build($_builder, $_state, ", forwarded,
- ", ::cir::fastMathForBuilder($_builder));");
-}
-
class CIR_Op<string mnemonic, list<Trait> traits = []> :
Op<CIR_Dialect, mnemonic, traits>, LLVMLoweringInfo {
// Should we generate an ABI lowering pattern for this op?
@@ -2182,29 +2160,18 @@ def CIR_FNegOp : CIR_UnaryOp<"fneg", CIR_AnyFloatOrVecOfFloatType> {
The `cir.fneg` operation negates the operand. The operand and result must
have the same type.
- The optional `fastmath` attribute carries LLVM fast-math flags for this
- operation. `-ffp-contract=fast` sets `contract`.
-
Example:
```
%1 = cir.fneg %0 : !cir.float
- %3 = cir.fneg %2 : !cir.double {fastmath = #cir.fastmath<contract>}
+ %3 = cir.fneg %2 : !cir.double
%5 = cir.fneg %4 : !cir.vector<4 x !cir.float>
```
}];
- let arguments = !con(commonArgs,
- (ins OptionalAttr<CIR_FastMathFlagsAttr>:$fastmath));
-
- let builders = [
- CIR_DefaultFastMathBuilder<(ins "mlir::Value":$input), "input">
- ];
-
let hasFolder = 1;
let llvmOp = "FNegOp";
- let propagateFastMathFlags = true;
}
//===----------------------------------------------------------------------===//
@@ -2659,8 +2626,7 @@ def CIR_CmpOp : CIR_Op<"cmp",
CIR_CmpOpKindAttr:$kind,
CIR_ComparableType:$lhs,
CIR_ComparableType:$rhs,
- OptionalAttr<CIR_FenvAttr>:$fenv,
- OptionalAttr<CIR_FastMathFlagsAttr>:$fastmath
+ OptionalAttr<CIR_FenvAttr>:$fenv
);
let results = (outs CIR_BoolType:$result);
@@ -2672,13 +2638,11 @@ def CIR_CmpOp : CIR_Op<"cmp",
let builders = [
OpBuilder<(ins "cir::CmpOpKind":$kind, "mlir::Value":$lhs,
"mlir::Value":$rhs), [{
- build($_builder, $_state, kind, lhs, rhs, cir::FenvAttr{},
- cir::FastMathFlagsAttr{});
+ build($_builder, $_state, kind, lhs, rhs, cir::FenvAttr{});
}]>,
OpBuilder<(ins "mlir::Type":$result, "cir::CmpOpKind":$kind,
"mlir::Value":$lhs, "mlir::Value":$rhs), [{
- build($_builder, $_state, result, kind, lhs, rhs, cir::FenvAttr{},
- cir::FastMathFlagsAttr{});
+ build($_builder, $_state, result, kind, lhs, rhs, cir::FenvAttr{});
}]>
];
@@ -2977,22 +2941,23 @@ def CIR_RemOp : CIR_BinaryOp<"rem", CIR_AnyIntOrVecOfIntType> {
// and result must all be the same floating-point scalar or vector type.
//
// The optional `fenv` attribute describes constraints on the floating-point
-// handling of the operation. The optional `fastmath` attribute carries LLVM
-// fast-math flags; `-ffp-contract=fast` sets `contract` here rather than
-// forming `cir.fmuladd`.
+// handling of the operation. The optional `fastmath_flags` attribute holds the
+// fast-math flags of the operation.
class CIR_FPBinaryOp<string mnemonic, list<Trait> traits = []>
: CIR_BinaryOp<mnemonic, CIR_AnyFloatOrVecOfFloatType,
!listconcat(CIR_FenvOpTraits, traits),
CIR_DynamicMemoryEffects> {
let arguments = !con(commonArgs, (ins
OptionalAttr<CIR_FenvAttr>:$fenv,
- OptionalAttr<CIR_FastMathFlagsAttr>:$fastmath));
+ OptionalAttr<CIR_FastMathFlagsAttr>:$fastmath_flags));
let constrainedLLVMIntrinsic = mnemonic;
let builders = [
- CIR_DefaultFPAttrsBuilder<(ins "mlir::Value":$lhs, "mlir::Value":$rhs),
- "lhs, rhs">
+ OpBuilder<(ins "mlir::Value":$lhs, "mlir::Value":$rhs), [{
+ build($_builder, $_state, lhs, rhs, cir::FenvAttr{},
+ cir::FastMathFlagsAttr{});
+ }]>
];
}
@@ -6133,8 +6098,7 @@ def CIR_VecCmpOp : CIR_Op<"vec.cmp",
CIR_CmpOpKindAttr:$kind,
CIR_VectorType:$lhs,
CIR_VectorType:$rhs,
- OptionalAttr<CIR_FenvAttr>:$fenv,
- OptionalAttr<CIR_FastMathFlagsAttr>:$fastmath
+ OptionalAttr<CIR_FenvAttr>:$fenv
);
let results = (outs CIR_VectorType:$result);
@@ -6147,8 +6111,7 @@ def CIR_VecCmpOp : CIR_Op<"vec.cmp",
let builders = [
OpBuilder<(ins "mlir::Type":$result, "cir::CmpOpKind":$kind,
"mlir::Value":$lhs, "mlir::Value":$rhs), [{
- build($_builder, $_state, result, kind, lhs, rhs, cir::FenvAttr{},
- cir::FastMathFlagsAttr{});
+ build($_builder, $_state, result, kind, lhs, rhs, cir::FenvAttr{});
}]>
];
@@ -7489,14 +7452,15 @@ class CIR_UnaryFPToFPBuiltinOp<string mnemonic, string llvmOpName>
!listconcat([SameOperandsAndResultType], CIR_FenvOpTraits)>
{
let arguments = (ins CIR_AnyFloatOrVecOfFloatType:$src,
- OptionalAttr<CIR_FenvAttr>:$fenv,
- OptionalAttr<CIR_FastMathFlagsAttr>:$fastmath);
+ OptionalAttr<CIR_FenvAttr>:$fenv);
let results = (outs CIR_AnyFloatOrVecOfFloatType:$result);
let assemblyFormat = "$src `:` type($src) attr-dict";
let builders = [
- CIR_DefaultFPAttrsBuilder<(ins "mlir::Value":$src), "src">
+ OpBuilder<(ins "mlir::Value":$src), [{
+ build($_builder, $_state, src, cir::FenvAttr{});
+ }]>
];
let llvmOp = llvmOpName;
@@ -7756,7 +7720,6 @@ def CIR_FAbsOp : CIR_UnaryFPToFPBuiltinOp<"fabs", "FAbsOp"> {
// fabs is exact and does not raise exceptions, so it is always lowered to
// the plain llvm.fabs intrinsic.
let constrainedLLVMIntrinsic = "";
- let propagateFastMathFlags = true;
}
def CIR_AbsOp : CIR_Op<"abs", [Pure, SameOperandsAndResultType]> {
@@ -7811,8 +7774,7 @@ class CIR_UnaryFPToIntBuiltinOp<string mnemonic, string llvmOpName>
: CIR_Op<mnemonic, CIR_FenvOpTraits>
{
let arguments = (ins CIR_AnyFloatType:$src,
- OptionalAttr<CIR_FenvAttr>:$fenv,
- OptionalAttr<CIR_FastMathFlagsAttr>:$fastmath);
+ OptionalAttr<CIR_FenvAttr>:$fenv);
let results = (outs CIR_IntType:$result);
let summary = [{
@@ -7825,8 +7787,9 @@ class CIR_UnaryFPToIntBuiltinOp<string mnemonic, string llvmOpName>
}];
let builders = [
- CIR_DefaultFPAttrsBuilder<(ins "mlir::Type":$result, "mlir::Value":$src),
- "result, src">
+ OpBuilder<(ins "mlir::Type":$result, "mlir::Value":$src), [{
+ build($_builder, $_state, result, src, cir::FenvAttr{});
+ }]>
];
let llvmOp = llvmOpName;
@@ -7880,8 +7843,7 @@ class CIR_BinaryFPToFPBuiltinOp<string mnemonic, string llvmOpName>
let arguments = (ins
CIR_AnyFloatOrVecOfFloatType:$lhs,
CIR_AnyFloatOrVecOfFloatType:$rhs,
- OptionalAttr<CIR_FenvAttr>:$fenv,
- OptionalAttr<CIR_FastMathFlagsAttr>:$fastmath
+ OptionalAttr<CIR_FenvAttr>:$fenv
);
let results = (outs CIR_AnyFloatOrVecOfFloatType:$result);
@@ -7891,9 +7853,10 @@ class CIR_BinaryFPToFPBuiltinOp<string mnemonic, string llvmOpName>
}];
let builders = [
- CIR_DefaultFPAttrsBuilder<(ins "mlir::Type":$result, "mlir::Value":$lhs,
- "mlir::Value":$rhs),
- "result, lhs, rhs">
+ OpBuilder<(ins "mlir::Type":$result, "mlir::Value":$lhs,
+ "mlir::Value":$rhs), [{
+ build($_builder, $_state, result, lhs, rhs, cir::FenvAttr{});
+ }]>
];
let llvmOp = llvmOpName;
@@ -7908,7 +7871,6 @@ def CIR_CopysignOp : CIR_BinaryFPToFPBuiltinOp<"copysign", "CopySignOp"> {
// copysign is exact and does not raise exceptions, so it is always lowered
// to the plain llvm.copysign intrinsic.
- let propagateFastMathFlags = true;
}
def CIR_FMaxNumOp : CIR_BinaryFPToFPBuiltinOp<"fmaxnum", "MaxNumOp"> {
@@ -8033,8 +7995,7 @@ class CIR_TernaryFPToFPBuiltinOp<string mnemonic, string llvmOpName>
CIR_AnyFloatOrVecOfFloatType:$a,
CIR_AnyFloatOrVecOfFloatType:$b,
CIR_AnyFloatOrVecOfFloatType:$c,
- OptionalAttr<CIR_FenvAttr>:$fenv,
- OptionalAttr<CIR_FastMathFlagsAttr>:$fastmath
+ OptionalAttr<CIR_FenvAttr>:$fenv
);
let results = (outs CIR_AnyFloatOrVecOfFloatType:$result);
@@ -8042,9 +8003,10 @@ class CIR_TernaryFPToFPBuiltinOp<string mnemonic, string llvmOpName>
let assemblyFormat = "$a `,` $b `,` $c `:` type($a) attr-dict";
let builders = [
- CIR_DefaultFPAttrsBuilder<(ins "mlir::Type":$result, "mlir::Value":$a,
- "mlir::Value":$b, "mlir::Value":$c),
- "result, a, b, c">
+ OpBuilder<(ins "mlir::Type":$result, "mlir::Value":$a, "mlir::Value":$b,
+ "mlir::Value":$c), [{
+ build($_builder, $_state, result, a, b, c, cir::FenvAttr{});
+ }]>
];
let llvmOp = llvmOpName;
diff --git a/clang/include/clang/CIR/MissingFeatures.h b/clang/include/clang/CIR/MissingFeatures.h
index 491ee9d9ae2ca6..4c012e993ce27f 100644
--- a/clang/include/clang/CIR/MissingFeatures.h
+++ b/clang/include/clang/CIR/MissingFeatures.h
@@ -242,7 +242,6 @@ struct MissingFeatures {
static bool isPPC_FP128Ty() { return false; }
// Fast math.
- static bool fastMathGuard() { return false; }
// Should be implemented with a moduleOp level attribute and directly
// mapped to LLVM - those can be set directly for every relevant LLVM IR
// dialect operation (log10, ...).
diff --git a/clang/lib/CIR/CodeGen/CIRGenBuiltin.cpp b/clang/lib/CIR/CodeGen/CIRGenBuiltin.cpp
index fd36a2f8ed2693..245708691b7d99 100644
--- a/clang/lib/CIR/CodeGen/CIRGenBuiltin.cpp
+++ b/clang/lib/CIR/CodeGen/CIRGenBuiltin.cpp
@@ -557,14 +557,14 @@ static RValue emitUnaryMaybeConstrainedFPBuiltin(CIRGenFunction &cgf,
CIRGenFunction::CIRGenFPOptionsRAII FPOptsRAII(cgf, &e);
- auto call = Operation::create(cgf.getBuilder(), arg.getLoc(), arg);
+ auto call = Operation::create(cgf.getBuilder(), arg.getLoc(), arg.getType(),
+ arg, cgf.getBuilder().getConstrainedFPAttr());
return RValue::get(call->getResult(0));
}
template <class Operation>
static RValue emitUnaryFPBuiltin(CIRGenFunction &cgf, const CallExpr &e) {
mlir::Value arg = cgf.emitScalarExpr(e.getArg(0));
- CIRGenFunction::CIRGenFPOptionsRAII FPOptsRAII(cgf, &e);
auto call = Operation::create(cgf.getBuilder(), arg.getLoc(), arg);
return RValue::get(call->getResult(0));
}
@@ -577,7 +577,8 @@ static RValue emitUnaryMaybeConstrainedFPToIntBuiltin(CIRGenFunction &cgf,
CIRGenFunction::CIRGenFPOptionsRAII FPOptsRAII(cgf, &e);
- auto call = Op::create(cgf.getBuilder(), src.getLoc(), resultType, src);
+ auto call = Op::create(cgf.getBuilder(), src.getLoc(), resultType, src,
+ cgf.getBuilder().getConstrainedFPAttr());
return RValue::get(call->getResult(0));
}
@@ -586,7 +587,6 @@ static RValue emitBinaryFPBuiltin(CIRGenFunction &cgf, const CallExpr &e) {
mlir::Value arg0 = cgf.emitScalarExpr(e.getArg(0));
mlir::Value arg1 = cgf.emitScalarExpr(e.getArg(1));
- CIRGenFunction::CIRGenFPOptionsRAII FPOptsRAII(cgf, &e);
mlir::Location loc = cgf.getLoc(e.getExprLoc());
mlir::Type ty = cgf.convertType(e.getType());
auto call = Op::create(cgf.getBuilder(), loc, ty, arg0, arg1);
@@ -620,7 +620,8 @@ static RValue emitTernaryMaybeConstrainedFPBuiltin(CIRGenFunction &cgf,
mlir::Location loc = cgf.getLoc(e.getExprLoc());
mlir::Type ty = cgf.convertType(e.getType());
- auto call = Op::create(cgf.getBuilder(), loc, ty, arg0, arg1, arg2);
+ auto call = Op::create(cgf.getBuilder(), loc, ty, arg0, arg1, arg2,
+ cgf.getBuilder().getConstrainedFPAttr());
return RValue::get(call->getResult(0));
}
@@ -635,7 +636,8 @@ static mlir::Value emitBinaryMaybeConstrainedFPBuiltin(CIRGenFunction &cgf,
mlir::Location loc = cgf.getLoc(e.getExprLoc());
mlir::Type ty = cgf.convertType(e.getType());
- auto call = Op::create(cgf.getBuilder(), loc, ty, arg0, arg1);
+ auto call = Op::create(cgf.getBuilder(), loc, ty, arg0, arg1,
+ cgf.getBuilder().getConstrainedFPAttr());
return call->getResult(0);
}
diff --git a/clang/lib/CIR/CodeGen/CIRGenExprScalar.cpp b/clang/lib/CIR/CodeGen/CIRGenExprScalar.cpp
index 539fd36ca3fced..fe26dc61adb45b 100644
--- a/clang/lib/CIR/CodeGen/CIRGenExprScalar.cpp
+++ b/clang/lib/CIR/CodeGen/CIRGenExprScalar.cpp
@@ -899,10 +899,8 @@ class ScalarExprEmitter : public StmtVisitor<ScalarExprEmitter, mlir::Value> {
mlir::Location loc = cgf.getLoc(e->getSourceRange().getBegin());
- if (cir::isFPOrVectorOfFPType(operand.getType())) {
- CIRGenFunction::CIRGenFPOptionsRAII FPOptsRAII(cgf, e);
+ if (cir::isFPOrVectorOfFPType(operand.getType()))
return builder.createOrFold<cir::FNegOp>(loc, operand);
- }
// TODO(cir): We might have to change this to support overflow trapping.
// Classic codegen routes unary minus through emitSub to ensure
@@ -1283,9 +1281,6 @@ class ScalarExprEmitter : public StmtVisitor<ScalarExprEmitter, mlir::Value> {
BinOpInfo boInfo = emitBinOps(e);
mlir::Value lhs = boInfo.lhs;
mlir::Value rhs = boInfo.rhs;
- std::optional<CIRGenFunction::CIRGenFPOptionsRAII> fpOpts;
- if (cir::isFPOrVectorOfFPType(lhs.getType()))
- fpOpts.emplace(cgf, boInfo.fpFeatures);
if (lhsTy->isVectorType()) {
if (!e->getType()->isVectorType()) {
@@ -1295,15 +1290,9 @@ class ScalarExprEmitter : public StmtVisitor<ScalarExprEmitter, mlir::Value> {
} else {
// Other kinds of vectors. Element-wise comparison returning
// a vector.
- result = cir::VecCmpOp::create(
- builder, cgf.getLoc(boInfo.loc), cgf.convertType(boInfo.fullType),
- kind, boInfo.lhs, boInfo.rhs,
- cir::isFPOrVectorOfFPType(boInfo.lhs.getType())
- ? builder.getConstrainedFPAttr()
- : cir::FenvAttr{},
- cir::isFPOrVectorOfFPType(boInfo.lhs.getType())
- ? builder.getFastMathFlagsAttr()
- : cir::FastMathFlagsAttr{});
+ result = cir::VecCmpOp::create(builder, cgf.getLoc(boInfo.loc),
+ cgf.convertType(boInfo.fullType), kind,
+ boInfo.lhs, boInfo.rhs);
}
} else if (boInfo.isFixedPointOp()) {
result = emitFixedPointBinOp(boInfo);
@@ -2080,9 +2069,9 @@ static mlir::Value buildFMulAdd(mlir::Location addLoc, cir::FMulOp mulOp,
// Carry the mul's fenv attribute so a constrained fmul yields a constrained
// fmuladd; the builder is under the add's FP options, not the mul's.
- mlir::Value fmuladd = cir::FMulAddOp::create(
- builder, loc, addend.getType(), mulOp0, mulOp1, addend,
- mulOp.getFenvAttr(), cir::FastMathFlagsAttr{});
+ mlir::Value fmuladd =
+ cir::FMulAddOp::create(builder, loc, addend.getType(), mulOp0, mulOp1,
+ addend, mulOp.getFenvAttr());
mulOp.erase();
return fmuladd;
}
@@ -2101,9 +2090,7 @@ static mlir::Value tryEmitFMulAdd(mlir::Location loc, const BinOpInfo &op,
"Only fadd/fsub can be the root of an fmuladd.");
// Check whether this op is fusable, i.e. -ffp-contract=on. -ffp-contract=fast
- // is not a cir.fmuladd: the builder stamps `contract` on the fmul and fadd,
- // which is what the backend fuses when it is no longer in Fast mode.
- assert(!cir::MissingFeatures::fastMathFlags());
+ // is represented by the `contract` flag on the fmul/fadd instead.
if (!op.fpFeatures.allowFPContractWithinStatement())
return nullptr;
diff --git a/clang/lib/CIR/CodeGen/CIRGenFunction.cpp b/clang/lib/CIR/CodeGen/CIRGenFunction.cpp
index c70b00eae1ec67..e1c2f33fa13bb8 100644
--- a/clang/lib/CIR/CodeGen/CIRGenFunction.cpp
+++ b/clang/lib/CIR/CodeGen/CIRGenFunction.cpp
@@ -44,17 +44,6 @@ static bool functionMightHaveBypass(const Stmt *s) {
return false;
}
-static cir::FastMathFlags
-fastMathFlagsFromFPOptions(clang::FPOptions fpFeatures) {
- // Other fast-math bits (nnan, ninf, reassoc, ...) are not modeled yet.
- // `contract` is the bit `-ffp-contract=fast` needs so a later backend in
- // Standard fusion mode can still form an FMA.
- cir::FastMathFlags flags = cir::FastMathFlags::none;
- if (fpFeatures.allowFPContractAcrossStatement())
- flags = flags | cir::FastMathFlags::contract;
- return flags;
-}
-
CIRGenFunction::CIRGenFunction(CIRGenModule &cgm, CIRGenBuilderTy &builder,
bool suppressNewContext)
: CIRGenTypeCache(cgm), cgm{cgm}, builder(builder),
@@ -62,11 +51,20 @@ CIRGenFunction::CIRGenFunction(CIRGenModule &cgm, CIRGenBuilderTy &builder,
ehStack.setCGF(this);
shouldEmitLifetimeMarkers = CodeGenUtils::shouldEmitLifetimeMarkers(
cgm.getCodeGenOpts(), getContext().getLangOpts());
- builder.setFastMathFlags(fastMathFlagsFromFPOptions(curFPFeatures));
+ setFastMathFlags(curFPFeatures);
}
CIRGenFunction::~CIRGenFunction() {}
+void CIRGenFunction::setFastMathFlags(FPOptions fpFeatures) {
+ // TODO(cir): set the remaining fast-math flags.
+ assert(!cir::MissingFeatures::fastMathFlags());
+ cir::FastMathFlags flags = cir::FastMathFlags::none;
+ if (fpFeatures.allowFPContractAcrossStatement())
+ flags = flags | cir::FastMathFlags::contract;
+ builder.setFastMathFlags(flags);
+}
+
// This is copied from clang/lib/CodeGen/CodeGenFunction.cpp
cir::TypeEvaluationKind CIRGenFunction::getEvaluationKind(QualType type) {
type = type.getCanonicalType();
@@ -1448,13 +1446,11 @@ void CIRGenFunction::CIRGenFPOptionsRAII::ConstructorHelper(
oldExcept = cgf.builder.getDefaultConstrainedExcept();
oldRounding = cgf.builder.getDefaultConstrainedRounding();
- oldFastMathFlags = cgf.builder.getFastMathFlags();
if (oldFPFeatures == fpFeatures)
return;
- // TODO(cir): create guard to restore fast math configurations.
- assert(!cir::MissingFeatures::fastMathGuard());
+ oldFastMathFlags = cgf.builder.getFastMathFlags();
llvm::RoundingMode newRoundingMode = fpFeatures.getRoundingMode();
LangOptions::FPExceptionModeKind newExceptionBehavior =
@@ -1462,11 +1458,8 @@ void CIRGenFunction::CIRGenFPOptionsRAII::ConstructorHelper(
cgf.builder.setDefaultConstrainedRounding(newRoundingMode);
cgf.builder.setDefaultConstrainedExcept(newExceptionBehavior);
- cgf.builder.setFastMathFlags(fastMathFlagsFromFPOptions(fpFeatures));
- restoredFastMathFlags = true;
- // nnan/ninf/reassoc/arcp/afn are still missing. `contract` is applied above.
- assert(!cir::MissingFeatures::fastMathFlags());
+ cgf.setFastMathFlags(fpFeatures);
assert((cgf.curFuncDecl == nullptr || cgf.builder.getIsFPConstrained() ||
isa<CXXConstructorDecl>(cgf.curFuncDecl) ||
@@ -1483,8 +1476,8 @@ CIRGenFunction::CIRGenFPOptionsRAII::~CIRGenFPOptionsRAII() {
cgf.curFPFeatures = oldFPFeatures;
cgf.builder.setDefaultConstrainedExcept(oldExcept);
cgf.builder.setDefaultConstrainedRounding(oldRounding);
- if (restoredFastMathFlags)
- cgf.builder.setFastMathFlags(oldFastMathFlags);
+ if (oldFastMathFlags)
+ cgf.builder.setFastMathFlags(*oldFastMathFlags);
}
// TODO(cir): should be shared with LLVM codegen.
diff --git a/clang/lib/CIR/CodeGen/CIRGenFunction.h b/clang/lib/CIR/CodeGen/CIRGenFunction.h
index b1230c79c8b611..9ceab0b0d6231b 100644
--- a/clang/lib/CIR/CodeGen/CIRGenFunction.h
+++ b/clang/lib/CIR/CodeGen/CIRGenFunction.h
@@ -84,15 +84,18 @@ class CIRGenFunction : public CIRGenTypeCache {
bool savedIsFPConstrained;
LangOptions::FPExceptionModeKind savedExcept;
llvm::RoundingMode savedRounding;
+ cir::FastMathFlags savedFastMathFlags;
explicit ConstrainedFPRAII(CIRGenBuilderTy &builder)
: builder(builder), savedIsFPConstrained(builder.getIsFPConstrained()),
savedExcept(builder.getDefaultConstrainedExcept()),
- savedRounding(builder.getDefaultConstrainedRounding()) {}
+ savedRounding(builder.getDefaultConstrainedRounding()),
+ savedFastMathFlags(builder.getFastMathFlags()) {}
~ConstrainedFPRAII() {
builder.setIsFPConstrained(savedIsFPConstrained);
builder.setDefaultConstrainedExcept(savedExcept);
builder.setDefaultConstrainedRounding(savedRounding);
+ builder.setFastMathFlags(savedFastMathFlags);
}
} constrainedFPState{builder};
@@ -337,11 +340,12 @@ class CIRGenFunction : public CIRGenTypeCache {
clang::FPOptions oldFPFeatures;
LangOptions::FPExceptionModeKind oldExcept;
llvm::RoundingMode oldRounding;
- cir::FastMathFlags oldFastMathFlags = cir::FastMathFlags::none;
- bool restoredFastMathFlags = false;
+ std::optional<cir::FastMathFlags> oldFastMathFlags;
};
clang::FPOptions curFPFeatures;
+ void setFastMathFlags(FPOptions fpFeatures);
+
/// The symbol table maps a variable name to a value in the current scope.
/// Entering a function creates a new scope, and the function arguments are
/// added to the mapping. When the processing of a function is terminated,
diff --git a/clang/lib/CIR/Dialect/IR/CIRDialect.cpp b/clang/lib/CIR/Dialect/IR/CIRDialect.cpp
index 9739c8de5d4708..ce406707f5942f 100644
--- a/clang/lib/CIR/Dialect/IR/CIRDialect.cpp
+++ b/clang/lib/CIR/Dialect/IR/CIRDialect.cpp
@@ -29,63 +29,14 @@
#include "clang/CIR/Dialect/IR/CIROpsDialect.cpp.inc"
#include "clang/CIR/Dialect/IR/CIROpsEnums.cpp.inc"
#include "clang/CIR/MissingFeatures.h"
-#include "llvm/ADT/DenseMap.h"
#include "llvm/ADT/SetOperations.h"
#include "llvm/ADT/SmallSet.h"
#include "llvm/ADT/TypeSwitch.h"
#include "llvm/Support/LogicalResult.h"
-#include "llvm/Support/Mutex.h"
using namespace mlir;
using namespace cir;
-namespace {
-struct CIRBuilderFPDefaults {
- FenvAttr (*fenv)(void *);
- FastMathFlagsAttr (*fastMath)(void *);
- void *self;
-};
-
-llvm::sys::SmartMutex<true> &cirBuilderFPMutex() {
- static llvm::sys::SmartMutex<true> mutex;
- return mutex;
-}
-
-llvm::DenseMap<mlir::OpBuilder *, CIRBuilderFPDefaults> &cirBuilderFPMap() {
- static llvm::DenseMap<mlir::OpBuilder *, CIRBuilderFPDefaults> map;
- return map;
-}
-} // namespace
-
-void cir::registerCIRBuilderFPDefaults(mlir::OpBuilder *builder,
- FenvAttr (*fenv)(void *),
- FastMathFlagsAttr (*fastMath)(void *),
- void *self) {
- llvm::sys::SmartScopedLock<true> lock(cirBuilderFPMutex());
- cirBuilderFPMap()[builder] = {fenv, fastMath, self};
-}
-
-void cir::unregisterCIRBuilderFPDefaults(mlir::OpBuilder *builder) {
- llvm::sys::SmartScopedLock<true> lock(cirBuilderFPMutex());
- cirBuilderFPMap().erase(builder);
-}
-
-FenvAttr cir::fenvForBuilder(mlir::OpBuilder &builder) {
- llvm::sys::SmartScopedLock<true> lock(cirBuilderFPMutex());
- auto it = cirBuilderFPMap().find(&builder);
- if (it == cirBuilderFPMap().end())
- return {};
- return it->second.fenv(it->second.self);
-}
-
-FastMathFlagsAttr cir::fastMathForBuilder(mlir::OpBuilder &builder) {
- llvm::sys::SmartScopedLock<true> lock(cirBuilderFPMutex());
- auto it = cirBuilderFPMap().find(&builder);
- if (it == cirBuilderFPMap().end())
- return {};
- return it->second.fastMath(it->second.self);
-}
-
//===----------------------------------------------------------------------===//
// CIR Dialect
//===----------------------------------------------------------------------===//
@@ -3906,9 +3857,6 @@ LogicalResult cir::CmpOp::verify() {
if (getFenvAttr() && !cir::isAnyFloatingPointType(getLhs().getType()))
return emitOpError()
<< "'fenv' is only valid for floating-point comparisons";
- if (getFastmathAttr() && !cir::isAnyFloatingPointType(getLhs().getType()))
- return emitOpError()
- << "'fastmath' is only valid for floating-point comparisons";
return success();
}
@@ -3920,9 +3868,6 @@ LogicalResult cir::VecCmpOp::verify() {
if (getFenvAttr() && !cir::isFPOrVectorOfFPType(getLhs().getType()))
return emitOpError()
<< "'fenv' is only valid for floating-point comparisons";
- if (getFastmathAttr() && !cir::isFPOrVectorOfFPType(getLhs().getType()))
- return emitOpError()
- << "'fastmath' is only valid for floating-point comparisons";
return success();
}
diff --git a/clang/lib/CIR/Lowering/DirectToLLVM/LowerToLLVM.cpp b/clang/lib/CIR/Lowering/DirectToLLVM/LowerToLLVM.cpp
index a5383045b14465..5b2172b42cbe76 100644
--- a/clang/lib/CIR/Lowering/DirectToLLVM/LowerToLLVM.cpp
+++ b/clang/lib/CIR/Lowering/DirectToLLVM/LowerToLLVM.cpp
@@ -550,45 +550,6 @@ static llvm::StringRef getConstrainedExceptMetadata(cir::FenvAttr fenv) {
return strictExcept.getValue() ? "fpexcept.strict" : "fpexcept.maytrap";
}
-static mlir::LLVM::FastmathFlags
-convertFastMathFlags(cir::FastMathFlags cirFlags) {
- mlir::LLVM::FastmathFlags llvmFlags{};
- const std::pair<cir::FastMathFlags, mlir::LLVM::FastmathFlags> flags[] = {
- {cir::FastMathFlags::nnan, mlir::LLVM::FastmathFlags::nnan},
- {cir::FastMathFlags::ninf, mlir::LLVM::FastmathFlags::ninf},
- {cir::FastMathFlags::nsz, mlir::LLVM::FastmathFlags::nsz},
- {cir::FastMathFlags::arcp, mlir::LLVM::FastmathFlags::arcp},
- {cir::FastMathFlags::contract, mlir::LLVM::FastmathFlags::contract},
- {cir::FastMathFlags::afn, mlir::LLVM::FastmathFlags::afn},
- {cir::FastMathFlags::reassoc, mlir::LLVM::FastmathFlags::reassoc},
- };
-
- for (auto [cirFlag, llvmFlag] : flags) {
- if (bitEnumContainsAny(cirFlags, cirFlag))
- llvmFlags = llvmFlags | llvmFlag;
- }
-
- return llvmFlags;
-}
-
-static mlir::LLVM::FastmathFlags readFastMathFlags(mlir::Operation *cirOp) {
- auto cirFlags = cirOp->getAttrOfType<cir::FastMathFlagsAttr>("fastmath");
- if (!cirFlags)
- return {};
- return convertFastMathFlags(cirFlags.getValue());
-}
-
-void propagateFastMathFlags(mlir::Operation *cirOp, mlir::Operation *llvmOp) {
- mlir::LLVM::FastmathFlags flags = readFastMathFlags(cirOp);
- if (flags == mlir::LLVM::FastmathFlags::none)
- return;
- auto fmfOp = dyn_cast<mlir::LLVM::FastmathFlagsInterface>(llvmOp);
- if (!fmfOp)
- return;
- fmfOp.setFastmathAttr(
- mlir::LLVM::FastmathFlagsAttr::get(llvmOp->getContext(), flags));
-}
-
static mlir::Value
createFenvMetadataValue(mlir::ConversionPatternRewriter &rewriter,
mlir::Location loc, llvm::StringRef str) {
@@ -616,28 +577,58 @@ mlir::LogicalResult lowerToConstrainedFPIntrinsic(
return mlir::success();
}
+static mlir::LLVM::FastmathFlags
+convertFastMathFlags(cir::FastMathFlags cirFlags);
+
template <typename LLVMOp>
mlir::LogicalResult lowerConstrainableFPOp(
mlir::Operation *op, mlir::ValueRange operands, cir::FenvAttr fenv,
- const mlir::TypeConverter &typeConverter,
+ cir::FastMathFlagsAttr fastmath, const mlir::TypeConverter &typeConverter,
mlir::ConversionPatternRewriter &rewriter,
llvm::StringRef constrainedMnemonic, bool hasRoundingMode) {
mlir::Type llvmResTy = typeConverter.convertType(op->getResultTypes()[0]);
if (!llvmResTy)
return op->emitError("expected LLVM result type for floating-point op");
+ mlir::LLVM::FastmathFlags fastmathFlags = {};
+ if (fastmath)
+ fastmathFlags = convertFastMathFlags(fastmath.getValue());
+
if (!fenv) {
- LLVMOp llvmOp = LLVMOp::create(
- rewriter, op->getLoc(), mlir::TypeRange{llvmResTy}, operands,
+ auto newOp = rewriter.replaceOpWithNewOp<LLVMOp>(
+ op, mlir::TypeRange{llvmResTy}, operands,
cir::getDefaultProperties<LLVMOp>(op->getContext()));
- propagateFastMathFlags(op, llvmOp);
- rewriter.replaceOp(op, llvmOp.getResult());
+ if (fastmathFlags != mlir::LLVM::FastmathFlags::none)
+ mlir::cast<mlir::LLVM::FastmathFlagsInterface>(newOp.getOperation())
+ .setFastmathAttr(mlir::LLVM::FastmathFlagsAttr::get(
+ rewriter.getContext(), fastmathFlags));
return mlir::success();
}
return lowerToConstrainedFPIntrinsic(op, operands, fenv, llvmResTy, rewriter,
constrainedMnemonic, hasRoundingMode,
- readFastMathFlags(op));
+ fastmathFlags);
+}
+
+static mlir::LLVM::FastmathFlags
+convertFastMathFlags(cir::FastMathFlags cirFlags) {
+ mlir::LLVM::FastmathFlags llvmFlags{};
+ const std::pair<cir::FastMathFlags, mlir::LLVM::FastmathFlags> flags[] = {
+ {cir::FastMathFlags::nnan, mlir::LLVM::FastmathFlags::nnan},
+ {cir::FastMathFlags::ninf, mlir::LLVM::FastmathFlags::ninf},
+ {cir::FastMathFlags::nsz, mlir::LLVM::FastmathFlags::nsz},
+ {cir::FastMathFlags::arcp, mlir::LLVM::FastmathFlags::arcp},
+ {cir::FastMathFlags::contract, mlir::LLVM::FastmathFlags::contract},
+ {cir::FastMathFlags::afn, mlir::LLVM::FastmathFlags::afn},
+ {cir::FastMathFlags::reassoc, mlir::LLVM::FastmathFlags::reassoc},
+ };
+
+ for (auto [cirFlag, llvmFlag] : flags) {
+ if (bitEnumContainsAny(cirFlags, cirFlag))
+ llvmFlags = llvmFlags | llvmFlag;
+ }
+
+ return llvmFlags;
}
mlir::LogicalResult CIRToLLVMLLVMIntrinsicCallOpLowering::matchAndRewrite(
@@ -2112,15 +2103,13 @@ mlir::LogicalResult CIRToLLVMFMaxNumOpLowering::matchAndRewrite(
cir::FMaxNumOp op, OpAdaptor adaptor,
mlir::ConversionPatternRewriter &rewriter) const {
mlir::Type resTy = typeConverter->convertType(op.getType());
- mlir::LLVM::FastmathFlags flags = static_cast<mlir::LLVM::FastmathFlags>(
- static_cast<uint32_t>(mlir::LLVM::FastmathFlags::nsz) |
- static_cast<uint32_t>(readFastMathFlags(op)));
if (cir::FenvAttr fenv = op.getFenvAttr())
- return lowerToConstrainedFPIntrinsic(op, adaptor.getOperands(), fenv, resTy,
- rewriter, "maxnum",
- /*hasRoundingMode=*/false, flags);
- rewriter.replaceOpWithNewOp<mlir::LLVM::MaxNumOp>(op, resTy, adaptor.getLhs(),
- adaptor.getRhs(), flags);
+ return lowerToConstrainedFPIntrinsic(
+ op, adaptor.getOperands(), fenv, resTy, rewriter, "maxnum",
+ /*hasRoundingMode=*/false, mlir::LLVM::FastmathFlags::nsz);
+ rewriter.replaceOpWithNewOp<mlir::LLVM::MaxNumOp>(
+ op, resTy, adaptor.getLhs(), adaptor.getRhs(),
+ mlir::LLVM::FastmathFlags::nsz);
return mlir::success();
}
@@ -2128,15 +2117,13 @@ mlir::LogicalResult CIRToLLVMFMinNumOpLowering::matchAndRewrite(
cir::FMinNumOp op, OpAdaptor adaptor,
mlir::ConversionPatternRewriter &rewriter) const {
mlir::Type resTy = typeConverter->convertType(op.getType());
- mlir::LLVM::FastmathFlags flags = static_cast<mlir::LLVM::FastmathFlags>(
- static_cast<uint32_t>(mlir::LLVM::FastmathFlags::nsz) |
- static_cast<uint32_t>(readFastMathFlags(op)));
if (cir::FenvAttr fenv = op.getFenvAttr())
- return lowerToConstrainedFPIntrinsic(op, adaptor.getOperands(), fenv, resTy,
- rewriter, "minnum",
- /*hasRoundingMode=*/false, flags);
- rewriter.replaceOpWithNewOp<mlir::LLVM::MinNumOp>(op, resTy, adaptor.getLhs(),
- adaptor.getRhs(), flags);
+ return lowerToConstrainedFPIntrinsic(
+ op, adaptor.getOperands(), fenv, resTy, rewriter, "minnum",
+ /*hasRoundingMode=*/false, mlir::LLVM::FastmathFlags::nsz);
+ rewriter.replaceOpWithNewOp<mlir::LLVM::MinNumOp>(
+ op, resTy, adaptor.getLhs(), adaptor.getRhs(),
+ mlir::LLVM::FastmathFlags::nsz);
return mlir::success();
}
@@ -3634,10 +3621,11 @@ static bool isSignalingConstrainedFCmp(cir::CmpOpKind kind) {
llvm_unreachable("Unknown CmpOpKind");
}
-static mlir::LLVM::CallIntrinsicOp createConstrainedFCmpCall(
- mlir::ConversionPatternRewriter &rewriter, mlir::Location loc,
- mlir::Value lhs, mlir::Value rhs, cir::CmpOpKind kind, cir::FenvAttr fenv,
- mlir::Type llvmResTy, mlir::LLVM::FastmathFlags fastmathFlags = {}) {
+static mlir::LLVM::CallIntrinsicOp
+createConstrainedFCmpCall(mlir::ConversionPatternRewriter &rewriter,
+ mlir::Location loc, mlir::Value lhs, mlir::Value rhs,
+ cir::CmpOpKind kind, cir::FenvAttr fenv,
+ mlir::Type llvmResTy) {
llvm::SmallVector<mlir::Value, 4> callOperands = {
lhs, rhs,
createFenvMetadataValue(rewriter, loc,
@@ -3648,7 +3636,7 @@ static mlir::LLVM::CallIntrinsicOp createConstrainedFCmpCall(
? "llvm.experimental.constrained.fcmps"
: "llvm.experimental.constrained.fcmp";
return createCallLLVMIntrinsicOp(rewriter, loc, intrinsicName, llvmResTy,
- callOperands, fastmathFlags);
+ callOperands);
}
mlir::LogicalResult CIRToLLVMCmpOpLowering::matchAndRewrite(
@@ -3683,16 +3671,14 @@ mlir::LogicalResult CIRToLLVMCmpOpLowering::matchAndRewrite(
if (cir::FenvAttr fenv = cmpOp.getFenvAttr()) {
mlir::LLVM::CallIntrinsicOp call = createConstrainedFCmpCall(
rewriter, cmpOp.getLoc(), adaptor.getLhs(), adaptor.getRhs(),
- cmpOp.getKind(), fenv, llvmResTy, readFastMathFlags(cmpOp));
+ cmpOp.getKind(), fenv, llvmResTy);
rewriter.replaceOp(cmpOp, call.getResult(0));
return mlir::success();
}
mlir::LLVM::FCmpPredicate kind =
convertCmpKindToFCmpPredicate(cmpOp.getKind());
- auto fcmp = mlir::LLVM::FCmpOp::create(rewriter, cmpOp.getLoc(), kind,
- adaptor.getLhs(), adaptor.getRhs());
- propagateFastMathFlags(cmpOp, fcmp);
- rewriter.replaceOp(cmpOp, fcmp.getResult());
+ rewriter.replaceOpWithNewOp<mlir::LLVM::FCmpOp>(
+ cmpOp, kind, adaptor.getLhs(), adaptor.getRhs());
return mlir::success();
}
@@ -3728,8 +3714,6 @@ mlir::LogicalResult CIRToLLVMCmpOpLowering::matchAndRewrite(
rewriter, loc, mlir::LLVM::FCmpPredicate::oeq, lhsReal, rhsReal);
auto imagCmp = mlir::LLVM::FCmpOp::create(
rewriter, loc, mlir::LLVM::FCmpPredicate::oeq, lhsImag, rhsImag);
- propagateFastMathFlags(cmpOp, realCmp);
- propagateFastMathFlags(cmpOp, imagCmp);
rewriter.replaceOpWithNewOp<mlir::LLVM::AndOp>(cmpOp, realCmp, imagCmp);
return mlir::success();
}
@@ -3748,8 +3732,6 @@ mlir::LogicalResult CIRToLLVMCmpOpLowering::matchAndRewrite(
rewriter, loc, mlir::LLVM::FCmpPredicate::une, lhsReal, rhsReal);
auto imagCmp = mlir::LLVM::FCmpOp::create(
rewriter, loc, mlir::LLVM::FCmpPredicate::une, lhsImag, rhsImag);
- propagateFastMathFlags(cmpOp, realCmp);
- propagateFastMathFlags(cmpOp, imagCmp);
rewriter.replaceOpWithNewOp<mlir::LLVM::OrOp>(cmpOp, realCmp, imagCmp);
return mlir::success();
}
@@ -5084,16 +5066,14 @@ mlir::LogicalResult CIRToLLVMVecCmpOpLowering::matchAndRewrite(
if (cir::FenvAttr fenv = op.getFenvAttr()) {
auto i1VecTy = mlir::VectorType::get(op.getLhs().getType().getSize(),
rewriter.getI1Type());
- bitResult = createConstrainedFCmpCall(
- rewriter, op.getLoc(), adaptor.getLhs(), adaptor.getRhs(),
- op.getKind(), fenv, i1VecTy, readFastMathFlags(op))
+ bitResult = createConstrainedFCmpCall(rewriter, op.getLoc(),
+ adaptor.getLhs(), adaptor.getRhs(),
+ op.getKind(), fenv, i1VecTy)
.getResult(0);
} else {
- auto fcmp = mlir::LLVM::FCmpOp::create(
+ bitResult = mlir::LLVM::FCmpOp::create(
rewriter, op.getLoc(), convertCmpKindToFCmpPredicate(op.getKind()),
adaptor.getLhs(), adaptor.getRhs());
- propagateFastMathFlags(op, fcmp);
- bitResult = fcmp.getResult();
}
} else {
return op.emitError() << "unsupported type for VecCmpOp: " << elementType;
diff --git a/clang/lib/CIR/Lowering/DirectToLLVM/LowerToLLVM.h b/clang/lib/CIR/Lowering/DirectToLLVM/LowerToLLVM.h
index 65f52cec486c63..e73fce6ac6e6f2 100644
--- a/clang/lib/CIR/Lowering/DirectToLLVM/LowerToLLVM.h
+++ b/clang/lib/CIR/Lowering/DirectToLLVM/LowerToLLVM.h
@@ -77,10 +77,6 @@ struct LLVMBlockAddressInfo {
int32_t blockTagOpIndex;
};
-/// Copy a CIR `fastmath` attribute onto an LLVM dialect operation that
-/// implements `FastmathFlagsInterface`. No-op when the CIR op has no flags.
-void propagateFastMathFlags(mlir::Operation *cirOp, mlir::Operation *llvmOp);
-
mlir::LogicalResult lowerToConstrainedFPIntrinsic(
mlir::Operation *op, mlir::ValueRange operands, cir::FenvAttr fenv,
mlir::Type llvmResTy, mlir::ConversionPatternRewriter &rewriter,
@@ -90,7 +86,7 @@ mlir::LogicalResult lowerToConstrainedFPIntrinsic(
template <typename LLVMOp>
mlir::LogicalResult lowerConstrainableFPOp(
mlir::Operation *op, mlir::ValueRange operands, cir::FenvAttr fenv,
- const mlir::TypeConverter &typeConverter,
+ cir::FastMathFlagsAttr fastmath, const mlir::TypeConverter &typeConverter,
mlir::ConversionPatternRewriter &rewriter,
llvm::StringRef constrainedMnemonic, bool hasRoundingMode);
diff --git a/clang/test/CIR/CodeGen/fp-contract-fast.c b/clang/test/CIR/CodeGen/fp-contract-fast.c
deleted file mode 100644
index 9c1b57e6a12354..00000000000000
--- a/clang/test/CIR/CodeGen/fp-contract-fast.c
+++ /dev/null
@@ -1,138 +0,0 @@
-// -ffp-contract=fast does not form cir.fmuladd. It stamps `contract` on the
-// individual floating-point ops so a backend in Standard fusion mode can
-// still contract them, including across statements.
-// -ffp-contract=on still forms cir.fmuladd and does not set `contract`.
-// -ffp-contract=off does neither.
-
-// RUN: %clang_cc1 -triple x86_64-unknown-linux-gnu -Wno-unused-value -fclangir -ffp-contract=fast -emit-cir %s -o %t-fast.cir
-// RUN: FileCheck --input-file=%t-fast.cir %s -check-prefix=CIR-FAST
-// RUN: %clang_cc1 -triple x86_64-unknown-linux-gnu -Wno-unused-value -fclangir -ffp-contract=fast -emit-llvm %s -o %t-fast.ll
-// RUN: FileCheck --input-file=%t-fast.ll %s -check-prefix=LLVM-FAST
-
-// RUN: %clang_cc1 -triple x86_64-unknown-linux-gnu -Wno-unused-value -fclangir -ffp-contract=on -emit-cir %s -o %t-on.cir
-// RUN: FileCheck --input-file=%t-on.cir %s -check-prefix=CIR-ON
-// RUN: %clang_cc1 -triple x86_64-unknown-linux-gnu -Wno-unused-value -fclangir -ffp-contract=off -emit-cir %s -o %t-off.cir
-// RUN: FileCheck --input-file=%t-off.cir %s -check-prefix=CIR-OFF
-
-// RUN: %clang_cc1 -triple x86_64-unknown-linux-gnu -Wno-unused-value -ffp-contract=fast -emit-llvm %s -o %t-og.ll
-// RUN: FileCheck --input-file=%t-og.ll %s -check-prefix=LLVM-FAST
-
-float same_stmt(float a, float b, float c) { return a * b + c; }
-// CIR-FAST-LABEL: cir.func {{.*}}@same_stmt
-// CIR-FAST: cir.fmul {{.*}} : !cir.float {fastmath = #cir.fastmath<contract>}
-// CIR-FAST: cir.fadd {{.*}} : !cir.float {fastmath = #cir.fastmath<contract>}
-// CIR-FAST-NOT: cir.fmuladd
-
-// CIR-ON-LABEL: cir.func {{.*}}@same_stmt
-// CIR-ON: cir.fmuladd
-// CIR-ON-NOT: #cir.fastmath
-
-// CIR-OFF-LABEL: cir.func {{.*}}@same_stmt
-// CIR-OFF: cir.fmul {{.*}} : !cir.float
-// CIR-OFF: cir.fadd {{.*}} : !cir.float
-// CIR-OFF-NOT: #cir.fastmath
-// CIR-OFF-NOT: cir.fmuladd
-
-// LLVM-FAST-LABEL: @same_stmt
-// LLVM-FAST: fmul contract float
-// LLVM-FAST: fadd contract float
-// LLVM-FAST-NOT: @llvm.fmuladd
-
-float across_stmt(float a, float b, float c) {
- float t = a * b;
- return t + c;
-}
-// CIR-FAST-LABEL: cir.func {{.*}}@across_stmt
-// CIR-FAST: cir.fmul {{.*}} : !cir.float {fastmath = #cir.fastmath<contract>}
-// CIR-FAST: cir.fadd {{.*}} : !cir.float {fastmath = #cir.fastmath<contract>}
-// CIR-FAST-NOT: cir.fmuladd
-
-// CIR-ON-LABEL: cir.func {{.*}}@across_stmt
-// CIR-ON: cir.fmul {{.*}} : !cir.float
-// CIR-ON: cir.fadd {{.*}} : !cir.float
-// CIR-ON-NOT: cir.fmuladd
-// CIR-ON-NOT: #cir.fastmath
-
-// LLVM-FAST-LABEL: @across_stmt
-// LLVM-FAST: fmul contract float
-// LLVM-FAST: fadd contract float
-
-float sub_stmt(float a, float b, float c) { return a * b - c; }
-// CIR-FAST-LABEL: cir.func {{.*}}@sub_stmt
-// CIR-FAST: cir.fmul {{.*}} : !cir.float {fastmath = #cir.fastmath<contract>}
-// CIR-FAST: cir.fsub {{.*}} : !cir.float {fastmath = #cir.fastmath<contract>}
-// CIR-FAST-NOT: cir.fmuladd
-
-// CIR-ON-LABEL: cir.func {{.*}}@sub_stmt
-// CIR-ON: cir.fmuladd
-// CIR-ON-NOT: #cir.fastmath
-
-// CIR-OFF-LABEL: cir.func {{.*}}@sub_stmt
-// CIR-OFF: cir.fmul {{.*}} : !cir.float
-// CIR-OFF: cir.fsub {{.*}} : !cir.float
-// CIR-OFF-NOT: #cir.fastmath
-// CIR-OFF-NOT: cir.fmuladd
-
-// LLVM-FAST-LABEL: @sub_stmt
-// LLVM-FAST: fmul contract float
-// LLVM-FAST: fsub contract float
-
-float neg(float a) { return -a; }
-// CIR-FAST-LABEL: cir.func {{.*}}@neg
-// CIR-FAST: cir.fneg {{.*}} : !cir.float {fastmath = #cir.fastmath<contract>}
-// CIR-ON-LABEL: cir.func {{.*}}@neg
-// CIR-ON: cir.fneg {{.*}} : !cir.float
-// CIR-ON-NOT: #cir.fastmath
-// CIR-OFF-LABEL: cir.func {{.*}}@neg
-// CIR-OFF: cir.fneg {{.*}} : !cir.float
-// CIR-OFF-NOT: #cir.fastmath
-// LLVM-FAST-LABEL: @neg
-// LLVM-FAST: fneg contract float
-
-int cmp(float a, float b) { return a < b; }
-// CIR-FAST-LABEL: cir.func {{.*}}@cmp
-// CIR-FAST: cir.cmp lt {{.*}} : !cir.float {fastmath = #cir.fastmath<contract>}
-// CIR-ON-LABEL: cir.func {{.*}}@cmp
-// CIR-ON: cir.cmp lt {{.*}} : !cir.float
-// CIR-ON-NOT: #cir.fastmath
-// CIR-OFF-LABEL: cir.func {{.*}}@cmp
-// CIR-OFF: cir.cmp lt {{.*}} : !cir.float
-// CIR-OFF-NOT: #cir.fastmath
-// LLVM-FAST-LABEL: @cmp
-// LLVM-FAST: fcmp contract olt float
-
-float rem(float a, float b) { return __builtin_fmodf(a, b); }
-// CIR-FAST-LABEL: cir.func {{.*}}@rem
-// CIR-FAST: cir.fmod {{.*}} : !cir.float {fastmath = #cir.fastmath<contract>}
-// CIR-ON-LABEL: cir.func {{.*}}@rem
-// CIR-ON: cir.fmod {{.*}} : !cir.float
-// CIR-ON-NOT: #cir.fastmath
-// CIR-OFF-LABEL: cir.func {{.*}}@rem
-// CIR-OFF: cir.fmod {{.*}} : !cir.float
-// CIR-OFF-NOT: #cir.fastmath
-// LLVM-FAST-LABEL: @rem
-// LLVM-FAST: frem contract float
-
-float sq(float a) { return __builtin_sqrtf(a); }
-// CIR-FAST-LABEL: cir.func {{.*}}@sq
-// CIR-FAST: cir.sqrt {{.*}} : !cir.float {fastmath = #cir.fastmath<contract>}
-// CIR-ON-LABEL: cir.func {{.*}}@sq
-// CIR-ON: cir.sqrt {{.*}} : !cir.float
-// CIR-ON-NOT: #cir.fastmath
-// CIR-OFF-LABEL: cir.func {{.*}}@sq
-// CIR-OFF: cir.sqrt {{.*}} : !cir.float
-// CIR-OFF-NOT: #cir.fastmath
-// LLVM-FAST-LABEL: @sq
-// LLVM-FAST: call contract float @llvm.sqrt.f32
-
-float absf(float a) { return __builtin_fabsf(a); }
-// CIR-FAST-LABEL: cir.func {{.*}}@absf
-// CIR-FAST: cir.fabs {{.*}} : !cir.float {fastmath = #cir.fastmath<contract>}
-// CIR-ON-LABEL: cir.func {{.*}}@absf
-// CIR-ON: cir.fabs {{.*}} : !cir.float
-// CIR-ON-NOT: #cir.fastmath
-// CIR-OFF-LABEL: cir.func {{.*}}@absf
-// CIR-OFF: cir.fabs {{.*}} : !cir.float
-// CIR-OFF-NOT: #cir.fastmath
-// LLVM-FAST-LABEL: @absf
-// LLVM-FAST: call contract float @llvm.fabs.f32
diff --git a/clang/test/CIR/CodeGen/fp-contract.c b/clang/test/CIR/CodeGen/fp-contract.c
index 3abe498a0cfb73..d20587c087e5a5 100644
--- a/clang/test/CIR/CodeGen/fp-contract.c
+++ b/clang/test/CIR/CodeGen/fp-contract.c
@@ -1,6 +1,7 @@
// Test that -ffp-contract=on fuses a*b+c / a*b-c into cir.fmuladd and that
-// -ffp-contract=off does not. The CIR-lowered and classic CodeGen LLVM IR
-// match here, so both feed the LLVM-* prefixes.
+// -ffp-contract=off does not. -ffp-contract=fast sets `contract` on the fmul and
+// fadd instead. The CIR-lowered and classic CodeGen LLVM IR match here, so both
+// feed the LLVM-* prefixes.
// RUN: %clang_cc1 -triple x86_64-unknown-linux-gnu -Wno-unused-value -fclangir -ffp-contract=on -emit-cir %s -o %t.cir
// RUN: FileCheck --input-file=%t.cir %s -check-prefix=CIR-ON
@@ -17,6 +18,21 @@
// RUN: %clang_cc1 -triple x86_64-unknown-linux-gnu -Wno-unused-value -ffp-contract=off -emit-llvm %s -o %t-off.ll
// RUN: FileCheck --input-file=%t-off.ll %s -check-prefix=LLVM-OFF
+// RUN: %clang_cc1 -triple x86_64-unknown-linux-gnu -Wno-unused-value -fclangir -ffp-contract=fast -emit-cir %s -o %t-fast.cir
+// RUN: FileCheck --input-file=%t-fast.cir %s -check-prefix=CIR-FAST
+// RUN: %clang_cc1 -triple x86_64-unknown-linux-gnu -Wno-unused-value -fclangir -ffp-contract=fast -emit-llvm %s -o %t-fast.ll
+// RUN: FileCheck --input-file=%t-fast.ll %s -check-prefix=LLVM-FAST
+// RUN: %clang_cc1 -triple x86_64-unknown-linux-gnu -Wno-unused-value -ffp-contract=fast -emit-llvm %s -o %t-fast-ogcg.ll
+// RUN: FileCheck --input-file=%t-fast-ogcg.ll %s -check-prefix=LLVM-FAST
+
+// -ffp-contract=fast-honor-pragmas matches -ffp-contract=fast here.
+// RUN: %clang_cc1 -triple x86_64-unknown-linux-gnu -Wno-unused-value -fclangir -ffp-contract=fast-honor-pragmas -emit-cir %s -o %t-fhp.cir
+// RUN: FileCheck --input-file=%t-fhp.cir %s -check-prefix=CIR-FAST
+// RUN: %clang_cc1 -triple x86_64-unknown-linux-gnu -Wno-unused-value -fclangir -ffp-contract=fast-honor-pragmas -emit-llvm %s -o %t-fhp.ll
+// RUN: FileCheck --input-file=%t-fhp.ll %s -check-prefix=LLVM-FAST
+// RUN: %clang_cc1 -triple x86_64-unknown-linux-gnu -Wno-unused-value -ffp-contract=fast-honor-pragmas -emit-llvm %s -o %t-fhp-ogcg.ll
+// RUN: FileCheck --input-file=%t-fhp-ogcg.ll %s -check-prefix=LLVM-FAST
+
// Under strict FP the fused op carries an fenv attribute and lowers to the
// constrained fmuladd intrinsic.
// RUN: %clang_cc1 -triple x86_64-unknown-linux-gnu -Wno-unused-value -fclangir -ffp-contract=on -fexperimental-strict-floating-point -ffp-exception-behavior=strict -emit-cir %s -o %t-strict.cir
@@ -26,6 +42,15 @@
// RUN: %clang_cc1 -triple x86_64-unknown-linux-gnu -Wno-unused-value -ffp-contract=on -fexperimental-strict-floating-point -ffp-exception-behavior=strict -emit-llvm %s -o %t-strict-ogcg.ll
// RUN: FileCheck --input-file=%t-strict-ogcg.ll %s -check-prefix=LLVM-STRICT
+// Under strict FP with -ffp-contract=fast, the constrained intrinsics carry
+// `contract`.
+// RUN: %clang_cc1 -triple x86_64-unknown-linux-gnu -Wno-unused-value -fclangir -ffp-contract=fast -fexperimental-strict-floating-point -ffp-exception-behavior=strict -emit-cir %s -o %t-strict-fast.cir
+// RUN: FileCheck --input-file=%t-strict-fast.cir %s -check-prefix=CIR-STRICT-FAST
+// RUN: %clang_cc1 -triple x86_64-unknown-linux-gnu -Wno-unused-value -fclangir -ffp-contract=fast -fexperimental-strict-floating-point -ffp-exception-behavior=strict -emit-llvm %s -o %t-strict-fast.ll
+// RUN: FileCheck --input-file=%t-strict-fast.ll %s -check-prefix=LLVM-STRICT-FAST
+// RUN: %clang_cc1 -triple x86_64-unknown-linux-gnu -Wno-unused-value -ffp-contract=fast -fexperimental-strict-floating-point -ffp-exception-behavior=strict -emit-llvm %s -o %t-strict-fast-ogcg.ll
+// RUN: FileCheck --input-file=%t-strict-fast-ogcg.ll %s -check-prefix=LLVM-STRICT-FAST
+
// a * b + c => fmuladd(a, b, c)
float fmuladd_add(float a, float b, float c) {
return a * b + c;
@@ -45,6 +70,15 @@ float fmuladd_add(float a, float b, float c) {
// LLVM-OFF: fmul float
// LLVM-OFF: fadd float
+// CIR-FAST-LABEL: cir.func {{.*}}@fmuladd_add
+// CIR-FAST: cir.fmul %{{.*}}, %{{.*}} : !cir.float {fastmath_flags = #cir.fastmath<contract>}
+// CIR-FAST: cir.fadd %{{.*}}, %{{.*}} : !cir.float {fastmath_flags = #cir.fastmath<contract>}
+// CIR-FAST-NOT: cir.fmuladd
+
+// LLVM-FAST-LABEL: @fmuladd_add
+// LLVM-FAST: fmul contract float
+// LLVM-FAST: fadd contract float
+
// c + a * b => fmuladd(a, b, c) (mul on the RHS)
float fmuladd_add_rhs(float a, float b, float c) {
@@ -95,6 +129,14 @@ float4 fmuladd_vec(float4 a, float4 b, float4 c) {
// LLVM-ON-LABEL: @fmuladd_vec
// LLVM-ON: call <4 x float> @llvm.fmuladd.v4f32
+// CIR-FAST-LABEL: cir.func {{.*}}@fmuladd_vec
+// CIR-FAST: cir.fmul %{{.*}}, %{{.*}} : !cir.vector<4 x !cir.float> {fastmath_flags = #cir.fastmath<contract>}
+// CIR-FAST: cir.fadd %{{.*}}, %{{.*}} : !cir.vector<4 x !cir.float> {fastmath_flags = #cir.fastmath<contract>}
+
+// LLVM-FAST-LABEL: @fmuladd_vec
+// LLVM-FAST: fmul contract <4 x float>
+// LLVM-FAST: fadd contract <4 x float>
+
// Strict FP: fused op carries an fenv attr, lowering to the constrained
// fmuladd intrinsic.
float fmuladd_strict(float a, float b, float c) {
@@ -105,6 +147,16 @@ float fmuladd_strict(float a, float b, float c) {
// LLVM-STRICT-LABEL: @fmuladd_strict
// LLVM-STRICT: call float @llvm.experimental.constrained.fmuladd.f32
+// CIR-STRICT-FAST-LABEL: cir.func {{.*}}@fmuladd_strict
+// CIR-STRICT-FAST: cir.fmul %{{.*}}, %{{.*}} : !cir.float
+// CIR-STRICT-FAST-SAME: fastmath_flags = #cir.fastmath<contract>
+// CIR-STRICT-FAST: cir.fadd %{{.*}}, %{{.*}} : !cir.float
+// CIR-STRICT-FAST-SAME: fastmath_flags = #cir.fastmath<contract>
+
+// LLVM-STRICT-FAST-LABEL: @fmuladd_strict
+// LLVM-STRICT-FAST: call contract float @llvm.experimental.constrained.fmul.f32
+// LLVM-STRICT-FAST: call contract float @llvm.experimental.constrained.fadd.f32
+
// Strict FP with a negated addend: the fmuladd carries the mul's fenv while
// the fneg (which takes none) lowers to a plain fneg.
float fmuladd_sub_strict(float a, float b, float c) {
@@ -138,3 +190,70 @@ float fmuladd_sub_assign(float x, float a, float b) {
// LLVM-ON-LABEL: @fmuladd_sub_assign
// LLVM-ON: fneg float
// LLVM-ON: call float @llvm.fmuladd.f32
+
+// The pragma turns contraction off for this function only.
+float contract_pragma_off(float a, float b, float c) {
+#pragma clang fp contract(off)
+ return a * b + c;
+}
+// CIR-FAST-LABEL: cir.func {{.*}}@contract_pragma_off
+// CIR-FAST-NOT: #cir.fastmath
+// CIR-FAST: cir.return
+
+// LLVM-FAST-LABEL: @contract_pragma_off
+// LLVM-FAST: fmul float
+// LLVM-FAST: fadd float
+
+// -ffp-contract=fast also allows contraction across statements.
+float contract_across_stmt(float a, float b, float c) {
+ float t = a * b;
+ return t + c;
+}
+// CIR-FAST-LABEL: cir.func {{.*}}@contract_across_stmt
+// CIR-FAST: cir.fmul %{{.*}}, %{{.*}} : !cir.float {fastmath_flags = #cir.fastmath<contract>}
+// CIR-FAST: cir.fadd %{{.*}}, %{{.*}} : !cir.float {fastmath_flags = #cir.fastmath<contract>}
+
+// LLVM-FAST-LABEL: @contract_across_stmt
+// LLVM-FAST: fmul contract float
+// LLVM-FAST: fadd contract float
+
+// Nested pragmas: each scope sets its own flags and the enclosing ones are
+// restored on exit.
+float nested_pragmas(float a, float b, float c) {
+ float r;
+ {
+#pragma STDC FP_CONTRACT OFF
+ r = a * b + c;
+ {
+#pragma clang fp contract(fast)
+ r = r * a + c;
+ }
+ r = r * b + c;
+ }
+ {
+#pragma float_control(precise, on)
+ r = r * a + b;
+ }
+ return r * c + a;
+}
+// CIR-FAST-LABEL: cir.func {{.*}}@nested_pragmas
+// CIR-FAST: cir.fmul %{{.*}}, %{{.*}} : !cir.float{{$}}
+// CIR-FAST: cir.fadd %{{.*}}, %{{.*}} : !cir.float{{$}}
+// CIR-FAST: cir.fmul %{{.*}}, %{{.*}} : !cir.float {fastmath_flags = #cir.fastmath<contract>}
+// CIR-FAST: cir.fadd %{{.*}}, %{{.*}} : !cir.float {fastmath_flags = #cir.fastmath<contract>}
+// CIR-FAST: cir.fmul %{{.*}}, %{{.*}} : !cir.float{{$}}
+// CIR-FAST: cir.fadd %{{.*}}, %{{.*}} : !cir.float{{$}}
+// CIR-FAST: cir.fmuladd %{{.*}}, %{{.*}}, %{{.*}} : !cir.float{{$}}
+// CIR-FAST: cir.fmul %{{.*}}, %{{.*}} : !cir.float {fastmath_flags = #cir.fastmath<contract>}
+// CIR-FAST: cir.fadd %{{.*}}, %{{.*}} : !cir.float {fastmath_flags = #cir.fastmath<contract>}
+
+// LLVM-FAST-LABEL: @nested_pragmas
+// LLVM-FAST: fmul float
+// LLVM-FAST: fadd float
+// LLVM-FAST: fmul contract float
+// LLVM-FAST: fadd contract float
+// LLVM-FAST: fmul float
+// LLVM-FAST: fadd float
+// LLVM-FAST: call float @llvm.fmuladd.f32
+// LLVM-FAST: fmul contract float
+// LLVM-FAST: fadd contract float
diff --git a/clang/test/CIR/CodeGen/fp-math-precision-opts.c b/clang/test/CIR/CodeGen/fp-math-precision-opts.c
index e365d594f33c31..4c04eb14c6309b 100644
--- a/clang/test/CIR/CodeGen/fp-math-precision-opts.c
+++ b/clang/test/CIR/CodeGen/fp-math-precision-opts.c
@@ -57,10 +57,10 @@ float test_fast(float f) {
// Should produce an intrinsic at -O1
return __builtin_cosf(f);
// ALL: test_fast
- // CIR-ERRNO-O1: cir.cos {{.*}} {fastmath = #cir.fastmath<contract>}
- // CIR-NO-ERRNO-O1: cir.cos {{.*}} {fastmath = #cir.fastmath<contract>}
- // LLVM-ERRNO-O1: call contract float @llvm.cos.f32
- // LLVM-NO-ERRNO-O1: call contract float @llvm.cos.f32
+ // CIR-ERRNO-O1: cir.cos
+ // CIR-NO-ERRNO-O1: cir.cos
+ // LLVM-ERRNO-O1: call float @llvm.cos.f32
+ // LLVM-NO-ERRNO-O1: call float @llvm.cos.f32
// OGCG-ERRNO-O1: call {{.*}} float @llvm.cos.f32
// OGCG-NO-ERRNO-O1: call {{.*}} float @llvm.cos.f32
}
diff --git a/clang/test/CIR/CodeGenCUDA/fp-contract.cu b/clang/test/CIR/CodeGenCUDA/fp-contract.cu
index 5b993e358f086f..26080b38804ac4 100644
--- a/clang/test/CIR/CodeGenCUDA/fp-contract.cu
+++ b/clang/test/CIR/CodeGenCUDA/fp-contract.cu
@@ -1,57 +1,23 @@
-// CUDA's default contract mode is -ffp-contract=fast. CIR records that as
-// `contract` on fmul/fadd, not as cir.fmuladd. The lowered LLVM IR must carry
-// the same flag so a later Standard-fusion backend still emits FMA.
+// CUDA device compilation defaults to -ffp-contract=fast.
// RUN: %clang_cc1 -triple nvptx64-nvidia-cuda -x cuda -fcuda-is-device \
// RUN: -fclangir -emit-cir %s -o %t.cir
-// RUN: FileCheck --input-file=%t.cir %s -check-prefix=CIR-FAST
+// RUN: FileCheck --input-file=%t.cir %s -check-prefix=CIR
// RUN: %clang_cc1 -triple nvptx64-nvidia-cuda -x cuda -fcuda-is-device \
-// RUN: -fclangir -emit-llvm %s -o %t.ll
-// RUN: FileCheck --input-file=%t.ll %s -check-prefix=LLVM-FAST
-
-// RUN: %clang_cc1 -triple nvptx64-nvidia-cuda -x cuda -fcuda-is-device \
-// RUN: -ffp-contract=on -fclangir -emit-cir %s -o %t-on.cir
-// RUN: FileCheck --input-file=%t-on.cir %s -check-prefix=CIR-ON
-
+// RUN: -fclangir -emit-llvm %s -o %t-cir.ll
+// RUN: FileCheck --input-file=%t-cir.ll %s -check-prefix=LLVM
// RUN: %clang_cc1 -triple nvptx64-nvidia-cuda -x cuda -fcuda-is-device \
-// RUN: -emit-llvm %s -o %t-og.ll
-// RUN: FileCheck --input-file=%t-og.ll %s -check-prefix=LLVM-FAST
+// RUN: -emit-llvm %s -o %t.ll
+// RUN: FileCheck --input-file=%t.ll %s -check-prefix=LLVM
#include "Inputs/cuda.h"
-__host__ __device__ float same_stmt(float a, float b, float c) {
- return a * b + c;
-}
-// CIR-FAST-LABEL: cir.func {{.*}}@_Z9same_stmtfff
-// CIR-FAST: cir.fmul {{.*}} : !cir.float {fastmath = #cir.fastmath<contract>}
-// CIR-FAST: cir.fadd {{.*}} : !cir.float {fastmath = #cir.fastmath<contract>}
-// CIR-FAST-NOT: cir.fmuladd
-
-// CIR-ON-LABEL: cir.func {{.*}}@_Z9same_stmtfff
-// CIR-ON: cir.fmuladd
-// CIR-ON-NOT: #cir.fastmath
-
-// LLVM-FAST-LABEL: @_Z9same_stmtfff
-// LLVM-FAST: fmul contract float
-// LLVM-FAST: fadd contract float
-// LLVM-FAST-NOT: @llvm.fmuladd
-
-__host__ __device__ float across_stmt(float a, float b, float c) {
- float t = a * b;
- return t + c;
-}
-// CIR-FAST-LABEL: cir.func {{.*}}@_Z11across_stmtfff
-// CIR-FAST: cir.fmul {{.*}} : !cir.float {fastmath = #cir.fastmath<contract>}
-// CIR-FAST: cir.fadd {{.*}} : !cir.float {fastmath = #cir.fastmath<contract>}
-// CIR-FAST-NOT: cir.fmuladd
+__device__ float axpy(float a, float x, float y) { return a * x + y; }
-// CIR-ON-LABEL: cir.func {{.*}}@_Z11across_stmtfff
-// CIR-ON: cir.fmul {{.*}} : !cir.float
-// CIR-ON: cir.fadd {{.*}} : !cir.float
-// CIR-ON-NOT: cir.fmuladd
-// CIR-ON-NOT: #cir.fastmath
-// CIR-ON: cir.return
+// CIR-LABEL: cir.func {{.*}}@_Z4axpyfff
+// CIR: cir.fmul %{{.*}}, %{{.*}} : !cir.float {fastmath_flags = #cir.fastmath<contract>}
+// CIR: cir.fadd %{{.*}}, %{{.*}} : !cir.float {fastmath_flags = #cir.fastmath<contract>}
-// LLVM-FAST-LABEL: @_Z11across_stmtfff
-// LLVM-FAST: fmul contract float
-// LLVM-FAST: fadd contract float
+// LLVM-LABEL: @_Z4axpyfff
+// LLVM: fmul contract float
+// LLVM: fadd contract float
diff --git a/clang/test/CIR/Lowering/fastmath-contract.cir b/clang/test/CIR/Lowering/fastmath-contract.cir
deleted file mode 100644
index a75a9391936dbe..00000000000000
--- a/clang/test/CIR/Lowering/fastmath-contract.cir
+++ /dev/null
@@ -1,47 +0,0 @@
-// RUN: cir-opt %s --verify-roundtrip | FileCheck %s -check-prefix=CIR
-// RUN: cir-opt %s -cir-to-llvm -o - | FileCheck %s -check-prefix=MLIR
-// RUN: cir-translate %s -cir-to-llvmir --target nvptx64-nvidia-cuda --disable-cc-lowering | FileCheck %s -check-prefix=LLVM
-
-!f32 = !cir.float
-
-module attributes {cir.triple = "nvptx64-nvidia-cuda"} {
- cir.func @contract(%a: !f32, %b: !f32, %c: !f32) -> !f32 {
- %m = cir.fmul %a, %b : !f32 {fastmath = #cir.fastmath<contract>}
- %n = cir.fneg %c : !f32 {fastmath = #cir.fastmath<contract>}
- %s = cir.fadd %m, %n : !f32 {fastmath = #cir.fastmath<contract>}
- cir.return %s : !f32
- }
-
- cir.func @plain(%a: !f32, %b: !f32, %c: !f32) -> !f32 {
- %m = cir.fmul %a, %b : !f32
- %s = cir.fadd %m, %c : !f32
- cir.return %s : !f32
- }
-}
-
-// CIR-LABEL: cir.func @contract
-// CIR: cir.fmul {{.*}} {fastmath = #cir.fastmath<contract>}
-// CIR: cir.fneg {{.*}} {fastmath = #cir.fastmath<contract>}
-// CIR: cir.fadd {{.*}} {fastmath = #cir.fastmath<contract>}
-// CIR-LABEL: cir.func @plain
-// CIR: cir.fmul {{.*}} : !cir.float
-// CIR-NOT: #cir.fastmath
-// CIR: cir.fadd {{.*}} : !cir.float
-
-// MLIR-LABEL: llvm.func @contract
-// MLIR: llvm.fmul {{.*}} fastmath<contract> : f32
-// MLIR: llvm.fneg {{.*}} fastmath<contract> : f32
-// MLIR: llvm.fadd {{.*}} fastmath<contract> : f32
-// MLIR-LABEL: llvm.func @plain
-// MLIR: llvm.fmul {{.*}} : f32
-// MLIR-NOT: fastmath<
-// MLIR: llvm.fadd {{.*}} : f32
-
-// LLVM-LABEL: @contract
-// LLVM: fmul contract float
-// LLVM: fneg contract float
-// LLVM: fadd contract float
-// LLVM-LABEL: @plain
-// LLVM: fmul float
-// LLVM-NOT: contract
-// LLVM: fadd float
diff --git a/clang/test/CIR/Lowering/fenv.cir b/clang/test/CIR/Lowering/fenv.cir
index 415cf19ce95426..6c7b9942f6ab91 100644
--- a/clang/test/CIR/Lowering/fenv.cir
+++ b/clang/test/CIR/Lowering/fenv.cir
@@ -251,4 +251,13 @@ module {
%4 = cir.lround %a : !cir.double -> !s64i
cir.return
}
+
+ // CHECK-LABEL: llvm.func @fastmath_flags
+ cir.func @fastmath_flags(%a: !cir.float, %b: !cir.float) {
+ // CHECK: llvm.fmul %{{.*}}, %{{.*}} fastmath<contract> : f32
+ %0 = cir.fmul %a, %b : !cir.float {fastmath_flags = #cir.fastmath<contract>}
+ // CHECK: llvm.call_intrinsic "llvm.experimental.constrained.fadd"(%{{.*}}, %{{.*}}, %{{.*}}, %{{.*}}) {fastmathFlags = #llvm.fastmath<contract>} : (f32, f32, !llvm.metadata, !llvm.metadata) -> f32
+ %1 = cir.fadd %0, %b : !cir.float {fenv = #cir.fenv<>, fastmath_flags = #cir.fastmath<contract>}
+ cir.return
+ }
}
diff --git a/clang/utils/TableGen/CIRLoweringEmitter.cpp b/clang/utils/TableGen/CIRLoweringEmitter.cpp
index 978f0484627b58..b533c239635a61 100644
--- a/clang/utils/TableGen/CIRLoweringEmitter.cpp
+++ b/clang/utils/TableGen/CIRLoweringEmitter.cpp
@@ -147,12 +147,14 @@ void GenerateABILoweringPattern(llvm::StringRef OpName,
CXXABILoweringPatterns.push_back(std::move(CodeBuffer));
}
-void GenerateLLVMLoweringPattern(
- llvm::StringRef OpName, llvm::StringRef PatternName, bool IsRecursive,
- llvm::StringRef ExtraDecl, const Record *CustomCtorRec,
- llvm::StringRef LLVMOp, llvm::StringRef ConstrainedLLVMIntrinsic,
- bool ConstrainedHasRoundingMode, bool PropagateFastMathFlags,
- bool HasZeroResult) {
+void GenerateLLVMLoweringPattern(llvm::StringRef OpName,
+ llvm::StringRef PatternName, bool IsRecursive,
+ llvm::StringRef ExtraDecl,
+ const Record *CustomCtorRec,
+ llvm::StringRef LLVMOp,
+ llvm::StringRef ConstrainedLLVMIntrinsic,
+ bool ConstrainedHasRoundingMode,
+ bool HasFastMathFlags, bool HasZeroResult) {
std::optional<CustomLoweringCtor> CustomCtor =
parseCustomLoweringCtor(CustomCtorRec);
std::string CodeBuffer;
@@ -212,6 +214,10 @@ void GenerateLLVMLoweringPattern(
Code << " return lowerConstrainableFPOp<mlir::LLVM::" << LLVMOp
<< ">(\n";
Code << " op, adaptor.getOperands(), op.getFenvAttr(),\n";
+ Code << " "
+ << (HasFastMathFlags ? "op.getFastmathFlagsAttr()"
+ : "cir::FastMathFlagsAttr{}")
+ << ",\n";
Code << " *getTypeConverter(), rewriter, \""
<< ConstrainedLLVMIntrinsic << "\",\n";
Code << " /*hasRoundingMode=*/"
@@ -230,15 +236,6 @@ void GenerateLLVMLoweringPattern(
Code << " rewriter.replaceOpWithNewOp<mlir::LLVM::" << LLVMOp
<< ">(op, mlir::TypeRange{}, adaptor.getOperands(), " << Properties
<< ");\n";
- } else if (PropagateFastMathFlags) {
- Code << " mlir::Type resTy = "
- "typeConverter->convertType(op.getType());\n";
- Code << " auto lowered = mlir::LLVM::" << LLVMOp
- << "::create(rewriter, op.getLoc(), mlir::TypeRange{resTy}, "
- "adaptor.getOperands(), "
- << Properties << ");\n";
- Code << " propagateFastMathFlags(op, lowered);\n";
- Code << " rewriter.replaceOp(op, lowered.getResult());\n";
} else {
Code << " mlir::Type resTy = "
"typeConverter->convertType(op.getType());\n";
@@ -287,8 +284,11 @@ void Generate(const Record *OpRecord) {
OpRecord->getValueAsString("constrainedLLVMIntrinsic");
bool ConstrainedHasRoundingMode =
OpRecord->getValueAsBit("constrainedLLVMIntrinsicHasRoundingMode");
- bool PropagateFastMathFlags =
- OpRecord->getValueAsBit("propagateFastMathFlags");
+ const DagInit *ArgsDag = OpRecord->getValueAsDag("arguments");
+ bool HasFastMathFlags =
+ llvm::any_of(ArgsDag->getArgNames(), [](const StringInit *Name) {
+ return Name && Name->getValue() == "fastmath_flags";
+ });
if (!LLVMOp.empty() && CustomCtor)
PrintFatalError(OpRecord->getLoc(),
@@ -306,8 +306,8 @@ void Generate(const Record *OpRecord) {
bool IsZeroResult = ResultsDag->getNumArgs() == 0;
GenerateLLVMLoweringPattern(OpName, PatternName, IsRecursive, ExtraDecl,
CustomCtor, LLVMOp, ConstrainedLLVMIntrinsic,
- ConstrainedHasRoundingMode,
- PropagateFastMathFlags, IsZeroResult);
+ ConstrainedHasRoundingMode, HasFastMathFlags,
+ IsZeroResult);
// Only automatically register patterns that use the default constructor.
// Patterns with a custom constructor must be manually registered by the
// lowering pass.
More information about the cfe-commits
mailing list