[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
Thu Sep 24 18:53:36 PDT 2026
https://github.com/RiverDave updated https://github.com/llvm/llvm-project/pull/226334
>From 97d1d14c90250eca873d62411fcf1dad0a943a25 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/3] [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 c1ba78eea28354..9a83faff745d24 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{});
}]>
];
}
@@ -6084,7 +6109,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);
@@ -6097,7 +6123,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{});
}]>
];
@@ -7438,14 +7465,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{});
}]>
];
@@ -7706,6 +7735,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]> {
@@ -7760,7 +7790,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 = [{
@@ -7774,7 +7805,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{});
}]>
];
@@ -7829,7 +7861,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);
@@ -7841,7 +7874,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{});
}]>
];
@@ -7857,6 +7891,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"> {
@@ -7981,7 +8016,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);
@@ -7991,7 +8027,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 2b223c8ae19392..2d2ed0cceefbeb 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 bd529f71b38edf..1cc2fe68bf3799 100644
--- a/clang/lib/CIR/Lowering/DirectToLLVM/LowerToLLVM.cpp
+++ b/clang/lib/CIR/Lowering/DirectToLLVM/LowerToLLVM.cpp
@@ -534,6 +534,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) {
@@ -572,14 +624,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));
}
mlir::LogicalResult CIRToLLVMLLVMIntrinsicCallOpLowering::matchAndRewrite(
@@ -2051,13 +2106,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();
}
@@ -2065,13 +2122,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();
}
@@ -3573,7 +3632,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,
@@ -3584,7 +3644,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(
@@ -3619,14 +3679,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();
}
@@ -3662,6 +3724,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();
}
@@ -3680,6 +3744,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();
}
@@ -5014,14 +5080,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 88710e86238b34857b6f69b231895a89522626e8 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/3] [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 1cc2fe68bf3799..71042159166660 100644
--- a/clang/lib/CIR/Lowering/DirectToLLVM/LowerToLLVM.cpp
+++ b/clang/lib/CIR/Lowering/DirectToLLVM/LowerToLLVM.cpp
@@ -535,8 +535,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) ==
@@ -2110,11 +2109,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();
}
@@ -2126,11 +2125,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();
}
@@ -3628,12 +3627,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,
@@ -3685,8 +3682,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 138694b6d46c732c5d891545df7b8be1b832fe23 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/3] [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
More information about the cfe-commits
mailing list