[llvm] [NVPTX] Constant fold blockDim when reqntid is specified (PR #191575)

via llvm-commits llvm-commits at lists.llvm.org
Fri Apr 10 16:49:15 PDT 2026


https://github.com/Chengjunp updated https://github.com/llvm/llvm-project/pull/191575

>From 6fcb5043986e1c88c2a71ffd9d8b0944506b03e1 Mon Sep 17 00:00:00 2001
From: chengjunp <chengjunp at nvidia.com>
Date: Fri, 10 Apr 2026 23:36:10 +0000
Subject: [PATCH 1/2] Extend NVVMIntrRange to const fold reqntid

---
 llvm/lib/Target/NVPTX/NVVMIntrRange.cpp       | 54 ++++++++++---
 llvm/lib/Target/NVPTX/NVVMProperties.cpp      |  2 +-
 llvm/lib/Target/NVPTX/NVVMProperties.h        |  1 +
 llvm/test/CodeGen/NVPTX/intr-range.ll         | 15 ++--
 llvm/test/CodeGen/NVPTX/reqntid-const-fold.ll | 78 +++++++++++++++++++
 5 files changed, 129 insertions(+), 21 deletions(-)
 create mode 100644 llvm/test/CodeGen/NVPTX/reqntid-const-fold.ll

diff --git a/llvm/lib/Target/NVPTX/NVVMIntrRange.cpp b/llvm/lib/Target/NVPTX/NVVMIntrRange.cpp
index 2fd305efa57cf..5d4cb157aaea0 100644
--- a/llvm/lib/Target/NVPTX/NVVMIntrRange.cpp
+++ b/llvm/lib/Target/NVPTX/NVVMIntrRange.cpp
@@ -57,6 +57,14 @@ static bool addRangeAttr(uint64_t Low, uint64_t High, IntrinsicInst *II) {
   return true;
 }
 
+// Replace all uses of the intrinsic call with a constant value.
+static bool replaceWithConstant(uint64_t Val, IntrinsicInst *II, SmallVector<IntrinsicInst *, 8> &ToErase) {
+  Constant *C = ConstantInt::get(II->getType(), Val);
+  II->replaceAllUsesWith(C);
+  ToErase.push_back(II);
+  return true;
+}
+
 static bool runNVVMIntrRange(Function &F) {
   struct Vector3 {
     unsigned X, Y, Z;
@@ -66,7 +74,8 @@ static bool runNVVMIntrRange(Function &F) {
   if (!isKernelFunction(F))
     return false;
 
-  const auto OverallReqNTID = getOverallReqNTID(F);
+  const auto ReqNTID = getReqNTID(F);
+  const auto OverallReqNTID = getVectorProduct(ReqNTID);
   const auto OverallMaxNTID = getOverallMaxNTID(F);
   const auto OverallClusterRank = getOverallClusterRank(F);
 
@@ -90,23 +99,44 @@ static bool runNVVMIntrRange(Function &F) {
                                std::min(0xffffu, FunctionClusterRank),
                                std::min(0xffffu, FunctionClusterRank)};
 
-  const auto ProccessIntrinsic = [&](IntrinsicInst *II) -> bool {
+  // When reqntid is specified, ntid (blockDim) values are exact compile-time
+  // constants. Get per-dimension values for constant folding.
+  const Vector3 ReqBlockDim =
+      !ReqNTID.empty()
+          ? Vector3{ReqNTID[0], ReqNTID.size() > 1 ? ReqNTID[1] : 1,
+                    ReqNTID.size() > 2 ? ReqNTID[2] : 1}
+          : Vector3{};
+  // Only fold when reqntid values are within hardware limits.
+  const bool HasValidReqNTID =
+      (ReqBlockDim.X >= 1 && ReqBlockDim.X <= 1024) &&
+      (ReqBlockDim.Y >= 1 && ReqBlockDim.Y <= 1024) &&
+      (ReqBlockDim.Z >= 1 && ReqBlockDim.Z <= 64) &&
+      (ReqBlockDim.X * ReqBlockDim.Y * ReqBlockDim.Z <= 1024);
+
+  const auto ProcessIntrinsic =
+      [&](IntrinsicInst *II, SmallVector<IntrinsicInst *, 8> &ToErase) -> bool {
     switch (II->getIntrinsicID()) {
     // Index within block
     case Intrinsic::nvvm_read_ptx_sreg_tid_x:
-      return addRangeAttr(0, MaxBlockSize.X, II);
+      return addRangeAttr(0, HasValidReqNTID ? ReqBlockDim.X : MaxBlockSize.X,
+                          II);
     case Intrinsic::nvvm_read_ptx_sreg_tid_y:
-      return addRangeAttr(0, MaxBlockSize.Y, II);
+      return addRangeAttr(0, HasValidReqNTID ? ReqBlockDim.Y : MaxBlockSize.Y,
+                          II);
     case Intrinsic::nvvm_read_ptx_sreg_tid_z:
-      return addRangeAttr(0, MaxBlockSize.Z, II);
+      return addRangeAttr(0, HasValidReqNTID ? ReqBlockDim.Z : MaxBlockSize.Z,
+                          II);
 
-    // Block size
+    // Block size: replace with constants when reqntid is specified.
     case Intrinsic::nvvm_read_ptx_sreg_ntid_x:
-      return addRangeAttr(1, MaxBlockSize.X + 1, II);
+      return HasValidReqNTID ? replaceWithConstant(ReqBlockDim.X, II, ToErase)
+                             : addRangeAttr(1, MaxBlockSize.X + 1, II);
     case Intrinsic::nvvm_read_ptx_sreg_ntid_y:
-      return addRangeAttr(1, MaxBlockSize.Y + 1, II);
+      return HasValidReqNTID ? replaceWithConstant(ReqBlockDim.Y, II, ToErase)
+                             : addRangeAttr(1, MaxBlockSize.Y + 1, II);
     case Intrinsic::nvvm_read_ptx_sreg_ntid_z:
-      return addRangeAttr(1, MaxBlockSize.Z + 1, II);
+      return HasValidReqNTID ? replaceWithConstant(ReqBlockDim.Z, II, ToErase)
+                             : addRangeAttr(1, MaxBlockSize.Z + 1, II);
 
     // Cluster size
     case Intrinsic::nvvm_read_ptx_sreg_cluster_ctaid_x:
@@ -138,10 +168,12 @@ static bool runNVVMIntrRange(Function &F) {
 
   // Go through the calls in this function.
   bool Changed = false;
+  SmallVector<IntrinsicInst *, 8> ToErase;
   for (Instruction &I : instructions(F))
     if (IntrinsicInst *II = dyn_cast<IntrinsicInst>(&I))
-      Changed |= ProccessIntrinsic(II);
-
+      Changed |= ProcessIntrinsic(II, ToErase);
+  for (auto *II : ToErase)
+    II->eraseFromParent();
   return Changed;
 }
 
diff --git a/llvm/lib/Target/NVPTX/NVVMProperties.cpp b/llvm/lib/Target/NVPTX/NVVMProperties.cpp
index d68c5aaf4fe5f..5bc70c790d249 100644
--- a/llvm/lib/Target/NVPTX/NVVMProperties.cpp
+++ b/llvm/lib/Target/NVPTX/NVVMProperties.cpp
@@ -201,7 +201,7 @@ static SmallVector<unsigned, 3> getFnAttrParsedVector(const Function &F,
   return V;
 }
 
-static std::optional<uint64_t> getVectorProduct(ArrayRef<unsigned> V) {
+std::optional<uint64_t> getVectorProduct(ArrayRef<unsigned> V) {
   if (V.empty())
     return std::nullopt;
 
diff --git a/llvm/lib/Target/NVPTX/NVVMProperties.h b/llvm/lib/Target/NVPTX/NVVMProperties.h
index 6ccd6f8a20075..37475afdcd3b1 100644
--- a/llvm/lib/Target/NVPTX/NVVMProperties.h
+++ b/llvm/lib/Target/NVPTX/NVVMProperties.h
@@ -47,6 +47,7 @@ SmallVector<unsigned, 3> getMaxNTID(const Function &);
 SmallVector<unsigned, 3> getReqNTID(const Function &);
 SmallVector<unsigned, 3> getClusterDim(const Function &);
 
+std::optional<uint64_t> getVectorProduct(ArrayRef<unsigned> V);
 std::optional<uint64_t> getOverallMaxNTID(const Function &);
 std::optional<uint64_t> getOverallReqNTID(const Function &);
 std::optional<uint64_t> getOverallClusterRank(const Function &);
diff --git a/llvm/test/CodeGen/NVPTX/intr-range.ll b/llvm/test/CodeGen/NVPTX/intr-range.ll
index 48fa3e06629b4..726a82be1fbbc 100644
--- a/llvm/test/CodeGen/NVPTX/intr-range.ll
+++ b/llvm/test/CodeGen/NVPTX/intr-range.ll
@@ -36,17 +36,14 @@ define ptx_kernel i32 @test_reqntid() "nvvm.reqntid"="20" {
 ; CHECK-LABEL: define ptx_kernel i32 @test_reqntid(
 ; CHECK-SAME: ) #[[ATTR1:[0-9]+]] {
 ; CHECK-NEXT:    [[TMP1:%.*]] = call range(i32 0, 20) i32 @llvm.nvvm.read.ptx.sreg.tid.x()
-; CHECK-NEXT:    [[TMP5:%.*]] = call range(i32 0, 20) i32 @llvm.nvvm.read.ptx.sreg.tid.y()
-; CHECK-NEXT:    [[TMP2:%.*]] = call range(i32 0, 20) i32 @llvm.nvvm.read.ptx.sreg.tid.z()
-; CHECK-NEXT:    [[TMP4:%.*]] = call range(i32 1, 21) i32 @llvm.nvvm.read.ptx.sreg.ntid.x()
-; CHECK-NEXT:    [[TMP3:%.*]] = call range(i32 1, 21) i32 @llvm.nvvm.read.ptx.sreg.ntid.y()
-; CHECK-NEXT:    [[TMP6:%.*]] = call range(i32 1, 21) i32 @llvm.nvvm.read.ptx.sreg.ntid.z()
+; CHECK-NEXT:    [[TMP5:%.*]] = call range(i32 0, 1) i32 @llvm.nvvm.read.ptx.sreg.tid.y()
+; CHECK-NEXT:    [[TMP2:%.*]] = call range(i32 0, 1) i32 @llvm.nvvm.read.ptx.sreg.tid.z()
 ; CHECK-NEXT:    [[TMP7:%.*]] = add i32 [[TMP1]], [[TMP5]]
 ; CHECK-NEXT:    [[TMP8:%.*]] = add i32 [[TMP7]], [[TMP2]]
-; CHECK-NEXT:    [[TMP9:%.*]] = add i32 [[TMP8]], [[TMP4]]
-; CHECK-NEXT:    [[TMP10:%.*]] = add i32 [[TMP9]], [[TMP3]]
-; CHECK-NEXT:    [[TMP11:%.*]] = add i32 [[TMP10]], [[TMP6]]
-; CHECK-NEXT:    ret i32 [[TMP3]]
+; CHECK-NEXT:    [[TMP9:%.*]] = add i32 [[TMP8]], 20
+; CHECK-NEXT:    [[TMP10:%.*]] = add i32 [[TMP9]], 1
+; CHECK-NEXT:    [[TMP11:%.*]] = add i32 [[TMP10]], 1
+; CHECK-NEXT:    ret i32 1
 ;
   %1 = call i32 @llvm.nvvm.read.ptx.sreg.tid.x()
   %2 = call i32 @llvm.nvvm.read.ptx.sreg.tid.y()
diff --git a/llvm/test/CodeGen/NVPTX/reqntid-const-fold.ll b/llvm/test/CodeGen/NVPTX/reqntid-const-fold.ll
new file mode 100644
index 0000000000000..7d00312daa02e
--- /dev/null
+++ b/llvm/test/CodeGen/NVPTX/reqntid-const-fold.ll
@@ -0,0 +1,78 @@
+; RUN: opt < %s -S -mtriple=nvptx-nvidia-cuda -mcpu=sm_20 -passes=nvvm-intr-range | FileCheck %s
+
+; When .reqntid specifies 3D dimensions, ntid.x/y/z should be replaced with
+; constants and tid.x/y/z should get per-dimension ranges.
+; Product 128*4*2 = 1024 is within the hardware limit.
+define ptx_kernel i32 @test_reqntid_3d() "nvvm.reqntid"="128,4,2" {
+; CHECK-LABEL: define ptx_kernel i32 @test_reqntid_3d(
+; CHECK-NEXT:    [[TID_X:%.*]] = call range(i32 0, 128) i32 @llvm.nvvm.read.ptx.sreg.tid.x()
+; CHECK-NEXT:    [[TID_Y:%.*]] = call range(i32 0, 4) i32 @llvm.nvvm.read.ptx.sreg.tid.y()
+; CHECK-NEXT:    [[TID_Z:%.*]] = call range(i32 0, 2) i32 @llvm.nvvm.read.ptx.sreg.tid.z()
+; CHECK-NEXT:    [[A:%.*]] = add i32 [[TID_X]], [[TID_Y]]
+; CHECK-NEXT:    [[B:%.*]] = add i32 [[A]], [[TID_Z]]
+; CHECK-NEXT:    [[C:%.*]] = add i32 [[B]], 128
+; CHECK-NEXT:    [[D:%.*]] = add i32 [[C]], 4
+; CHECK-NEXT:    [[E:%.*]] = add i32 [[D]], 2
+; CHECK-NEXT:    ret i32 [[E]]
+;
+  %tid.x = call i32 @llvm.nvvm.read.ptx.sreg.tid.x()
+  %tid.y = call i32 @llvm.nvvm.read.ptx.sreg.tid.y()
+  %tid.z = call i32 @llvm.nvvm.read.ptx.sreg.tid.z()
+  %ntid.x = call i32 @llvm.nvvm.read.ptx.sreg.ntid.x()
+  %ntid.y = call i32 @llvm.nvvm.read.ptx.sreg.ntid.y()
+  %ntid.z = call i32 @llvm.nvvm.read.ptx.sreg.ntid.z()
+  %a = add i32 %tid.x, %tid.y
+  %b = add i32 %a, %tid.z
+  %c = add i32 %b, %ntid.x
+  %d = add i32 %c, %ntid.y
+  %e = add i32 %d, %ntid.z
+  ret i32 %e
+}
+
+; When .reqntid specifies only 1D, y and z default to 1.
+define ptx_kernel i32 @test_reqntid_1d() "nvvm.reqntid"="128" {
+; CHECK-LABEL: define ptx_kernel i32 @test_reqntid_1d(
+; CHECK-NEXT:    [[TID_X:%.*]] = call range(i32 0, 128) i32 @llvm.nvvm.read.ptx.sreg.tid.x()
+; CHECK-NEXT:    [[TID_Y:%.*]] = call range(i32 0, 1) i32 @llvm.nvvm.read.ptx.sreg.tid.y()
+; CHECK-NEXT:    [[TID_Z:%.*]] = call range(i32 0, 1) i32 @llvm.nvvm.read.ptx.sreg.tid.z()
+; CHECK-NEXT:    [[A:%.*]] = add i32 [[TID_X]], [[TID_Y]]
+; CHECK-NEXT:    [[B:%.*]] = add i32 [[A]], [[TID_Z]]
+; CHECK-NEXT:    [[C:%.*]] = add i32 [[B]], 128
+; CHECK-NEXT:    [[D:%.*]] = add i32 [[C]], 1
+; CHECK-NEXT:    [[E:%.*]] = add i32 [[D]], 1
+; CHECK-NEXT:    ret i32 [[E]]
+;
+  %tid.x = call i32 @llvm.nvvm.read.ptx.sreg.tid.x()
+  %tid.y = call i32 @llvm.nvvm.read.ptx.sreg.tid.y()
+  %tid.z = call i32 @llvm.nvvm.read.ptx.sreg.tid.z()
+  %ntid.x = call i32 @llvm.nvvm.read.ptx.sreg.ntid.x()
+  %ntid.y = call i32 @llvm.nvvm.read.ptx.sreg.ntid.y()
+  %ntid.z = call i32 @llvm.nvvm.read.ptx.sreg.ntid.z()
+  %a = add i32 %tid.x, %tid.y
+  %b = add i32 %a, %tid.z
+  %c = add i32 %b, %ntid.x
+  %d = add i32 %c, %ntid.y
+  %e = add i32 %d, %ntid.z
+  ret i32 %e
+}
+
+; When .reqntid exceeds hardware limits, no folding — fall back to range attrs.
+define ptx_kernel i32 @test_reqntid_invalid() "nvvm.reqntid"="2048" {
+; CHECK-LABEL: define ptx_kernel i32 @test_reqntid_invalid(
+; CHECK-NEXT:    [[TID_X:%.*]] = call range(i32 0, 1024) i32 @llvm.nvvm.read.ptx.sreg.tid.x()
+; CHECK-NEXT:    [[NTID_X:%.*]] = call range(i32 1, 1025) i32 @llvm.nvvm.read.ptx.sreg.ntid.x()
+; CHECK-NEXT:    [[A:%.*]] = add i32 [[TID_X]], [[NTID_X]]
+; CHECK-NEXT:    ret i32 [[A]]
+;
+  %tid.x = call i32 @llvm.nvvm.read.ptx.sreg.tid.x()
+  %ntid.x = call i32 @llvm.nvvm.read.ptx.sreg.ntid.x()
+  %a = add i32 %tid.x, %ntid.x
+  ret i32 %a
+}
+
+declare i32 @llvm.nvvm.read.ptx.sreg.tid.x()
+declare i32 @llvm.nvvm.read.ptx.sreg.tid.y()
+declare i32 @llvm.nvvm.read.ptx.sreg.tid.z()
+declare i32 @llvm.nvvm.read.ptx.sreg.ntid.x()
+declare i32 @llvm.nvvm.read.ptx.sreg.ntid.y()
+declare i32 @llvm.nvvm.read.ptx.sreg.ntid.z()

>From 1a027a03f7ce3a1333d443515a7c75438b39dcf6 Mon Sep 17 00:00:00 2001
From: chengjunp <chengjunp at nvidia.com>
Date: Fri, 10 Apr 2026 23:49:02 +0000
Subject: [PATCH 2/2] Format

---
 llvm/lib/Target/NVPTX/NVVMIntrRange.cpp | 3 ++-
 1 file changed, 2 insertions(+), 1 deletion(-)

diff --git a/llvm/lib/Target/NVPTX/NVVMIntrRange.cpp b/llvm/lib/Target/NVPTX/NVVMIntrRange.cpp
index 5d4cb157aaea0..1e7966fb9aec2 100644
--- a/llvm/lib/Target/NVPTX/NVVMIntrRange.cpp
+++ b/llvm/lib/Target/NVPTX/NVVMIntrRange.cpp
@@ -58,7 +58,8 @@ static bool addRangeAttr(uint64_t Low, uint64_t High, IntrinsicInst *II) {
 }
 
 // Replace all uses of the intrinsic call with a constant value.
-static bool replaceWithConstant(uint64_t Val, IntrinsicInst *II, SmallVector<IntrinsicInst *, 8> &ToErase) {
+static bool replaceWithConstant(uint64_t Val, IntrinsicInst *II,
+                                SmallVector<IntrinsicInst *, 8> &ToErase) {
   Constant *C = ConstantInt::get(II->getType(), Val);
   II->replaceAllUsesWith(C);
   ToErase.push_back(II);



More information about the llvm-commits mailing list