[llvm] [NVVM][NVPTX] Add tensor map override support in S2G and Reduce intrinsics (PR #215503)

Rajat Bajpai via llvm-commits llvm-commits at lists.llvm.org
Wed Aug 12 02:15:06 PDT 2026


================
@@ -1268,6 +1268,93 @@ class NVVMPureIntrinsic<list<LLVMType> ret_types,
   DefaultAttrsIntrinsic<ret_types, param_types,
                         intr_properties # [IntrNoMem, IntrSpeculatable], name> {}
 
+// Class to handle TMA Copy Override Support
+class TMAOverrideHandler<int dim, string mode> {
+  // TMA Override Support:
+  // +----------------+------------------+-------------------------+
+  // | Mode           | Dimensions       | Supported Overrides     |
+  // +----------------+------------------+-------------------------+
+  // | Tile           | 1D               | override_addr           |
+  // |                |                  | override_addr_dim       |
+  // +----------------+------------------+-------------------------+
+  // | Tile           | 2D and higher    | override_addr           |
+  // |                |                  | override_addr_dim_stride|
+  // +----------------+------------------+-------------------------+
+  // | Im2col         | 3D and higher    | override_addr           |
+  // +----------------+------------------+-------------------------+
+  list<string> common_overrides = ["override_addr"];
+  list<string> tile_overrides_1d = !listconcat(common_overrides, ["override_addr_dim"]);
+  list<string> tile_overrides_nd = !listconcat(common_overrides, ["override_addr_dim_stride"]);
+
+  // Return the applicable overrides based on mode and dimension
+  list<string> applicable_overrides = !if(!eq(mode, "tile"),
+                                !if(!eq(dim, 1), tile_overrides_1d, tile_overrides_nd),
+                                common_overrides);
+}
+
+class CP_ASYNC_BULK_TENSOR_OVERRIDE_BASE<int dim, string mode, string override> {
+  list<LLVMType> OverrideTy = [llvm_global_ptr_ty];
+  list<LLVMType> TensorSizeTy = !listsplat(llvm_i16_ty, dim);
+  list<LLVMType> LowerStrideTy = !listsplat(llvm_i32_ty, !add(dim, -1));
+  list<LLVMType> UpperStrideTy = [llvm_i16_ty];
+  list<LLVMType> OverrideDimTy = !listconcat(OverrideTy, TensorSizeTy);
+  list<LLVMType> OverrideDimStrideTy = !listconcat(OverrideTy, TensorSizeTy, LowerStrideTy, UpperStrideTy);
+  list<LLVMType> OverrideAddrTy = !cond(!eq(override, "override_addr"): OverrideTy,
+                                        !eq(override, "override_addr_dim"): OverrideDimTy,
+                                        !eq(override, "override_addr_dim_stride"): OverrideDimStrideTy);
+  list<LLVMType> TensorDimsTy = !listsplat(llvm_i32_ty, dim);
+
+  // For im2col_w/w128 mode, NumIm2ColOffsets is always 2 irrespective of the
+  // dim (wHalo and wOffset)
+  int NumIm2ColOffsets =
+      !cond(!eq(mode, "im2col_w") : 2,
+            !eq(mode, "im2col_w_128") : 2,
+            !eq(mode, "im2col") : !add(dim, -2),
+            true : 0);
+  list<LLVMType> Im2ColOffsetsTy = !listsplat(llvm_i16_ty, NumIm2ColOffsets);
+
+  string Prefix = "int_nvvm_cp_async_bulk_tensor_";
+  string Suffix = "_" # mode # "_" # override # "_" # dim # "d";
+}
+
+class CP_ASYNC_BULK_TENSOR_S2G_OVERRIDE_INTR<int dim, string mode, string override>
+      : CP_ASYNC_BULK_TENSOR_OVERRIDE_BASE<dim, mode, override> {
+  string Name = Prefix # "s2g" # Suffix;
+
+  list<LLVMType> ArgsTy = !listconcat(
+                          [llvm_shared_ptr_ty,  // src_smem_ptr
+                           llvm_ptr_ty],        // tensormap_ptr
+                           OverrideAddrTy,      // override_addr
----------------
rajatbajpai wrote:

sure, added.

https://github.com/llvm/llvm-project/pull/215503


More information about the llvm-commits mailing list