https://github.com/RiverDave created 
https://github.com/llvm/llvm-project/pull/226334

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.

>From 6df97aa3ec547f4b27c4d471597b24c6f953bf99 Mon Sep 17 00:00:00 2001
From: David Rivera <[email protected]>
Date: Thu, 24 Sep 2026 20:54:26 -0400
Subject: [PATCH] [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 <[email protected]>
---
 .../CIR/Dialect/Builder/CIRBaseBuilder.h      |  38 ++++-
 .../clang/CIR/Dialect/IR/CIREnumAttr.td       |  31 ++++
 clang/include/clang/CIR/Dialect/IR/CIROps.td  |  71 ++++++---
 clang/lib/CIR/CodeGen/CIRGenBuiltin.cpp       |  22 ++-
 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 | 103 ++++++++++---
 .../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   |  16 +-
 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 45cb18a50eeee..d2f778004cece 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) {
@@ -841,32 +853,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) {
@@ -876,7 +895,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) {
@@ -892,7 +911,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,
@@ -906,7 +928,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 966ab85698325..f356e4d8d642c 100644
--- a/clang/include/clang/CIR/Dialect/IR/CIREnumAttr.td
+++ b/clang/include/clang/CIR/Dialect/IR/CIREnumAttr.td
@@ -47,6 +47,37 @@ class CIR_DefaultValuedEnumParameter<EnumAttrInfo info, 
string value = "">
   let defaultValue = 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_I32BitEnumAttr<
+    "FastMathFlags", "fast-math flags", [
+  I32BitEnumAttrCaseNone<"none">,
+  I32BitEnumAttrCaseBit<"nnan", 0>,
+  I32BitEnumAttrCaseBit<"ninf", 1>,
+  I32BitEnumAttrCaseBit<"nsz", 2>,
+  I32BitEnumAttrCaseBit<"arcp", 3>,
+  I32BitEnumAttrCaseBit<"contract", 4>,
+  I32BitEnumAttrCaseBit<"afn", 5>,
+  I32BitEnumAttrCaseBit<"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 = ", ";
+  let genSpecializedAttr = 0;
+}
+
+def CIR_FastMathFlagsAttr : CIR_EnumAttr<CIR_FastMathFlags, "fastmath"> {
+  let summary = "Fast-math flags for a floating-point operation";
+}
+
 def CIR_LangAddressSpace : CIR_I32EnumAttr<
   "LangAddressSpace", "language address space kind", [
   I32EnumAttrCase<"Default", 0, "default">,
diff --git a/clang/include/clang/CIR/Dialect/IR/CIROps.td 
b/clang/include/clang/CIR/Dialect/IR/CIROps.td
index 5c2c948742f08..ab3682285d49f 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> {
@@ -2100,18 +2104,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;
 }
 
 
//===----------------------------------------------------------------------===//
@@ -2564,7 +2581,8 @@ def CIR_CmpOp : CIR_Op<"cmp",
     CIR_CmpOpKind:$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);
@@ -2576,11 +2594,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{});
     }]>
   ];
 
@@ -2879,18 +2899,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{});
     }]>
   ];
 }
@@ -5922,7 +5947,8 @@ def CIR_VecCmpOp : CIR_Op<"vec.cmp",
     CIR_CmpOpKind:$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);
@@ -5935,7 +5961,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{});
     }]>
   ];
 
@@ -7206,14 +7233,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{});
     }]>
   ];
 
@@ -7474,6 +7503,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]> {
@@ -7528,7 +7558,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 = [{
@@ -7542,7 +7573,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{});
     }]>
   ];
 
@@ -7597,7 +7629,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);
@@ -7609,7 +7642,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{});
     }]>
   ];
 
@@ -7625,6 +7659,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"> {
@@ -7749,7 +7784,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);
@@ -7759,7 +7795,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 61cb04828e272..37dfa40decd3d 100644
--- a/clang/lib/CIR/CodeGen/CIRGenBuiltin.cpp
+++ b/clang/lib/CIR/CodeGen/CIRGenBuiltin.cpp
@@ -563,15 +563,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.getType(), 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));
 }
 
@@ -584,7 +587,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));
 }
 
@@ -593,9 +597,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));
 }
@@ -613,7 +619,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));
 }
 
@@ -629,7 +636,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 f7a39a9b4a8c9..cd00a8e59a0f2 100644
--- a/clang/lib/CIR/CodeGen/CIRGenExprScalar.cpp
+++ b/clang/lib/CIR/CodeGen/CIRGenExprScalar.cpp
@@ -815,8 +815,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
@@ -1197,6 +1200,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()) {
@@ -1206,9 +1212,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);
@@ -1986,9 +1998,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;
 }
@@ -2007,8 +2019,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 8301627ad8123..765a2dce86a3e 100644
--- a/clang/lib/CIR/CodeGen/CIRGenFunction.cpp
+++ b/clang/lib/CIR/CodeGen/CIRGenFunction.cpp
@@ -60,6 +60,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),
@@ -67,6 +77,7 @@ CIRGenFunction::CIRGenFunction(CIRGenModule &cgm, 
CIRGenBuilderTy &builder,
   ehStack.setCGF(this);
   shouldEmitLifetimeMarkers = CIRGen::shouldEmitLifetimeMarkers(
       cgm.getCodeGenOpts(), getContext().getLangOpts());
+  builder.setFastMathFlags(fastMathFlagsFromFPOptions(curFPFeatures));
 }
 
 CIRGenFunction::~CIRGenFunction() {}
@@ -1423,6 +1434,7 @@ void 
CIRGenFunction::CIRGenFPOptionsRAII::ConstructorHelper(
 
   oldExcept = cgf.builder.getDefaultConstrainedExcept();
   oldRounding = cgf.builder.getDefaultConstrainedRounding();
+  oldFastMathFlags = cgf.builder.getFastMathFlags();
 
   if (oldFPFeatures == fpFeatures)
     return;
@@ -1436,8 +1448,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() ||
@@ -1455,6 +1469,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 e738c3b1fb72d..0b5ea77ea2290 100644
--- a/clang/lib/CIR/CodeGen/CIRGenFunction.h
+++ b/clang/lib/CIR/CodeGen/CIRGenFunction.h
@@ -286,6 +286,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 c3838413984a6..ea45282530406 100644
--- a/clang/lib/CIR/Dialect/IR/CIRDialect.cpp
+++ b/clang/lib/CIR/Dialect/IR/CIRDialect.cpp
@@ -3573,6 +3573,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();
 }
 
@@ -3584,6 +3587,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 6f7509c363fd3..5472abcee71fd 100644
--- a/clang/lib/CIR/Lowering/DirectToLLVM/LowerToLLVM.cpp
+++ b/clang/lib/CIR/Lowering/DirectToLLVM/LowerToLLVM.cpp
@@ -420,6 +420,7 @@ mlir::Value lowerCirAttrAsValue(mlir::Operation *parentOp,
   return value;
 }
 
+
 void convertSideEffectForCall(mlir::Operation *callOp, bool isNothrow,
                               cir::SideEffect sideEffect,
                               mlir::LLVM::MemoryEffectsAttr &memoryEffect,
@@ -523,6 +524,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) {
@@ -561,12 +614,15 @@ mlir::LogicalResult lowerConstrainableFPOp(
     return op->emitError("expected LLVM result type for floating-point op");
 
   if (!fenv) {
-    rewriter.replaceOpWithNewOp<LLVMOp>(op, llvmResTy, operands);
+    LLVMOp llvmOp = LLVMOp::create(rewriter, op->getLoc(), llvmResTy, 
operands);
+    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(
@@ -2052,13 +2108,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();
 }
 
@@ -2066,13 +2124,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();
 }
 
@@ -3507,7 +3567,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,
@@ -3518,7 +3579,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(
@@ -3553,14 +3614,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();
   }
 
@@ -3596,6 +3659,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();
     }
@@ -3614,6 +3679,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();
     }
@@ -4915,14 +4982,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 146b31b907fcc..68d4b70fa618f 100644
--- a/clang/lib/CIR/Lowering/DirectToLLVM/LowerToLLVM.h
+++ b/clang/lib/CIR/Lowering/DirectToLLVM/LowerToLLVM.h
@@ -82,6 +82,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 0000000000000..9c1b57e6a1235
--- /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 4c04eb14c6309..e365d594f33c3 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 0000000000000..5b993e358f086
--- /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 0000000000000..0ece658a12f8d
--- /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 f67c37b9e1870..36bc1ad215871 100644
--- a/clang/utils/TableGen/CIRLoweringEmitter.cpp
+++ b/clang/utils/TableGen/CIRLoweringEmitter.cpp
@@ -138,7 +138,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;
@@ -212,6 +213,14 @@ void GenerateLLVMLoweringPattern(
     if (HasZeroResult) {
       Code << "    rewriter.replaceOpWithNewOp<mlir::LLVM::" << LLVMOp
            << ">(op, mlir::TypeRange{}, adaptor.getOperands());\n";
+    } else if (PropagateFastMathFlags) {
+      Code << "    mlir::Type resTy = "
+              "typeConverter->convertType(op.getType());\n";
+      Code << "    auto lowered = mlir::LLVM::" << LLVMOp
+           << "::create(rewriter, op.getLoc(), resTy, "
+              "adaptor.getOperands());\n";
+      Code << "    propagateFastMathFlags(op, lowered);\n";
+      Code << "    rewriter.replaceOp(op, lowered.getResult());\n";
     } else {
       Code << "    mlir::Type resTy = "
               "typeConverter->convertType(op.getType());\n";
@@ -259,6 +268,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(),
@@ -276,7 +287,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.

_______________________________________________
cfe-commits mailing list
[email protected]
https://lists.llvm.org/cgi-bin/mailman/listinfo/cfe-commits

Reply via email to