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
