[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