[llvm] [NFC][NVPTX] Introduce AS aware address matching in TableGen (PR #210309)
via llvm-commits
llvm-commits at lists.llvm.org
Fri Jul 17 04:39:34 PDT 2026
llvmorg-github-actions[bot] wrote:
<!--LLVM PR SUMMARY COMMENT-->
@llvm/pr-subscribers-backend-nvptx
Author: Rajat Bajpai (rajatbajpai)
<details>
<summary>Changes</summary>
Presently, there is no mechanism in NVPTX backend tablegen to match an intrinsic based on the address space of its pointer argument. The scenarios that have such a requirement (wmma, prefetch.tensormap, and tensormap.replace) currently roll their own PatFrags to match on address space. This problem would become more pronounced when PR #<!-- -->172442 lands, because it opens up more motivation for intrinsics overloaded on the pointer's address space.
This change introduces an `addr_as<AS>` ComplexPattern in tablegen, along with specialized helpers for easier use (`addr_generic`, `addr_global`, `addr_shared`, etc.).
These patterns require the corresponding intrinsics to be registered in `getTgtMemIntrinsic` so that the address space information is preserved in the SelectionDAG.
---
Full diff: https://github.com/llvm/llvm-project/pull/210309.diff
3 Files Affected:
- (modified) llvm/lib/Target/NVPTX/NVPTXISelDAGToDAG.h (+10)
- (modified) llvm/lib/Target/NVPTX/NVPTXInstrInfo.td (+19)
- (modified) llvm/lib/Target/NVPTX/NVPTXIntrinsics.td (+37-115)
``````````diff
diff --git a/llvm/lib/Target/NVPTX/NVPTXISelDAGToDAG.h b/llvm/lib/Target/NVPTX/NVPTXISelDAGToDAG.h
index fcb5700dcb6d4..8618a58ec65b0 100644
--- a/llvm/lib/Target/NVPTX/NVPTXISelDAGToDAG.h
+++ b/llvm/lib/Target/NVPTX/NVPTXISelDAGToDAG.h
@@ -102,6 +102,16 @@ class LLVM_LIBRARY_VISIBILITY NVPTXDAGToDAGISel : public SelectionDAGISel {
NVPTX::Scope getAtomicScope(const MemSDNode *N) const;
bool SelectADDR(SDValue Addr, SDValue &Base, SDValue &Offset);
+ // Match an address only when its parent memory operation is in address
+ // space AS.
+ template <unsigned AS>
+ bool SelectADDRInAS(SDNode *Parent, SDValue Addr, SDValue &Base,
+ SDValue &Offset) {
+ const auto *MemN = dyn_cast_or_null<MemSDNode>(Parent);
+ if (!MemN || MemN->getAddressSpace() != AS)
+ return false;
+ return SelectADDR(Addr, Base, Offset);
+ }
SDValue getPTXCmpMode(const CondCodeSDNode &CondCode);
SDValue selectPossiblyImm(SDValue V);
diff --git a/llvm/lib/Target/NVPTX/NVPTXInstrInfo.td b/llvm/lib/Target/NVPTX/NVPTXInstrInfo.td
index dcf89f9a73956..3c64e018d5ebe 100644
--- a/llvm/lib/Target/NVPTX/NVPTXInstrInfo.td
+++ b/llvm/lib/Target/NVPTX/NVPTXInstrInfo.td
@@ -1694,6 +1694,25 @@ def SETP_bf16x2rr :
def addr : ComplexPattern<pAny, 2, "SelectADDR">;
+// Match an address only when the parent MemSDNode has the requested address
+// space.
+class ADDR_AS<string ASName>
+ : ComplexPattern<pAny, 2,
+ "SelectADDRInAS<NVPTX::AddressSpace::" # ASName # ">"> {
+ let WantsParent = true;
+}
+
+// Using these address space patterns requires the corresponding intrinsics to
+// be registered in getTgtMemIntrinsic so that SelectionDAG preserves the
+// address space (AS) information.
+def addr_generic : ADDR_AS<"Generic">;
+def addr_global : ADDR_AS<"Global">;
+def addr_shared : ADDR_AS<"Shared">;
+def addr_const : ADDR_AS<"Const">;
+def addr_local : ADDR_AS<"Local">;
+def addr_shared_cluster : ADDR_AS<"SharedCluster">;
+def addr_param : ADDR_AS<"EntryParam">;
+
def ADDR_base : Operand<pAny>;
def ADDR : Operand<pAny> {
let PrintMethod = "printMemOperand";
diff --git a/llvm/lib/Target/NVPTX/NVPTXIntrinsics.td b/llvm/lib/Target/NVPTX/NVPTXIntrinsics.td
index 5af5ae2fe9e52..62df97d68e5d5 100644
--- a/llvm/lib/Target/NVPTX/NVPTXIntrinsics.td
+++ b/llvm/lib/Target/NVPTX/NVPTXIntrinsics.td
@@ -6,28 +6,6 @@
//
//===----------------------------------------------------------------------===//
-def AS_match {
- code generic = [{
- return cast<MemSDNode>(N)->getAddressSpace() == llvm::ADDRESS_SPACE_GENERIC;
- }];
- code shared = [{
- return cast<MemSDNode>(N)->getAddressSpace() == llvm::ADDRESS_SPACE_SHARED;
- }];
- code shared_cluster = [{
- return cast<MemSDNode>(N)->getAddressSpace() == llvm::ADDRESS_SPACE_SHARED_CLUSTER;
- }];
- code global = [{
- return cast<MemSDNode>(N)->getAddressSpace() == llvm::ADDRESS_SPACE_GLOBAL;
- }];
- code const = [{
- return cast<MemSDNode>(N)->getAddressSpace() == llvm::ADDRESS_SPACE_CONST;
- }];
- code param = [{
- return cast<MemSDNode>(N)->getAddressSpace() == llvm::ADDRESS_SPACE_ENTRY_PARAM;
- }];
-}
-
-
//===----------------------------------------------------------------------===//
// NVPTX Scope Constants
// These map to the Scope enum in NVPTX.h
@@ -953,26 +931,17 @@ defm TMA_TENSOR_PF_TILE_GATHER4_2D : TMA_TENSOR_PREFETCH_INTR<5, "tile_gather4",
//Prefetchu and Prefetch
-defvar frag_pat = (int_nvvm_prefetch_tensormap node:$addr);
-
-multiclass PREFETCH_TENSORMAP_PATFRAG<string suffix, code predicate> {
- def !tolower(suffix) : PatFrag<!setdagop(frag_pat, ops), frag_pat, predicate>;
-}
-
-defm prefetch_tensormap_ : PREFETCH_TENSORMAP_PATFRAG<"CONST", AS_match.const>;
-defm prefetch_tensormap_ : PREFETCH_TENSORMAP_PATFRAG<"GENERIC", AS_match.generic>;
-defm prefetch_tensormap_ : PREFETCH_TENSORMAP_PATFRAG<"PARAM", AS_match.param>;
-
-multiclass PREFETCH_TENSORMAP_INST<string addrspace_name, PatFrag pattern_frag> {
+multiclass PREFETCH_TENSORMAP_INST<string addrspace_name,
+ ComplexPattern addr_pattern> {
def "" : BasicNVPTXInst<(outs), (ins ADDR:$addr),
"prefetch" # addrspace_name # ".tensormap",
- [(pattern_frag addr:$addr)]>,
+ [(int_nvvm_prefetch_tensormap addr_pattern:$addr)]>,
Requires<[hasPTX<80>, hasSM<90>]>;
}
-defm PREFETCH_CONST_TENSORMAP : PREFETCH_TENSORMAP_INST<".const", prefetch_tensormap_const>;
-defm PREFETCH_GENERIC_TENSORMAP : PREFETCH_TENSORMAP_INST<"", prefetch_tensormap_generic>;
-defm PREFETCH_PARAM_TENSORMAP : PREFETCH_TENSORMAP_INST<".param", prefetch_tensormap_param>;
+defm PREFETCH_CONST_TENSORMAP : PREFETCH_TENSORMAP_INST<".const", addr_const>;
+defm PREFETCH_GENERIC_TENSORMAP : PREFETCH_TENSORMAP_INST<"", addr_generic>;
+defm PREFETCH_PARAM_TENSORMAP : PREFETCH_TENSORMAP_INST<".param", addr_param>;
class PREFETCH_INTRS<string InstName, Intrinsic Intr> :
BasicNVPTXInst<(outs), (ins ADDR:$addr),
@@ -5041,30 +5010,24 @@ class WMMA_REGINFO<WMMA_REGS r, string op, string metadata = "",
}
// Convert dag of arguments into a dag to match given intrinsic.
-class BuildPatternI<Intrinsic Intr, dag Ins> {
+class BuildPatternI<Intrinsic Intr, dag Ins, ComplexPattern Address> {
// Build a dag pattern that matches the intrinsic call.
dag ret = !foreach(tmp, Ins,
- !subst(ADDR, addr,
+ !subst(ADDR, Address,
!subst(ins, Intr,
!subst(i32imm, timm, tmp))));
}
-// Same as above, but uses PatFrag instead of an Intrinsic.
-class BuildPatternPF<PatFrag Intr, dag Ins> {
- // Build a dag pattern that matches the intrinsic call.
- dag ret = !foreach(tmp, Ins,
- !subst(ADDR, addr,
- !subst(ins, Intr, tmp)));
-}
-
// Common WMMA-related fields used for building patterns for all MMA instructions.
-class WMMA_INSTR<string _Intr, list<dag> _Args>
+class WMMA_INSTR<string _Intr, list<dag> _Args,
+ ComplexPattern Address = addr>
: NVPTXInst<(outs), (ins), "?", []> {
Intrinsic Intr = !cast<Intrinsic>(_Intr);
// Concatenate all arguments into a single dag.
dag Args = !foldl((ins), _Args, a, b, !con(a, b));
// Pre-build the pattern to match (intrinsic arg0, arg1, ...).
- dag IntrinsicPattern = BuildPatternI<!cast<Intrinsic>(Intr), Args>.ret;
+ dag IntrinsicPattern =
+ BuildPatternI<!cast<Intrinsic>(Intr), Args, Address>.ret;
}
//
@@ -5074,22 +5037,11 @@ class WMMA_INSTR<string _Intr, list<dag> _Args>
class WMMA_LOAD<WMMA_REGINFO Frag, string Layout, string Space, bit WithStride>
: WMMA_INSTR<WMMA_NAME_LDST<"load", Frag, Layout, WithStride>.record_name,
[!con((ins ADDR:$src),
- !if(WithStride, (ins B32:$ldm), (ins)))]>,
+ !if(WithStride, (ins B32:$ldm), (ins)))],
+ !cond(!eq(Space, ".shared"): addr_shared,
+ !eq(Space, ".global"): addr_global,
+ true: addr_generic)>,
Requires<Frag.Predicates> {
- // Load/store intrinsics are overloaded on pointer's address space.
- // To match the right intrinsic, we need to build AS-constrained PatFrag.
- // Operands is a dag equivalent in shape to Args, but using (ops node:$name, .....).
- dag PFOperands = !if(WithStride, (ops node:$src, node:$ldm), (ops node:$src));
- dag PFOperandsIntr = !if(WithStride, (Intr node:$src, node:$ldm), (Intr node:$src));
- // Build PatFrag that only matches particular address space.
- PatFrag IntrFrag = PatFrag<PFOperands,
- PFOperandsIntr,
- !cond(!eq(Space, ".shared"): AS_match.shared,
- !eq(Space, ".global"): AS_match.global,
- true: AS_match.generic)>;
- // Build AS-constrained pattern.
- let IntrinsicPattern = BuildPatternPF<IntrFrag, Args>.ret;
-
let OutOperandList = Frag.Outs;
let InOperandList = !con(Args, (ins MmaCode:$ptx));
let AsmString = "wmma.load."
@@ -5114,24 +5066,11 @@ class WMMA_STORE_D<WMMA_REGINFO Frag, string Layout, string Space,
: WMMA_INSTR<WMMA_NAME_LDST<"store", Frag, Layout, WithStride>.record_name,
[!con((ins ADDR:$dst),
Frag.Ins,
- !if(WithStride, (ins B32:$ldm), (ins)))]>,
+ !if(WithStride, (ins B32:$ldm), (ins)))],
+ !cond(!eq(Space, ".shared"): addr_shared,
+ !eq(Space, ".global"): addr_global,
+ true: addr_generic)>,
Requires<Frag.Predicates> {
-
- // Load/store intrinsics are overloaded on pointer's address space.
- // To match the right intrinsic, we need to build AS-constrained PatFrag.
- // Operands is a dag equivalent in shape to Args, but using (ops node:$name, .....).
- dag PFOperands = !con((ops node:$dst),
- !dag(ops, !listsplat(node, !size(Frag.regs)), Frag.reg_names),
- !if(WithStride, (ops node:$ldm), (ops)));
- // Build PatFrag that only matches particular address space.
- PatFrag IntrFrag = PatFrag<PFOperands,
- !foreach(tmp, PFOperands, !subst(ops, Intr, tmp)),
- !cond(!eq(Space, ".shared"): AS_match.shared,
- !eq(Space, ".global"): AS_match.global,
- true: AS_match.generic)>;
- // Build AS-constrained pattern.
- let IntrinsicPattern = BuildPatternPF<IntrFrag, Args>.ret;
-
let InOperandList = !con(Args, (ins MmaCode:$ptx));
let OutOperandList = (outs);
let AsmString = "wmma.store.d.sync"
@@ -5478,15 +5417,9 @@ defset list<WMMA_INSTR> MMA_SP_BLOCK_SCALEs = {
// ldmatrix.sync.aligned.m8n8[|.trans][|.shared].b16
//
class LDMATRIX<WMMA_REGINFO Frag, bit Transposed, string Space>
- : WMMA_INSTR<LDMATRIX_NAME<Frag, Transposed>.record_name, [(ins ADDR:$src)]>,
+ : WMMA_INSTR<LDMATRIX_NAME<Frag, Transposed>.record_name, [(ins ADDR:$src)],
+ !if(!eq(Space, ".shared"), addr_shared, addr_generic)>,
Requires<Frag.Predicates> {
- // Build PatFrag that only matches particular address space.
- PatFrag IntrFrag = PatFrag<(ops node:$src), (Intr node:$src),
- !cond(!eq(Space, ".shared"): AS_match.shared,
- true: AS_match.generic)>;
- // Build AS-constrained pattern.
- let IntrinsicPattern = BuildPatternPF<IntrFrag, Args>.ret;
-
let OutOperandList = Frag.Outs;
let InOperandList = !con(Args, (ins MmaCode:$ptx));
let AsmString = "ldmatrix.sync.aligned."
@@ -5513,17 +5446,10 @@ defset list<WMMA_INSTR> LDMATRIXs = {
// stmatrix.sync.aligned.m8n8[|.trans][|.shared].b16
//
class STMATRIX<WMMA_REGINFO Frag, bit Transposed, string Space>
- : WMMA_INSTR<STMATRIX_NAME<Frag, Transposed>.record_name, [!con((ins ADDR:$dst), Frag.Ins)]>,
+ : WMMA_INSTR<STMATRIX_NAME<Frag, Transposed>.record_name,
+ [!con((ins ADDR:$dst), Frag.Ins)],
+ !if(!eq(Space, ".shared"), addr_shared, addr_generic)>,
Requires<Frag.Predicates> {
- // Build PatFrag that only matches particular address space.
- dag PFOperands = !con((ops node:$dst),
- !dag(ops, !listsplat(node, !size(Frag.regs)), Frag.reg_names));
- PatFrag IntrFrag = PatFrag<PFOperands,
- !foreach(tmp, PFOperands, !subst(ops, Intr, tmp)),
- !cond(!eq(Space, ".shared"): AS_match.shared,
- true: AS_match.generic)>;
- // Build AS-constrained pattern.
- let IntrinsicPattern = BuildPatternPF<IntrFrag, Args>.ret;
let OutOperandList = (outs);
let InOperandList = !con(Args, (ins MmaCode:$ptx));
let AsmString = "stmatrix.sync.aligned."
@@ -6505,61 +6431,57 @@ foreach sparse = [0, 1] in {
class TensormapReplaceInst_2<string state_space, string field_name,
string regclass_name, NVPTXRegClass val_RC, ValueType ValTy, Intrinsic Intrin,
- code predicate> :
+ ComplexPattern addr_pattern> :
BasicNVPTXInst<(outs),
(ins ADDR:$addr, val_RC:$val),
"tensormap.replace.tile." # field_name # "." # state_space # ".b1024." # regclass_name,
- [(PatFrag<(ops node:$addr, node:$val),
- (Intrin node:$addr, node:$val), predicate>
- addr:$addr, ValTy:$val)]>;
+ [(Intrin addr_pattern:$addr, ValTy:$val)]>;
class TensormapReplaceInst_3<string state_space, string field_name,
string regclass_name, NVPTXRegClass val_RC, ValueType ValTy, Intrinsic Intrin,
- code predicate> :
+ ComplexPattern addr_pattern> :
BasicNVPTXInst<(outs),
(ins ADDR:$addr, B32:$ord, val_RC:$val),
"tensormap.replace.tile." # field_name # "." # state_space # ".b1024." # regclass_name,
- [(PatFrag<(ops node:$addr, node:$ord, node:$val),
- (Intrin node:$addr, node:$ord, node:$val), predicate>
- addr:$addr, i32:$ord, ValTy:$val)]>;
+ [(Intrin addr_pattern:$addr, i32:$ord, ValTy:$val)]>;
foreach ss = ["GLOBAL", "SHARED_CTA"] in {
- defvar pred = !if(!eq(ss, "GLOBAL"), AS_match.global, AS_match.shared);
+ defvar addr_pattern = !if(!eq(ss, "GLOBAL"), addr_global, addr_shared);
defvar ss_ptx = !tolower(!subst("_", "::", ss));
let Predicates = [callSubtarget<"hasTensormapReplaceSupport">] in {
def TENSORMAP_REPLACE_TILE_GLOBAL_ADDRESS_ # ss :
TensormapReplaceInst_2<ss_ptx, "global_address", "b64", B64, i64,
- int_nvvm_tensormap_replace_global_address, pred>;
+ int_nvvm_tensormap_replace_global_address, addr_pattern>;
foreach field_name = ["INTERLEAVE_LAYOUT", "FILL_MODE", "RANK"] in {
defvar intrin = !cast<Intrinsic>("int_nvvm_tensormap_replace_" # !tolower(field_name));
def TENSORMAP_REPLACE_TILE_ # field_name # _ # ss :
TensormapReplaceInst_2<ss_ptx, !tolower(field_name), "b32", B32, i32,
- intrin, pred>;
+ intrin, addr_pattern>;
} // field_name
def TENSORMAP_REPLACE_TILE_GLOBAL_STRIDE_ # ss :
TensormapReplaceInst_3<ss_ptx, "global_stride", "b64", B64, i64,
- int_nvvm_tensormap_replace_global_stride, pred>;
+ int_nvvm_tensormap_replace_global_stride, addr_pattern>;
foreach field_name = ["BOX_DIM", "GLOBAL_DIM", "ELEMENT_STRIDE"] in {
defvar intrin = !cast<Intrinsic>("int_nvvm_tensormap_replace_" # !tolower(field_name));
def TENSORMAP_REPLACE_TILE_ # field_name # _ # ss :
TensormapReplaceInst_3<ss_ptx, !tolower(field_name), "b32", B32, i32,
- intrin, pred>;
+ intrin, addr_pattern>;
} // field_name
} // hasTensormapReplaceSupport
def TENSORMAP_REPLACE_TILE_ELEMTYPE_ # ss :
TensormapReplaceInst_2<ss_ptx, "elemtype", "b32", B32, i32,
- int_nvvm_tensormap_replace_elemtype, pred>;
+ int_nvvm_tensormap_replace_elemtype, addr_pattern>;
def TENSORMAP_REPLACE_SWIZZLE_ATOMICITY_ # ss :
TensormapReplaceInst_2<ss_ptx, "swizzle_atomicity", "b32", B32, i32,
- int_nvvm_tensormap_replace_swizzle_atomicity, pred>,
+ int_nvvm_tensormap_replace_swizzle_atomicity, addr_pattern>,
Requires<[callSubtarget<"hasTensormapReplaceSwizzleAtomicitySupport">]>;
def TENSORMAP_REPLACE_SWIZZLE_MODE_ # ss :
TensormapReplaceInst_2<ss_ptx, "swizzle_mode", "b32", B32, i32,
- int_nvvm_tensormap_replace_swizzle_mode, pred>;
+ int_nvvm_tensormap_replace_swizzle_mode, addr_pattern>;
} // state_space
``````````
</details>
https://github.com/llvm/llvm-project/pull/210309
More information about the llvm-commits
mailing list