[llvm] fe62ed5 - [NVPTX] Add support for ldmatrix extensions introduced in PTX 9.4 (#224264)

via llvm-commits llvm-commits at lists.llvm.org
Mon Sep 21 00:01:27 PDT 2026


Author: Dharuni R Acharya
Date: 2026-09-21T09:01:21+02:00
New Revision: fe62ed5134de05e1ac19ce914b15797c95eee0dd

URL: https://github.com/llvm/llvm-project/commit/fe62ed5134de05e1ac19ce914b15797c95eee0dd
DIFF: https://github.com/llvm/llvm-project/commit/fe62ed5134de05e1ac19ce914b15797c95eee0dd.diff

LOG: [NVPTX] Add support for ldmatrix extensions introduced in PTX 9.4 (#224264)

This patch adds support for `.s8.s4` types for `ldmatrix` instruction 
for `.m8n16` shape introduced in PTX 9.4.

PTX ISA Reference:
https://docs.nvidia.com/cuda/parallel-thread-execution/index.html#warp-level-matrix-instructions-ldmatrix

---------

Signed-off-by: DharuniRAcharya <dharunira at nvidia.com>

Added: 
    llvm/test/CodeGen/NVPTX/wmma-ptx94-sm100f.py
    llvm/test/CodeGen/NVPTX/wmma-ptx94-sm110f.py
    llvm/test/CodeGen/NVPTX/wmma-ptx94-sm120f.py
    llvm/test/CodeGen/NVPTX/wmma-ptx94-sm90a.py

Modified: 
    llvm/include/llvm/IR/IntrinsicsNVVM.td
    llvm/lib/Target/NVPTX/NVPTXISelLowering.cpp
    llvm/lib/Target/NVPTX/NVPTXInstrInfo.td
    llvm/lib/Target/NVPTX/NVPTXIntrinsics.td
    llvm/test/CodeGen/NVPTX/wmma.py

Removed: 
    


################################################################################
diff  --git a/llvm/include/llvm/IR/IntrinsicsNVVM.td b/llvm/include/llvm/IR/IntrinsicsNVVM.td
index ff12175180d14..d9b73418394ed 100644
--- a/llvm/include/llvm/IR/IntrinsicsNVVM.td
+++ b/llvm/include/llvm/IR/IntrinsicsNVVM.td
@@ -481,7 +481,7 @@ class WMMA_REGS<string Geom, string Frag, string PtxEltType, bit IsSparse = fals
       !eq(gf,"m16n16:x1") : !listsplat(llvm_i32_ty, 2),
       !eq(gf,"m16n16:x2") : !listsplat(llvm_i32_ty, 4),
 
-      // ldmatrix b8x16.b6x16_p32, b8x16.b4x16_p64 -> s32 @ m8n16
+      // ldmatrix b8x16.b6x16_p32, b8x16.b4x16_p64, s8.s4 -> s32 @ m8n16
       !eq(gf,"m8n16:x1") : !listsplat(llvm_i32_ty, 1),
       !eq(gf,"m8n16:x2") : !listsplat(llvm_i32_ty, 2),
       !eq(gf,"m8n16:x4") : !listsplat(llvm_i32_ty, 4),
@@ -834,7 +834,7 @@ class NVVM_MMA_OPS {
     ["m16n16"], ["x1", "x2"], ["b8", "b8x16.b6x16_p32", "b8x16.b4x16_p64"]>.ret;
 
   list<WMMA_REGS> ldmatrix_geom_m8n16_ops = LDMATRIX_OPS<
-    ["m8n16"], ["x1", "x2", "x4"], ["b8x16.b6x16_p32", "b8x16.b4x16_p64"]>.ret;
+    ["m8n16"], ["x1", "x2", "x4"], ["b8x16.b6x16_p32", "b8x16.b4x16_p64", "s8.s4"]>.ret;
 
   list<WMMA_REGS> stmatrix_b16_ops = STMATRIX_OPS<
     ["m8n8"], ["x1", "x2", "x4"], ["b16"]>.ret;
@@ -1033,6 +1033,7 @@ class NVVM_LDMATRIX_SUPPORTED<WMMA_REGS frag, bit trans> {
     !and(!eq(g, "m8n16"), !eq(t, "b8"), !eq(trans, 0)): true,
     !and(!eq(g, "m8n16"), !eq(t, "b8x16.b6x16_p32"), !eq(trans, 0)): true,
     !and(!eq(g, "m8n16"), !eq(t, "b8x16.b4x16_p64"), !eq(trans, 0)): true,
+    !and(!eq(g, "m8n16"), !eq(t, "s8.s4"), !eq(trans, 0)): true,
     true: false
   );
 }

diff  --git a/llvm/lib/Target/NVPTX/NVPTXISelLowering.cpp b/llvm/lib/Target/NVPTX/NVPTXISelLowering.cpp
index c01506f94ad90..a3fc7407e595f 100644
--- a/llvm/lib/Target/NVPTX/NVPTXISelLowering.cpp
+++ b/llvm/lib/Target/NVPTX/NVPTXISelLowering.cpp
@@ -4424,7 +4424,8 @@ void NVPTXTargetLowering::getTgtMemIntrinsic(
   case Intrinsic::nvvm_ldmatrix_sync_aligned_m16n16_x2_trans_b8x16_b4x16_p64:
   case Intrinsic::nvvm_ldmatrix_sync_aligned_m16n16_x2_trans_b8x16_b6x16_p32:
   case Intrinsic::nvvm_ldmatrix_sync_aligned_m8n16_x4_b8x16_b4x16_p64:
-  case Intrinsic::nvvm_ldmatrix_sync_aligned_m8n16_x4_b8x16_b6x16_p32: {
+  case Intrinsic::nvvm_ldmatrix_sync_aligned_m8n16_x4_b8x16_b6x16_p32:
+  case Intrinsic::nvvm_ldmatrix_sync_aligned_m8n16_x4_s8_s4: {
     Info.opc = ISD::INTRINSIC_W_CHAIN;
     Info.memVT = MVT::v4i32;
     Info.ptrVal = I.getArgOperand(0);
@@ -4467,7 +4468,8 @@ void NVPTXTargetLowering::getTgtMemIntrinsic(
   case Intrinsic::nvvm_ldmatrix_sync_aligned_m8n8_x1_b16:
   case Intrinsic::nvvm_ldmatrix_sync_aligned_m8n8_x1_trans_b16:
   case Intrinsic::nvvm_ldmatrix_sync_aligned_m8n16_x1_b8x16_b4x16_p64:
-  case Intrinsic::nvvm_ldmatrix_sync_aligned_m8n16_x1_b8x16_b6x16_p32: {
+  case Intrinsic::nvvm_ldmatrix_sync_aligned_m8n16_x1_b8x16_b6x16_p32:
+  case Intrinsic::nvvm_ldmatrix_sync_aligned_m8n16_x1_s8_s4: {
     Info.opc = ISD::INTRINSIC_W_CHAIN;
     Info.memVT = MVT::i32;
     Info.ptrVal = I.getArgOperand(0);
@@ -4572,7 +4574,8 @@ void NVPTXTargetLowering::getTgtMemIntrinsic(
   case Intrinsic::nvvm_ldmatrix_sync_aligned_m16n16_x1_trans_b8x16_b4x16_p64:
   case Intrinsic::nvvm_ldmatrix_sync_aligned_m16n16_x1_trans_b8x16_b6x16_p32:
   case Intrinsic::nvvm_ldmatrix_sync_aligned_m8n16_x2_b8x16_b4x16_p64:
-  case Intrinsic::nvvm_ldmatrix_sync_aligned_m8n16_x2_b8x16_b6x16_p32: {
+  case Intrinsic::nvvm_ldmatrix_sync_aligned_m8n16_x2_b8x16_b6x16_p32:
+  case Intrinsic::nvvm_ldmatrix_sync_aligned_m8n16_x2_s8_s4: {
     Info.opc = ISD::INTRINSIC_W_CHAIN;
     Info.memVT = MVT::v2i32;
     Info.ptrVal = I.getArgOperand(0);

diff  --git a/llvm/lib/Target/NVPTX/NVPTXInstrInfo.td b/llvm/lib/Target/NVPTX/NVPTXInstrInfo.td
index c130a2f351ff3..f794002971905 100644
--- a/llvm/lib/Target/NVPTX/NVPTXInstrInfo.td
+++ b/llvm/lib/Target/NVPTX/NVPTXInstrInfo.td
@@ -258,6 +258,8 @@ def hasSetMaxNRegSupport : PredOr<[SM90a, SM100f, SM110f, SM120f]>;
 
 def hasLdStmatrixBlackwellSupport : PredOr<[SM100f, SM110f, SM120f]>;
 
+def hasLdmatrixS8S4Support : PredAnd<[PTX94, PredOr<[SM90a, SM100f, SM110f, SM120f]>]>;
+
 def hasConvertWithStochasticRounding
   : PredAnd<[PTX87, PredOr<[SM100a, SM103a, SM107a]>]>;
 

diff  --git a/llvm/lib/Target/NVPTX/NVPTXIntrinsics.td b/llvm/lib/Target/NVPTX/NVPTXIntrinsics.td
index f54fd1b348af8..89a9adf4426d8 100644
--- a/llvm/lib/Target/NVPTX/NVPTXIntrinsics.td
+++ b/llvm/lib/Target/NVPTX/NVPTXIntrinsics.td
@@ -5904,6 +5904,7 @@ class WMMA_REGINFO<WMMA_REGS r, string op, string metadata = "",
     !eq(ptx_elt_type, "b8") : B32,
     !eq(ptx_elt_type, "b8x16.b6x16_p32") : B32,
     !eq(ptx_elt_type, "b8x16.b4x16_p64") : B32,
+    !eq(ptx_elt_type, "s8.s4") : B32,
     !eq(ptx_elt_type, "s8") : B32,
     !eq(ptx_elt_type, "u8") : B32,
     !eq(ptx_elt_type, "s4") : B32,
@@ -6053,6 +6054,10 @@ class WMMA_REGINFO<WMMA_REGS r, string op, string metadata = "",
          !eq(ptx_elt_type, "b8x16.b4x16_p64"),
          !eq(geom, "m8n16")) : [hasLdStmatrixBlackwellSupport],
 
+    !and(!eq(op, "ldmatrix"),
+         !eq(ptx_elt_type, "s8.s4"),
+         !eq(geom, "m8n16")) : [hasLdmatrixS8S4Support],
+
     !and(!eq(op, "stmatrix"),!eq(ptx_elt_type, "b16"),
          !eq(geom, "m8n8")) : [SM90],
 

diff  --git a/llvm/test/CodeGen/NVPTX/wmma-ptx94-sm100f.py b/llvm/test/CodeGen/NVPTX/wmma-ptx94-sm100f.py
new file mode 100644
index 0000000000000..036ab5eba147e
--- /dev/null
+++ b/llvm/test/CodeGen/NVPTX/wmma-ptx94-sm100f.py
@@ -0,0 +1,14 @@
+# Check all variants of instructions supported by PTX94 on SM100f
+# RUN: %python %s --ptx=94 --gpu-arch=100f > %t-ptx94-sm_100f.ll
+# RUN: FileCheck %t-ptx94-sm_100f.ll < %t-ptx94-sm_100f.ll \
+# RUN:           --check-prefixes=PTX94LDMATRIX-DAG
+# RUN: llc < %t-ptx94-sm_100f.ll -mtriple=nvptx64 -mcpu=sm_100f -mattr=+ptx94 \
+# RUN:           | FileCheck %t-ptx94-sm_100f.ll
+# RUN: %if ptxas-sm_100f && ptxas-isa-9.4 %{                                  \
+# RUN: llc < %t-ptx94-sm_100f.ll -mtriple=nvptx64 -mcpu=sm_100f -mattr=+ptx94 \
+# RUN:           | %ptxas-verify -arch=sm_100f                              \
+# RUN: %}
+
+import wmma
+
+wmma.main()

diff  --git a/llvm/test/CodeGen/NVPTX/wmma-ptx94-sm110f.py b/llvm/test/CodeGen/NVPTX/wmma-ptx94-sm110f.py
new file mode 100644
index 0000000000000..680412241bd26
--- /dev/null
+++ b/llvm/test/CodeGen/NVPTX/wmma-ptx94-sm110f.py
@@ -0,0 +1,14 @@
+# Check all variants of instructions supported by PTX94 on SM110f
+# RUN: %python %s --ptx=94 --gpu-arch=110f > %t-ptx94-sm_110f.ll
+# RUN: FileCheck %t-ptx94-sm_110f.ll < %t-ptx94-sm_110f.ll \
+# RUN:           --check-prefixes=PTX94LDMATRIX-DAG
+# RUN: llc < %t-ptx94-sm_110f.ll -mtriple=nvptx64 -mcpu=sm_110f -mattr=+ptx94 \
+# RUN:           | FileCheck %t-ptx94-sm_110f.ll
+# RUN: %if ptxas-sm_110f && ptxas-isa-9.4 %{                                  \
+# RUN: llc < %t-ptx94-sm_110f.ll -mtriple=nvptx64 -mcpu=sm_110f -mattr=+ptx94 \
+# RUN:           | %ptxas-verify -arch=sm_110f                              \
+# RUN: %}
+
+import wmma
+
+wmma.main()

diff  --git a/llvm/test/CodeGen/NVPTX/wmma-ptx94-sm120f.py b/llvm/test/CodeGen/NVPTX/wmma-ptx94-sm120f.py
new file mode 100644
index 0000000000000..60d45cb0b5e05
--- /dev/null
+++ b/llvm/test/CodeGen/NVPTX/wmma-ptx94-sm120f.py
@@ -0,0 +1,14 @@
+# Check all variants of instructions supported by PTX94 on SM120f
+# RUN: %python %s --ptx=94 --gpu-arch=120f > %t-ptx94-sm_120f.ll
+# RUN: FileCheck %t-ptx94-sm_120f.ll < %t-ptx94-sm_120f.ll \
+# RUN:           --check-prefixes=PTX94LDMATRIX-DAG
+# RUN: llc < %t-ptx94-sm_120f.ll -mtriple=nvptx64 -mcpu=sm_120f -mattr=+ptx94 \
+# RUN:           | FileCheck %t-ptx94-sm_120f.ll
+# RUN: %if ptxas-sm_120f && ptxas-isa-9.4 %{                                  \
+# RUN: llc < %t-ptx94-sm_120f.ll -mtriple=nvptx64 -mcpu=sm_120f -mattr=+ptx94 \
+# RUN:           | %ptxas-verify -arch=sm_120f                              \
+# RUN: %}
+
+import wmma
+
+wmma.main()

diff  --git a/llvm/test/CodeGen/NVPTX/wmma-ptx94-sm90a.py b/llvm/test/CodeGen/NVPTX/wmma-ptx94-sm90a.py
new file mode 100644
index 0000000000000..cbffa8d000de7
--- /dev/null
+++ b/llvm/test/CodeGen/NVPTX/wmma-ptx94-sm90a.py
@@ -0,0 +1,14 @@
+# Check all variants of instructions supported by PTX94 on SM90a
+# RUN: %python %s --ptx=94 --gpu-arch=90a > %t-ptx94-sm_90a.ll
+# RUN: FileCheck %t-ptx94-sm_90a.ll < %t-ptx94-sm_90a.ll \
+# RUN:           --check-prefixes=PTX94LDMATRIX-DAG
+# RUN: llc < %t-ptx94-sm_90a.ll -mtriple=nvptx64 -mcpu=sm_90a -mattr=+ptx94 \
+# RUN:           | FileCheck %t-ptx94-sm_90a.ll
+# RUN: %if ptxas-sm_90a && ptxas-isa-9.4 %{                                  \
+# RUN: llc < %t-ptx94-sm_90a.ll -mtriple=nvptx64 -mcpu=sm_90a -mattr=+ptx94 \
+# RUN:           | %ptxas-verify -arch=sm_90a                              \
+# RUN: %}
+
+import wmma
+
+wmma.main()

diff  --git a/llvm/test/CodeGen/NVPTX/wmma.py b/llvm/test/CodeGen/NVPTX/wmma.py
index 2bd0796c68f52..4bdc219dd28cb 100644
--- a/llvm/test/CodeGen/NVPTX/wmma.py
+++ b/llvm/test/CodeGen/NVPTX/wmma.py
@@ -35,6 +35,7 @@ def __init__(self, ptx_type):
             "b8": "i32",
             "b8x16.b6x16_p32": "i32",
             "b8x16.b4x16_p64": "i32",
+            "s8.s4": "i32",
             "s8": "i32",
             "u8": "i32",
             "s4": "i32",
@@ -257,6 +258,9 @@ def __init__(self, geom, frag, ptx_elt_type, is_mma_sparse=False):
             "m8n16:x1:b8x16.b4x16_p64": 1,
             "m8n16:x2:b8x16.b4x16_p64": 2,
             "m8n16:x4:b8x16.b4x16_p64": 4,
+            "m8n16:x1:s8.s4": 1,
+            "m8n16:x2:s8.s4": 2,
+            "m8n16:x4:s8.s4": 4,
             # stmatrix
             "m8n8:x1:b16": 1,
             "m8n8:x2:b16": 2,
@@ -421,7 +425,9 @@ def get_ldmatrix_ops():
             ["m16n16"], ["x1", "x2"], ["b8", "b8x16.b6x16_p32", "b8x16.b4x16_p64"]
         )
         + make_ldmatrix_ops(
-            ["m8n16"], ["x1", "x2", "x4"], ["b8x16.b6x16_p32", "b8x16.b4x16_p64"]
+            ["m8n16"],
+            ["x1", "x2", "x4"],
+            ["b8x16.b6x16_p32", "b8x16.b4x16_p64", "s8.s4"],
         )
     )
 
@@ -641,7 +647,26 @@ def is_ldst_variant_supported(frag, layout):
     return True
 
 
+def is_ldmatrix_s8s4_supported():
+    if ptx_version < 94:
+        return False
+    # sm_90a
+    if sm_version == 90 and has_arch_accel_features():
+        return True
+    # sm_100f / sm_110f / sm_120f families
+    if sm_version in [100, 110, 120] and has_family_specific_features():
+        return True
+    return False
+
+
 def is_ldmatrix_variant_supported(frag, trans):
+    if frag.mma_type.ptx_type == "s8.s4":
+        return (
+            frag.geom == "m8n16"
+            and trans == ""
+            and frag.frag in ["x1", "x2", "x4"]
+            and is_ldmatrix_s8s4_supported()
+        )
     if not (
         is_type_supported(frag.mma_type.ptx_type)
         and is_ldmatrix_geom_supported(frag.geom)
@@ -1853,6 +1878,13 @@ def gen_check_unsupported_ops(items):
 ; PTX86LDMATRIX-DAG: ldmatrix.sync.aligned.m8n16.x4.b8x16.b6x16_p32
 ; PTX86LDMATRIX-DAG: ldmatrix.sync.aligned.m8n16.x4.b8x16.b4x16_p64
 
+; PTX94LDMATRIX-DAG: ldmatrix.sync.aligned.m8n16.x1.s8.s4
+; PTX94LDMATRIX-DAG: ldmatrix.sync.aligned.m8n16.x2.s8.s4
+; PTX94LDMATRIX-DAG: ldmatrix.sync.aligned.m8n16.x4.s8.s4
+; PTX94LDMATRIX-DAG: ldmatrix.sync.aligned.m8n16.x1.shared.s8.s4
+; PTX94LDMATRIX-DAG: ldmatrix.sync.aligned.m8n16.x2.shared.s8.s4
+; PTX94LDMATRIX-DAG: ldmatrix.sync.aligned.m8n16.x4.shared.s8.s4
+
 ; PTX78STMATRIX-DAG: stmatrix.sync.aligned.m8n8.x1.b16
 ; PTX78STMATRIX-DAG: stmatrix.sync.aligned.m8n8.x2.b16
 ; PTX78STMATRIX-DAG: stmatrix.sync.aligned.m8n8.x4.b16


        


More information about the llvm-commits mailing list