[llvm] [SPIRV] Implement NaN propation for FMINIMUM and FMAXIMUM. (PR #180797)
via llvm-commits
llvm-commits at lists.llvm.org
Tue Feb 10 10:01:41 PST 2026
llvmbot wrote:
<!--LLVM PR SUMMARY COMMENT-->
@llvm/pr-subscribers-backend-spir-v
Author: Faijul Amin (mdfaijul)
<details>
<summary>Changes</summary>
The LLVM intrinsic for float minimum and maximum, e.g., `llvm.{maximum,minimum}.f*` has NaN propagating semantic. This PR implements NaN propagating semantic using `SPIRVLegalizerInfo` and `SPIRVPostLegalizer`
---
Full diff: https://github.com/llvm/llvm-project/pull/180797.diff
2 Files Affected:
- (modified) llvm/lib/Target/SPIRV/SPIRVLegalizerInfo.cpp (+14-2)
- (modified) llvm/lib/Target/SPIRV/SPIRVPostLegalizer.cpp (+52)
``````````diff
diff --git a/llvm/lib/Target/SPIRV/SPIRVLegalizerInfo.cpp b/llvm/lib/Target/SPIRV/SPIRVLegalizerInfo.cpp
index dc5906cfa9ceb..607f0e18cdb00 100644
--- a/llvm/lib/Target/SPIRV/SPIRVLegalizerInfo.cpp
+++ b/llvm/lib/Target/SPIRV/SPIRVLegalizerInfo.cpp
@@ -464,12 +464,13 @@ SPIRVLegalizerInfo::SPIRVLegalizerInfo(const SPIRVSubtarget &ST) {
G_FNEARBYINT,
G_INTRINSIC_ROUND,
G_INTRINSIC_TRUNC,
- G_FMINIMUM,
- G_FMAXIMUM,
G_INTRINSIC_ROUNDEVEN})
.legalFor(allFloatScalarsAndVectors);
// clang-format on
+ getActionDefinitionsBuilder({G_FMINIMUM, G_FMAXIMUM})
+ .customFor(allFloatScalarsAndVectors);
+
getActionDefinitionsBuilder(G_FCOPYSIGN)
.legalForCartesianProduct(allFloatScalarsAndVectors,
allFloatScalarsAndVectors);
@@ -646,6 +647,14 @@ static bool legalizeStore(LegalizerHelper &Helper, MachineInstr &MI,
return true;
}
+// NaN progagation semantics for FMinimum and FMaximum provided by generic
+// lowering.
+static bool legalizeFMinimumMaximum(LegalizerHelper &Helper, MachineInstr &MI) {
+ // Generic lowering ignores SPIR-V types. They are assigned at PostLegalizer
+ // to intermediate registers by looking at their LLT.
+ return Helper.lowerFMinimumMaximum(MI) == LegalizerHelper::Legalized;
+}
+
bool SPIRVLegalizerInfo::legalizeCustom(
LegalizerHelper &Helper, MachineInstr &MI,
LostDebugLocObserver &LocObserver) const {
@@ -658,6 +667,9 @@ bool SPIRVLegalizerInfo::legalizeCustom(
return legalizeBitcast(Helper, MI);
case TargetOpcode::G_EXTRACT_VECTOR_ELT:
return legalizeExtractVectorElt(Helper, MI, GR);
+ case TargetOpcode::G_FMINIMUM:
+ case TargetOpcode::G_FMAXIMUM:
+ return legalizeFMinimumMaximum(Helper, MI);
case TargetOpcode::G_INSERT_VECTOR_ELT:
return legalizeInsertVectorElt(Helper, MI, GR);
case TargetOpcode::G_INTRINSIC:
diff --git a/llvm/lib/Target/SPIRV/SPIRVPostLegalizer.cpp b/llvm/lib/Target/SPIRV/SPIRVPostLegalizer.cpp
index 198654b0564b1..b113ab3bac1b3 100644
--- a/llvm/lib/Target/SPIRV/SPIRVPostLegalizer.cpp
+++ b/llvm/lib/Target/SPIRV/SPIRVPostLegalizer.cpp
@@ -52,6 +52,46 @@ static SPIRVType *deduceIntTypeFromResult(Register ResVReg,
return GR->getOrCreateSPIRVIntegerType(Ty.getScalarSizeInBits(), MIB);
}
+static SPIRVType *deduceFloatTypeFromResult(Register ResVReg,
+ MachineIRBuilder &MIB,
+ SPIRVGlobalRegistry *GR) {
+ const LLT &ResLLT = MIB.getMRI()->getType(ResVReg);
+ unsigned ScalarBits = ResLLT.getScalarSizeInBits();
+ LLVMContext &Ctx = MIB.getMF().getFunction().getContext();
+ Type *ScalarTy = nullptr;
+ switch (ScalarBits) {
+ case 16:
+ ScalarTy = Type::getHalfTy(Ctx);
+ break;
+ case 32:
+ ScalarTy = Type::getFloatTy(Ctx);
+ break;
+ case 64:
+ ScalarTy = Type::getDoubleTy(Ctx);
+ break;
+ default:
+ return nullptr;
+ }
+
+ SPIRVType *SpvScalarType = GR->getOrCreateSPIRVType(
+ ScalarTy, MIB, SPIRV::AccessQualifier::ReadWrite, true);
+ if (ResLLT.isVector())
+ return GR->getOrCreateSPIRVVectorType(SpvScalarType,
+ ResLLT.getNumElements(), MIB, false);
+ return SpvScalarType;
+}
+
+static SPIRVType *deduceBoolTypeFromResult(Register ResVReg,
+ MachineIRBuilder &MIB,
+ SPIRVGlobalRegistry *GR) {
+ const LLT &ResLLT = MIB.getMRI()->getType(ResVReg);
+ SPIRVType *SpvBoolType = GR->getOrCreateSPIRVBoolType(MIB, true);
+ if (ResLLT.isVector())
+ return GR->getOrCreateSPIRVVectorType(SpvBoolType, ResLLT.getNumElements(),
+ MIB, false);
+ return SpvBoolType;
+}
+
static SPIRVType *deduceTypeFromSingleOperand(MachineInstr *I,
MachineIRBuilder &MIB,
SPIRVGlobalRegistry *GR,
@@ -272,6 +312,11 @@ static SPIRVType *deduceResultTypeFromOperands(MachineInstr *I,
MachineIRBuilder &MIB) {
Register ResVReg = I->getOperand(0).getReg();
switch (I->getOpcode()) {
+ case TargetOpcode::G_FCMP:
+ case TargetOpcode::G_ICMP:
+ return deduceBoolTypeFromResult(ResVReg, MIB, GR);
+ case TargetOpcode::G_FCONSTANT:
+ return deduceFloatTypeFromResult(ResVReg, MIB, GR);
case TargetOpcode::G_CONSTANT:
case TargetOpcode::G_ANYEXT:
case TargetOpcode::G_SEXT:
@@ -281,6 +326,8 @@ static SPIRVType *deduceResultTypeFromOperands(MachineInstr *I,
return deduceTypeFromOperandRange(I, MIB, GR, 1, I->getNumOperands());
case TargetOpcode::G_SHUFFLE_VECTOR:
return deduceTypeFromOperandRange(I, MIB, GR, 1, 3);
+ case TargetOpcode::G_SELECT:
+ return deduceTypeFromOperandRange(I, MIB, GR, 2, 4);
case TargetOpcode::G_INTRINSIC_W_SIDE_EFFECTS:
case TargetOpcode::G_INTRINSIC: {
auto IntrinsicID = cast<GIntrinsic>(I)->getIntrinsicID();
@@ -381,6 +428,11 @@ static bool requiresSpirvType(MachineInstr &I, SPIRVGlobalRegistry *GR,
MachineRegisterInfo &MRI) {
LLVM_DEBUG(dbgs() << "Checking if instruction requires a SPIR-V type: "
<< I;);
+ // COPY does not have a property of PreIselOpCode. However, it needs an SPIR-V
+ // type for consumer instructions required by the SPIRVInstructionSelector.
+ if (I.getOpcode() == TargetOpcode::COPY)
+ return true;
+
if (I.getNumDefs() == 0) {
LLVM_DEBUG(dbgs() << "Instruction does not have a definition.\n");
return false;
``````````
</details>
https://github.com/llvm/llvm-project/pull/180797
More information about the llvm-commits
mailing list