[llvm] [NVPTX] Add support for ldmatrix extensions introduced in PTX 9.4 (PR #224264)
via llvm-commits
llvm-commits at lists.llvm.org
Thu Sep 17 04:09:26 PDT 2026
llvmorg-github-actions[bot] wrote:
<!--LLVM PR SUMMARY COMMENT-->
@llvm/pr-subscribers-llvm-ir
Author: Dharuni R Acharya (DharuniRAcharya)
<details>
<summary>Changes</summary>
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
---
Full diff: https://github.com/llvm/llvm-project/pull/224264.diff
6 Files Affected:
- (modified) llvm/include/llvm/IR/IntrinsicsNVVM.td (+3-2)
- (modified) llvm/lib/Target/NVPTX/NVPTXISelLowering.cpp (+6-3)
- (modified) llvm/lib/Target/NVPTX/NVPTXInstrInfo.td (+2)
- (modified) llvm/lib/Target/NVPTX/NVPTXIntrinsics.td (+5)
- (added) llvm/test/CodeGen/NVPTX/wmma-ptx94-sm90a.py (+14)
- (modified) llvm/test/CodeGen/NVPTX/wmma.py (+32-1)
``````````diff
diff --git a/llvm/include/llvm/IR/IntrinsicsNVVM.td b/llvm/include/llvm/IR/IntrinsicsNVVM.td
index ff12175180d149..d9b73418394ed6 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 c01506f94ad90a..a3fc7407e595fa 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 1ae80f4937fa88..3d435a4cee13cf 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 f54fd1b348af85..89a9adf4426d80 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-sm90a.py b/llvm/test/CodeGen/NVPTX/wmma-ptx94-sm90a.py
new file mode 100644
index 00000000000000..cbffa8d000de70
--- /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 2bd0796c68f52d..e645dc0d8035c1 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,8 @@ 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 +646,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 +1877,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
``````````
</details>
https://github.com/llvm/llvm-project/pull/224264
More information about the llvm-commits
mailing list