v4.4 release update v2. (#2999)

This commit is contained in:
Junkai-Wu
2026-02-03 20:48:31 -05:00
committed by GitHub
parent 1cfbb53a23
commit 6b3e607b85
91 changed files with 13242 additions and 1488 deletions
+32
View File
@@ -6,6 +6,18 @@
### CuTe DSL
* New features
- CuTe DSL now supports CUDA toolkit 13.1!
+ Set up with cutlass/python/CuTeDSL/setup.sh --cu13
+ Refer to https://docs.nvidia.com/cutlass/latest/media/docs/pythonDSL/quick_start.html for more details
- GB300 is now supported in CuTe DSL with CTK 13.1
+ Refer to [SM103 batched 3xFP4 blockscaled GEMM kernel](https://github.com/NVIDIA/cutlass/tree/main/examples/python/CuTeDSL/blackwell/sm103_dense_blockscaled_gemm_persistent.py) for example kernel
- cute.experimental: introduce a higher-level, composable layer on top of existing CuTe DSL APIs (not a separate abstraction), which can be mixed with existing Cute DSL building blocks.
+ Fragment-free programming model: copy/dot APIs take memrefs directly instead of descriptors/fragments.
+ Automatic TMA descriptor generation and update insertion.
+ Automatic vectorization and predication for SIMT copies.
+ New pipeline abstraction with convenience wrappers
+ New Partition ops to simplify partitioning logic.
+ Device-side TMA descriptor allocation, initialization, and management
- Ahead of Time (AoT) compilation is now available!
+ Refer to files under https://github.com/NVIDIA/cutlass/tree/main/examples/python/CuTeDSL/cute/export for example usage
- JAX support - you can now use CuTeDSL along with JAX
@@ -15,7 +27,11 @@
+ cutlass.CUDA_VERSION for a version class to tell the CUDA version used for DSL
- Added CopyDsmemStoreOp to store data to distributed shared memory with explicit synchronization.
* More examples of authorizing peak-performance kernels
- [SM103 batched 3xFP4 blockscaled GEMM kernel](https://github.com/NVIDIA/cutlass/tree/main/examples/python/CuTeDSL/blackwell/sm103_dense_blockscaled_gemm_persistent.py)
* Bug fixing and improvements
- Fixed an issue that both branches of if are executed
- Fixed `cute.printf` with f-string
- Fixed an issue that cutlass.cuda.initialize_cuda_context() silently kills python
@@ -26,6 +42,21 @@
- LdMatrix16x16x8bOp copy traits updated to be faithful to PTX without permutations. Permuted variant is renamed to LdMatrix16x8x8bOp.
- group_bulk_copy_modes in async bulk copy example is now deprecated, use group_modes directly instead.
- cute.arch.calc_packed_f32x2_op default enable ftz to default disable ftz
- In CuTe DSL with CTK 13.1, following APIs in cutlass.cute.arch now require string literal instead of enum as argument:
+ fence_proxy
+ fence_view_async_tmem_op
+ calc_packed_f32x2_op
+ warp_redux_sync
+ atomic_add
+ atomic_and
+ atomic_or
+ atomic_xor
+ atomic_max
+ atomic_min
+ atomic_exch
+ atomic_cas
+ store
+ load
### CUTLASS C++
* Add Hopper e2m1 to fp32 optimized conversion and e2m1 * TF32 tensor core GEMM.
@@ -54,6 +85,7 @@
- Fix a TMA descriptor bug where the CUDA driver is not properly setting the OOB address gen mode correctly.
- Fix memory fence for clc scheduler in Blackwell SM120 pingpong kernel.
- Fix missing SMEM alignment in Blackwell SM120 scale factors.
- Fix a PDL issue for grouped gemm.
* Fix some profiler issues:
- Refactor L1 functional test generation logic to reduce the L1 test cases to avoid timeout.
- Fix a core dump issue for nvfp4 grouped GEMM kernel.
+34 -2
View File
@@ -45,8 +45,20 @@ To get started quickly - please refer :
# What's New in CUTLASS 4.4
### CuTe DSL
## CuTe DSL
* New features
- CuTe DSL now supports CUDA toolkit 13.1!
+ Set up with cutlass/python/CuTeDSL/setup.sh --cu13
+ Refer to https://docs.nvidia.com/cutlass/latest/media/docs/pythonDSL/quick_start.html for more details
- GB300 is now supported in CuTe DSL with CTK 13.1
+ Refer to [SM103 batched 3xFP4 blockscaled GEMM kernel](https://github.com/NVIDIA/cutlass/tree/main/examples/python/CuTeDSL/blackwell/sm103_dense_blockscaled_gemm_persistent.py) for example kernel
- cute.experimental: introduce a higher-level, composable layer on top of existing CuTe DSL APIs (not a separate abstraction), which can be mixed with existing Cute DSL building blocks.
+ Fragment-free programming model: copy/dot APIs take memrefs directly instead of descriptors/fragments.
+ Automatic TMA descriptor generation and update insertion.
+ Automatic vectorization and predication for SIMT copies.
+ New pipeline abstraction with convenience wrappers
+ New Partition ops to simplify partitioning logic.
+ Device-side TMA descriptor allocation, initialization, and management
- Ahead of Time (AoT) compilation is now available!
+ Refer to files under https://github.com/NVIDIA/cutlass/tree/main/examples/python/CuTeDSL/cute/export for example usage
- JAX support - you can now use CuTeDSL along with JAX
@@ -56,7 +68,11 @@ To get started quickly - please refer :
+ cutlass.CUDA_VERSION for a version class to tell the CUDA version used for DSL
- Added CopyDsmemStoreOp to store data to distributed shared memory with explicit synchronization.
* More examples of authorizing peak-performance kernels
- [SM103 batched 3xFP4 blockscaled GEMM kernel](https://github.com/NVIDIA/cutlass/tree/main/examples/python/CuTeDSL/blackwell/sm103_dense_blockscaled_gemm_persistent.py)
* Bug fixing and improvements
- Fixed an issue that both branches of if are executed
- Fixed `cute.printf` with f-string
- Fixed an issue that cutlass.cuda.initialize_cuda_context() silently kills python
@@ -67,8 +83,23 @@ To get started quickly - please refer :
- LdMatrix16x16x8bOp copy traits updated to be faithful to PTX without permutations. Permuted variant is renamed to LdMatrix16x8x8bOp.
- group_bulk_copy_modes in async bulk copy example is now deprecated, use group_modes directly instead.
- cute.arch.calc_packed_f32x2_op default enable ftz to default disable ftz
- In CuTe DSL with CTK 13.1, following APIs in cutlass.cute.arch now require string literal instead of enum as argument:
+ fence_proxy
+ fence_view_async_tmem_op
+ calc_packed_f32x2_op
+ warp_redux_sync
+ atomic_add
+ atomic_and
+ atomic_or
+ atomic_xor
+ atomic_max
+ atomic_min
+ atomic_exch
+ atomic_cas
+ store
+ load
### CUTLASS C++
## CUTLASS C++
* Add Hopper e2m1 to fp32 optimized conversion and e2m1 * TF32 tensor core GEMM.
- Set MmaType to tfloat32_t for FP32 mode.
- TF32 provides FP32 inputs with reduced precision (19-bit vs 32-bit)
@@ -95,6 +126,7 @@ To get started quickly - please refer :
- Fix a TMA descriptor bug where the CUDA driver is not properly setting the OOB address gen mode correctly.
- Fix memory fence for clc scheduler in Blackwell SM120 pingpong kernel.
- Fix missing SMEM alignment in Blackwell SM120 scale factors.
- Fix a PDL issue for grouped gemm.
* Fix some profiler issues:
- Refactor L1 functional test generation logic to reduce the L1 test cases to avoid timeout.
- Fix a core dump issue for nvfp4 grouped GEMM kernel.
@@ -34,8 +34,8 @@ import cuda.bindings.driver as cuda
import cutlass
import cutlass.cute as cute
import cutlass.cute.testing as testing
from cutlass.cute.runtime import from_dlpack
import cutlass.utils as utils
from cutlass.utils import is_fp8_dtype, create_cute_tensor_for_fp8
import cutlass.pipeline as pipeline
from cutlass.pipeline import pipeline_init_arrive, pipeline_init_wait
from cutlass.cute.nvgpu import cpasync, tcgen05
@@ -1046,45 +1046,70 @@ class PersistentDenseGemmKernel:
# (MMA, MMA_M, MMA_N, STAGE)
tCtAcc_base = cute.make_tensor(tmem_ptr, tCtAcc_fake.layout)
#
# Persistent tile scheduling loop for epilogue
#
acc_consumer_state = pipeline.make_pipeline_state(
pipeline.PipelineUserType.Consumer, self.num_acc_stage
)
if cutlass.const_expr(self.use_tma_store):
assert tma_atom_c is not None and sC is not None
utils.gemm.sm100.epilogue_tma_store(
self,
tidx,
warp_idx,
acc_pipeline,
tiled_mma,
tma_atom_c,
tCtAcc_base,
sC,
tCgC,
epi_tile,
tile_sched,
epilogue_op,
clc_pipeline,
clc_consumer_state,
c_producer_group = pipeline.CooperativeGroup(
pipeline.Agent.Thread,
32 * len(self.epilogue_warp_id),
)
else:
utils.gemm.sm100.epilogue(
self,
tidx,
acc_pipeline,
tiled_mma,
tCtAcc_base,
tCgC,
epi_tile,
tile_sched,
epilogue_op,
tmem_dealloc_barrier,
None,
None,
clc_pipeline,
clc_consumer_state,
c_pipeline = pipeline.PipelineTmaStore.create(
num_stages=self.num_c_stage, producer_group=c_producer_group
)
while work_tile.is_valid_tile:
# Get tile coord from tile scheduler
cur_tile_coord = work_tile.tile_idx
mma_tile_coord_mnl = (
cur_tile_coord[0] // cute.size(tiled_mma.thr_id.shape),
cur_tile_coord[1],
cur_tile_coord[2],
)
num_tiles_executed = tile_sched.num_tiles_executed
if cutlass.const_expr(self.use_tma_store):
acc_consumer_state = utils.gemm.sm100.epilogue_tma_store(
self,
tidx,
warp_idx,
tma_atom_c,
tCtAcc_base,
sC,
tCgC,
epi_tile,
num_tiles_executed,
epilogue_op,
mma_tile_coord_mnl,
acc_consumer_state,
acc_pipeline,
c_pipeline,
)
else:
acc_consumer_state = utils.gemm.sm100.epilogue(
self,
tidx,
tCtAcc_base,
tCgC,
epi_tile,
epilogue_op,
mma_tile_coord_mnl,
acc_consumer_state,
acc_pipeline,
)
#
# Advance to next tile
#
clc_pipeline.consumer_wait(clc_consumer_state)
work_tile = tile_sched.get_current_work()
clc_pipeline.consumer_release(clc_consumer_state)
clc_consumer_state.advance()
if cutlass.const_expr(self.use_tma_store):
# Wait for C store complete
c_pipeline.producer_tail()
else:
# Synchronize before TMEM dealloc (done by the caller)
tmem_dealloc_barrier.arrive_and_wait()
#
# Dealloc the tensor memory buffer
#
@@ -1150,8 +1175,11 @@ class PersistentDenseGemmKernel:
return num_tmem_alloc_cols
def check_supported_dtypes(
self, ab_dtype: Type[cutlass.Numeric], c_dtype: Type[cutlass.Numeric]
) -> bool:
self,
a_dtype: Type[cutlass.Numeric],
b_dtype: Type[cutlass.Numeric],
c_dtype: Type[cutlass.Numeric],
):
"""
Check if the dtypes are valid
@@ -1173,8 +1201,10 @@ class PersistentDenseGemmKernel:
cutlass.Float8E4M3FN,
cutlass.Float8E5M2,
}
if ab_dtype not in valid_ab_dtypes:
raise testing.CantImplementError(f"Unsupported AB dtype: {ab_dtype}")
if a_dtype not in valid_ab_dtypes or b_dtype not in valid_ab_dtypes:
raise testing.CantImplementError(
f"Unsupported AB dtype: {a_dtype} and {b_dtype}"
)
if self.acc_dtype not in {cutlass.Float32, cutlass.Float16, cutlass.Int32}:
raise testing.CantImplementError(
@@ -1198,8 +1228,13 @@ class PersistentDenseGemmKernel:
cutlass.Int32: {cutlass.Uint8, cutlass.Int8},
}
# Check compatibility between accumulator type and AB type
if ab_dtype not in acc_ab_compatibility[self.acc_dtype]:
return False
if (
a_dtype not in acc_ab_compatibility[self.acc_dtype]
or b_dtype not in acc_ab_compatibility[self.acc_dtype]
):
raise testing.CantImplementError(
f"Unsupported AB dtype: {a_dtype} and {b_dtype} for accumulator dtype: {self.acc_dtype}"
)
# Define compatibility mapping between accumulator type and C type
acc_c_compatibility = {
@@ -1228,11 +1263,11 @@ class PersistentDenseGemmKernel:
}
# Check compatibility between accumulator type and C type
if c_dtype not in acc_c_compatibility[self.acc_dtype]:
return False
raise testing.CantImplementError(
f"Unsupported C dtype: {c_dtype} for accumulator dtype: {self.acc_dtype}"
)
return True
def check_mma_tiler_and_cluster_shape(self) -> bool:
def check_mma_tiler_and_cluster_shape(self):
"""Check if the mma tiler and cluster shape are valid.
:raises testing.CantImplementError: If the mma tiler and cluster shape are invalid
@@ -1273,12 +1308,13 @@ class PersistentDenseGemmKernel:
n: int,
k: int,
l: int,
ab_dtype: Type[cutlass.Numeric],
a_dtype: Type[cutlass.Numeric],
b_dtype: Type[cutlass.Numeric],
c_dtype: Type[cutlass.Numeric],
a_major: str,
b_major: str,
c_major: str,
) -> bool:
):
"""
Check if the tensor alignment is valid
@@ -1290,8 +1326,10 @@ class PersistentDenseGemmKernel:
:type k: int
:param l: The number of columns in the C tensor
:type l: int
:param ab_dtype: The data type of the A and B operands
:type ab_dtype: Type[cutlass.Numeric]
:param a_dtype: The data type of the A operand
:type a_dtype: Type[cutlass.Numeric]
:param b_dtype: The data type of the B operand
:type b_dtype: Type[cutlass.Numeric]
:param c_dtype: The data type of the output tensor
:type c_dtype: Type[cutlass.Numeric]
:param a_major: The major axis of the A tensor
@@ -1301,8 +1339,7 @@ class PersistentDenseGemmKernel:
:param c_major: The major axis of the C tensor
:type c_major: str
:return: True if the problem shape is valid, False otherwise
:rtype: bool
:raises testing.CantImplementError: If the tensor alignment is invalid
"""
# TODO: move to utils
@@ -1313,15 +1350,15 @@ class PersistentDenseGemmKernel:
return num_major_elements % num_contiguous_elements == 0
if (
not check_contiguous_16B_alignment(ab_dtype, a_major == "m", (m, k, l))
or not check_contiguous_16B_alignment(ab_dtype, b_major == "n", (n, k, l))
not check_contiguous_16B_alignment(a_dtype, a_major == "m", (m, k, l))
or not check_contiguous_16B_alignment(b_dtype, b_major == "n", (n, k, l))
or not check_contiguous_16B_alignment(c_dtype, c_major == "m", (m, n, l))
):
raise testing.CantImplementError(
f"Invalid tensor alignment: {m}, {n}, {k}, {l}, {ab_dtype}, {c_dtype}, {a_major}, {b_major}, {c_major}"
f"Invalid tensor alignment: {m}, {n}, {k}, {l}, {a_dtype}, {b_dtype}, {c_dtype}, {a_major}, {b_major}, {c_major}"
)
def check_epilog_store_option(self, m: int, n: int) -> bool:
def check_epilog_store_option(self, m: int, n: int):
"""
Check if the epilogue store option is valid
@@ -1346,7 +1383,8 @@ class PersistentDenseGemmKernel:
def can_implement(
self,
mnkl: Tuple[int, int, int, int],
ab_dtype: Type[cutlass.Numeric],
a_dtype: Type[cutlass.Numeric],
b_dtype: Type[cutlass.Numeric],
c_dtype: Type[cutlass.Numeric],
a_major: str,
b_major: str,
@@ -1357,8 +1395,10 @@ class PersistentDenseGemmKernel:
:param mnkl: Problem size as a tuple (M, N, K, L).
:type mnkl: Tuple[int, int, int, int]
:param ab_dtype: Data type for input tensors A and B.
:type ab_dtype: Type[cutlass.Numeric]
:param a_dtype: Data type for input tensors A.
:type a_dtype: Type[cutlass.Numeric]
:param b_dtype: Data type for input tensors B.
:type b_dtype: Type[cutlass.Numeric]
:param c_dtype: Data type for output tensor C.
:type c_dtype: Type[cutlass.Numeric]
:param a_major: Major dimension of the A tensor layout ("m" or "k").
@@ -1373,14 +1413,14 @@ class PersistentDenseGemmKernel:
try:
# Skip unsupported types
self.check_supported_dtypes(ab_dtype, c_dtype)
self.check_supported_dtypes(a_dtype, b_dtype, c_dtype)
# Skip invalid mma tile shape and cluster shape
self.check_mma_tiler_and_cluster_shape()
m, n, k, l = mnkl
self.check_tensor_alignment(
m, n, k, l, ab_dtype, c_dtype, a_major, b_major, c_major
m, n, k, l, a_dtype, b_dtype, c_dtype, a_major, b_major, c_major
)
self.check_epilog_store_option(m, n)
except testing.CantImplementError:
@@ -1432,43 +1472,72 @@ def bmm(
@lru_cache(maxsize=1)
def prepare_tensors(
mnkl: Tuple[int, int, int, int],
ab_dtype: Type[cutlass.Numeric],
a_dtype: Type[cutlass.Numeric],
b_dtype: Type[cutlass.Numeric],
c_dtype: Type[cutlass.Numeric],
a_major: str,
b_major: str,
c_major: str,
init_random: bool = True,
normal_mean: float = 0.0,
normal_std: float = 1.0,
):
"""Prepare tensors for GEMM.
Returns:
Tuple of (a_f32, b_f32, c_f32, a_storage, b_storage, c_storage):
- *_f32: Float32 tensors with the logical data (for reference and fp8 conversion)
- *_storage: Storage tensors for DLPack (uint8 for fp8, otherwise the target dtype)
"""
import torch
from cutlass.torch import dtype as torch_dtype
m, n, k, l = mnkl
if a_major == "k":
a = torch.empty((l, m, k), dtype=torch.float32, device="cuda")
a_f32 = torch.empty((l, m, k), dtype=torch.float32, device="cuda")
elif a_major == "m":
a = torch.empty((l, k, m), dtype=torch.float32, device="cuda").permute(0, 2, 1)
a_f32 = torch.empty((l, k, m), dtype=torch.float32, device="cuda").permute(
0, 2, 1
)
if b_major == "n":
b = torch.empty((l, k, n), dtype=torch.float32, device="cuda")
b_f32 = torch.empty((l, k, n), dtype=torch.float32, device="cuda")
elif b_major == "k":
b = torch.empty((l, n, k), dtype=torch.float32, device="cuda").permute(0, 2, 1)
b_f32 = torch.empty((l, n, k), dtype=torch.float32, device="cuda").permute(
0, 2, 1
)
if c_major == "n":
c = torch.empty((l, m, n), dtype=torch.float32, device="cuda")
c_f32 = torch.empty((l, m, n), dtype=torch.float32, device="cuda")
elif c_major == "m":
c = torch.empty((l, n, m), dtype=torch.float32, device="cuda").permute(0, 2, 1)
c_f32 = torch.empty((l, n, m), dtype=torch.float32, device="cuda").permute(
0, 2, 1
)
if init_random:
a.random_(-2, 3)
b.random_(-2, 3)
c.random_(-2, 3)
# Uniform random initialization in range [-2, 3)
a_f32.random_(-2, 3)
b_f32.random_(-2, 3)
c_f32.random_(-2, 3)
else:
# Normal (Gaussian) initialization with user-specified mean and std
a_f32.normal_(mean=normal_mean, std=normal_std)
b_f32.normal_(mean=normal_mean, std=normal_std)
c_f32.normal_(mean=normal_mean, std=normal_std)
return (
a.to(dtype=torch_dtype(ab_dtype)),
b.to(dtype=torch_dtype(ab_dtype)),
c.to(dtype=torch_dtype(c_dtype)),
)
# For float8 types, use uint8 as storage type to avoid dlpack limitation
# (dlpack doesn't support float8 types)
# For other types, convert to the target dtype
a_storage_dtype = torch.uint8 if is_fp8_dtype(a_dtype) else torch_dtype(a_dtype)
b_storage_dtype = torch.uint8 if is_fp8_dtype(b_dtype) else torch_dtype(b_dtype)
c_storage_dtype = torch.uint8 if is_fp8_dtype(c_dtype) else torch_dtype(c_dtype)
a_storage = a_f32.to(dtype=a_storage_dtype)
b_storage = b_f32.to(dtype=b_storage_dtype)
c_storage = c_f32.to(dtype=c_storage_dtype)
return (a_f32, b_f32, c_f32, a_storage, b_storage, c_storage)
@lru_cache(maxsize=1)
@@ -1499,7 +1568,7 @@ def compile_bmm(
)
# Check if configuration can be implemented
can_implement = gemm.can_implement(
mnkl, a.element_type, c.element_type, a_major, b_major, c_major
mnkl, a.element_type, b.element_type, c.element_type, a_major, b_major, c_major
)
if not can_implement:
raise testing.CantImplementError(
@@ -1594,21 +1663,30 @@ def run(
)
# Run and verify BMM with torch
a, b, c = prepare_tensors(mnkl, ab_dtype, c_dtype, a_major, b_major, c_major)
a_f32, b_f32, c_f32, a_storage, b_storage, c_storage = prepare_tensors(
mnkl, ab_dtype, ab_dtype, c_dtype, a_major, b_major, c_major
)
leading_dim_a = 2 if a_major == "k" else 1
leading_dim_b = 1 if b_major == "k" else 2
leading_dim_c = 2 if c_major == "n" else 1
a_ = from_dlpack(
a, assumed_align=16, force_tf32=ab_dtype == cutlass.TFloat32
).mark_layout_dynamic(leading_dim=leading_dim_a)
b_ = from_dlpack(
b, assumed_align=16, force_tf32=ab_dtype == cutlass.TFloat32
).mark_layout_dynamic(leading_dim=leading_dim_b)
c_ = from_dlpack(
c, assumed_align=16, force_tf32=c_dtype == cutlass.TFloat32
).mark_layout_dynamic(leading_dim=leading_dim_c)
# Create CuTe tensors, passing float32 source for fp8 conversion
a_ = create_cute_tensor_for_fp8(
a_storage, ab_dtype, leading_dim_a, source_f32_tensor=a_f32
)
b_ = create_cute_tensor_for_fp8(
b_storage, ab_dtype, leading_dim_b, source_f32_tensor=b_f32
)
c_ = create_cute_tensor_for_fp8(
c_storage, c_dtype, leading_dim_c, source_f32_tensor=c_f32
)
print("Compile Blackwell Persistent Dense GEMM with:")
print(f"ab_dtype: {ab_dtype}, c_dtype: {c_dtype}, acc_dtype: {acc_dtype}")
print(f"a_major: {a_major}, b_major: {b_major}, c_major: {c_major}")
print(f"mma_tiler_mn: {mma_tiler_mn}, cluster_shape_mn: {cluster_shape_mn}")
print(f"use_2cta_instrs: {use_2cta_instrs}, use_tma_store: {use_tma_store}")
compiled_fn = compile_bmm(
mnkl,
@@ -1640,13 +1718,15 @@ def run(
compiled_fn(a_, b_, c_, current_stream)
# Manually quantize to be comparable
# Use float32 source data for reference calculation
ref = (
torch.bmm(a.to(dtype=torch.float32), b.to(dtype=torch.float32))
torch.bmm(a_f32, b_f32)
.to(dtype=torch_dtype(c_dtype))
.to(dtype=torch.float32)
)
# Read back the result from CuTe tensor (c_storage was updated in-place)
torch.testing.assert_close(
c.to(dtype=torch.float32), ref, atol=tolerance, rtol=1e-03
c_storage.to(dtype=torch.float32), ref, atol=tolerance, rtol=1e-03
)
if not benchmark:
@@ -1654,32 +1734,33 @@ def run(
def generate_tensors():
init_normal = ab_dtype not in [cutlass.Int8, cutlass.Uint8]
a, b, c = prepare_tensors(
a_f32, b_f32, c_f32, a_st, b_st, c_st = prepare_tensors(
mnkl,
ab_dtype,
ab_dtype,
c_dtype,
a_major,
b_major,
c_major,
init_random=not init_normal,
)
a_ = from_dlpack(
a, assumed_align=16, force_tf32=ab_dtype == cutlass.TFloat32
).mark_layout_dynamic(leading_dim=leading_dim_a)
b_ = from_dlpack(
b, assumed_align=16, force_tf32=ab_dtype == cutlass.TFloat32
).mark_layout_dynamic(leading_dim=leading_dim_b)
c_ = from_dlpack(
c, assumed_align=16, force_tf32=c_dtype == cutlass.TFloat32
).mark_layout_dynamic(leading_dim=leading_dim_c)
a_ = create_cute_tensor_for_fp8(
a_st, ab_dtype, leading_dim_a, source_f32_tensor=a_f32
)
b_ = create_cute_tensor_for_fp8(
b_st, ab_dtype, leading_dim_b, source_f32_tensor=b_f32
)
c_ = create_cute_tensor_for_fp8(
c_st, c_dtype, leading_dim_c, source_f32_tensor=c_f32
)
return testing.JitArguments(a_, b_, c_, current_stream)
workspace_count = 1
if use_cold_l2:
one_workspace_bytes = (
a.numel() * a.element_size()
+ b.numel() * b.element_size()
+ c.numel() * c.element_size()
a_storage.numel() * a_storage.element_size()
+ b_storage.numel() * b_storage.element_size()
+ c_storage.numel() * c_storage.element_size()
)
workspace_count = testing.get_workspace_count(
one_workspace_bytes, warmup_iterations, iterations
@@ -1735,6 +1816,7 @@ def prepare_parser():
action="store_true",
help="Enable 2CTA MMA instructions feature",
)
parser.add_argument("--a_major", choices=["k", "m"], type=str, default="k")
parser.add_argument("--b_major", choices=["k", "n"], type=str, default="k")
parser.add_argument("--c_major", choices=["n", "m"], type=str, default="n")
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,155 @@
# Copyright (c) 2025 - 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: BSD-3-Clause
# Redistribution and use in source and binary forms, with or without
# modification, are permitted provided that the following conditions are met:
# 1. Redistributions of source code must retain the above copyright notice, this
# list of conditions and the following disclaimer.
# 2. Redistributions in binary form must reproduce the above copyright notice,
# this list of conditions and the following disclaimer in the documentation
# and/or other materials provided with the distribution.
# 3. Neither the name of the copyright holder nor the names of its
# contributors may be used to endorse or promote products derived from
# this software without specific prior written permission.
# THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
# AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
# IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
# DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
# FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
# DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
# SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
# CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
# OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
# OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
import argparse
import torch
import pytest
from cutlass import cute
from cutlass.cute import experimental as cute_ext
from cutlass.cute.runtime import from_dlpack
import cutlass.utils as utils
@cute.experimental.kernel
def memcpy_simt_universal_copy_kernel(
mA: cute.Tensor, mD: cute.Tensor, addend: cute.Float16
):
tile_mn = cute.core._pack_shape((128, 64))
gA = cute.zipped_divide(mA, tile_mn)
gD = cute.zipped_divide(mD, tile_mn)
cta_m, cta_n, cta_l = cute.arch.block_idx()
tid_x, _, _ = cute.arch.thread_idx()
gA_tile = gA[(None, None), (cta_m, cta_n, cta_l)]
gD_tile = gD[(None, None), (cta_m, cta_n, cta_l)]
buffer = cute_ext.allocate(
cute.Float16,
cute.AddressSpace.rmem,
cute.make_layout(((8, 1), (1, 8)), stride=((1, 8), (1, 8))),
alignment=16,
)
tCgA = cute_ext.partition(
gA_tile,
tid_x,
layout_tv=cute.make_layout(((16, 8), (8, 1)), stride=((8, 128), (1, 1024))),
tiler=cute.core._pack_tile((128, 8)),
)
tCgD = cute_ext.partition(
gD_tile,
tid_x,
layout_tv=cute.make_layout(((16, 8), (8, 1)), stride=((8, 128), (1, 1024))),
tiler=cute.core._pack_tile((128, 8)),
)
# cute_ext.copy() automatically computes predicates based on the shape of
# the tensor passed to the @cute.experimental.kernel argument
cute_ext.copy(
tCgA,
buffer,
copy_atom=cute.make_copy_atom(
cute.nvgpu.CopyUniversalOp(),
tCgD.element_type,
num_bits_per_copy=128,
),
)
# Update the RMEM tensor in place using elementwise addition.
buffer.store(buffer.load() + addend)
# cute_ext.copy() automatically computes predicates based on the shape of
# the tensor passed to the @cute.experimental.kernel argument
cute_ext.copy(
buffer,
tCgD,
copy_atom=cute.make_copy_atom(
cute.nvgpu.CopyUniversalOp(),
tCgD.element_type,
num_bits_per_copy=128,
),
)
@cute.experimental.jit
def memcpy_simt_universal_copy(
src: cute.Tensor, dst: cute.Tensor, addend: cute.Float16
):
tile_mn = cute.core._pack_shape((128, 64))
div = cute.tiled_divide(src, tile_mn)
grid = (div.shape[1], div.shape[2], div.shape[3])
memcpy_simt_universal_copy_kernel(src, dst, addend).launch(
grid=grid,
block=(128, 1, 1),
smem=cute.Int64(utils.get_smem_capacity_in_bytes("sm_80")),
)
def run_simt_universal_memcpy(M, N, L):
src = torch.randn(L, N, M).permute(2, 1, 0).to(torch.float16).cuda()
dst = torch.randn(L, N, M).permute(2, 1, 0).to(torch.float16).cuda()
mA = (
from_dlpack(src, assumed_align=16)
.mark_layout_dynamic(leading_dim=0)
.mark_compact_shape_dynamic(
mode=0, stride_order=src.dim_order(), divisibility=8
)
)
mD = (
from_dlpack(dst, assumed_align=16)
.mark_layout_dynamic(leading_dim=0)
.mark_compact_shape_dynamic(
mode=0, stride_order=dst.dim_order(), divisibility=8
)
)
addend = 5.0
memcpy_simt_universal_copy(
mA,
mD,
cute.Float16(addend),
no_cache=True,
)
torch.testing.assert_close(src.cpu() + addend, dst.cpu())
if __name__ == "__main__":
parser = argparse.ArgumentParser(
description="Example memory copy example using CuTe auto predication features."
)
parser.add_argument("--mnl", default=[136, 7, 9], nargs="+", type=int)
args = parser.parse_args()
M, N, L = tuple(args.mnl)
run_simt_universal_memcpy(M, N, L)
print("PASS")
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,522 @@
# SPDX-FileCopyrightText: Copyright (c) 2025 - 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: LicenseRef-NvidiaProprietary
#
# NVIDIA CORPORATION, its affiliates and licensors retain all intellectual
# property and proprietary rights in and to this material, related
# documentation and any modifications thereto. Any use, reproduction,
# disclosure or distribution of this material and related documentation
# without an express license agreement from NVIDIA CORPORATION or
# its affiliates is strictly prohibited.
"""
2SM Dense GEMM example using cute_ext decorators.
"""
import torch
import math
import cutlass
from cutlass import cute
from cutlass.cute import experimental as cute_ext
from cutlass.cute.runtime import from_dlpack
import cutlass.utils.blackwell_helpers as sm100_utils
import cutlass.utils as utils
from cutlass.base_dsl.typing import Numeric
from typing import Type
def create_gemm_tensors_torch(
M,
N,
K,
majors: tuple[
cute.nvgpu.tcgen05.OperandMajorMode,
cute.nvgpu.tcgen05.OperandMajorMode,
cute.nvgpu.tcgen05.OperandMajorMode,
],
dtypes: tuple[torch.dtype, torch.dtype, torch.dtype],
):
A = None
B = None
D = None
if majors[0] == cute.nvgpu.tcgen05.OperandMajorMode.MN:
A = torch.empty(K, M).random_(-4, 4).permute(1, 0).to(dtypes[0]).cuda()
elif majors[0] == cute.nvgpu.tcgen05.OperandMajorMode.K:
A = torch.empty(M, K).random_(-4, 4).permute(0, 1).to(dtypes[0]).cuda()
if majors[1] == cute.nvgpu.tcgen05.OperandMajorMode.MN:
B = torch.empty(K, N).random_(-4, 4).permute(1, 0).to(dtypes[1]).cuda()
elif majors[1] == cute.nvgpu.tcgen05.OperandMajorMode.K:
B = torch.empty(N, K).random_(-4, 4).permute(0, 1).to(dtypes[1]).cuda()
if majors[2] == cute.nvgpu.tcgen05.OperandMajorMode.MN:
D = torch.empty(N, M).random_(-4, 4).permute(1, 0).to(dtypes[2]).cuda()
elif majors[2] == cute.nvgpu.tcgen05.OperandMajorMode.K:
D = torch.empty(M, N).random_(-4, 4).permute(0, 1).to(dtypes[2]).cuda()
return A, B, D
def get_gemm_tensors(
M,
N,
K,
majors: tuple[
cute.nvgpu.tcgen05.OperandMajorMode,
cute.nvgpu.tcgen05.OperandMajorMode,
cute.nvgpu.tcgen05.OperandMajorMode,
],
dtypes: tuple[torch.dtype, torch.dtype, torch.dtype],
):
A, B, D = create_gemm_tensors_torch(M, N, K, majors, dtypes)
A_cute = from_dlpack(A, assumed_align=16).mark_layout_dynamic(
leading_dim=1 if majors[0] == cute.nvgpu.tcgen05.OperandMajorMode.K else 0
)
B_cute = from_dlpack(B, assumed_align=16).mark_layout_dynamic(
leading_dim=1 if majors[1] == cute.nvgpu.tcgen05.OperandMajorMode.K else 0
)
D_cute = from_dlpack(D, assumed_align=16).mark_layout_dynamic(
leading_dim=1 if majors[2] == cute.nvgpu.tcgen05.OperandMajorMode.K else 0
)
return A, B, D, A_cute, B_cute, D_cute
def sm100_4x4x1_kernel_builder(
use_tma_multicast: bool,
use_2cta_instrs: bool,
acc_dtype: Type[Numeric],
M: int,
N: int,
):
CLUSTER_SHAPE = (2, 1, 1)
GRID_SHAPE = (
math.ceil(M / 128),
math.ceil(N / 256),
1,
) # TODO (xpbowler): remove hard-code
NUM_WARPS_PER_CTA = 6
TMA_STORE_PIPE_DEPTH = 4
MAINLOOP_STAGE_DEPTH = 4 # pipeline depth of TMA->MMA
# pipeline depth of mainloop->epilogue. only useful if using persistent CTA
EPILOGUE_STAGE_DEPTH = 1
# m256n256k16 2SM MMA / m128n256k16 1SM MMA
mma_inst_shape_mnk = (256, 256, 16) if use_2cta_instrs else (128, 256, 16)
@cute_ext.kernel
def kernel(
mA: cute.Tensor,
mB: cute.Tensor,
mD: cute.Tensor,
):
d_layout = utils.LayoutEnum.from_tensor(mD)
d_dtype = mD.element_type
ab_dtype = mA.element_type
mma_inst_shape_m, mma_inst_shape_n, mma_inst_shape_k = mma_inst_shape_mnk
if cutlass.const_expr(use_2cta_instrs):
cta_group = cute.nvgpu.tcgen05.CtaGroup.TWO
else:
cta_group = cute.nvgpu.tcgen05.CtaGroup.ONE
tiled_mma = sm100_utils.make_trivial_tiled_mma(
ab_dtype,
utils.LayoutEnum.from_tensor(mA).mma_major_mode(),
utils.LayoutEnum.from_tensor(mB).mma_major_mode(),
acc_dtype,
cta_group,
(mma_inst_shape_m, mma_inst_shape_n),
)
mma_inst_tile_k = (
4 # 4 MMAs per MMA tile K. For 16b types, tcgen05.mma has K=16.
)
mma_inst_tile_m = mma_inst_tile_n = 1 # 1 MMAs per MMA tile M/N
bM = mma_inst_shape_m * mma_inst_tile_m
bN = mma_inst_shape_n * mma_inst_tile_n
bK = mma_inst_shape_k * mma_inst_tile_k
mnk_tiler = (bM, bN, bK)
cta_m, cta_n, _ = cute.arch.block_idx()
tid_x, _, _ = cute.arch.thread_idx()
warp_idx = cute.arch.warp_idx()
warp_idx = cute.arch.make_warp_uniform(warp_idx)
cluster_layout_vmnk = cute.tiled_divide(
cute.make_layout(CLUSTER_SHAPE),
cute.core._pack_shape((cute.size(tiled_mma.thr_id.shape),)),
)
cluster_layout_v_size = cute.size(cluster_layout_vmnk.shape[0])
mma_coord_vmnk = (
cta_m % cluster_layout_v_size,
cta_m // cluster_layout_v_size,
cta_n,
)
gA = cute.zipped_divide(mA, (bM, bK)) # ((bM, bK), (M/bM, K/bK))
gA_tma = cute.zipped_divide(
mA, (bM // cluster_layout_v_size, bK)
) # ((bM/2, bK), (2*M/bM, K/bK))
tAgA = gA_tma[(None, None), (cta_m, None)] # ((bM/2, bK), (1, K/bK))
gB_tma = cute.zipped_divide(
mB, (bN // cluster_layout_v_size, bK)
) # ((bN/2, bK), (2*M/bM, K/bK))
# ((bN/2, bK), (1, K/bK))
tBgB = gB_tma[
(None, None),
(cluster_layout_v_size * cta_n + cta_m % cluster_layout_v_size, None),
]
gD_tma = cute.zipped_divide(
mD, (bM // cluster_layout_v_size, bN)
) # ((bM/2, bN), (2*M/bM, N/bN))
tDgD = gD_tma[(None, None), (cta_m, cta_n)] # ((bM/2, bN), (1, 1))
a_smem_layout_staged = sm100_utils.make_smem_layout_a(
tiled_mma,
mnk_tiler,
ab_dtype,
MAINLOOP_STAGE_DEPTH,
)
b_smem_layout_staged = sm100_utils.make_smem_layout_b(
tiled_mma,
mnk_tiler,
ab_dtype,
MAINLOOP_STAGE_DEPTH,
)
cta_tile_shape_mnk = cute.shape_div(mnk_tiler, (cluster_layout_v_size, 1, 1))
epi_tile = sm100_utils.compute_epilogue_tile_shape(
cta_tile_shape_mnk,
use_2cta_instrs,
d_layout,
d_dtype,
)
sc_smem_layout_staged = sm100_utils.make_smem_layout_epi(
d_dtype,
d_layout,
epi_tile,
TMA_STORE_PIPE_DEPTH,
)
acc_shape = tiled_mma.partition_shape_C(mnk_tiler[:2])
tmem_layout = tiled_mma.make_fragment_C(
cute.append(acc_shape, EPILOGUE_STAGE_DEPTH)
).layout
bufferA = cute_ext.allocate(
ab_dtype,
cute.AddressSpace.smem,
a_smem_layout_staged,
alignment=1024,
)
bufferB = cute_ext.allocate(
ab_dtype,
cute.AddressSpace.smem,
b_smem_layout_staged,
alignment=1024,
)
bufferAcc = cute_ext.allocate(
acc_dtype,
cute.AddressSpace.tmem,
tmem_layout,
alignment=16,
is2cta=use_2cta_instrs,
)
bufferC = cute_ext.allocate(
d_dtype,
cute.AddressSpace.smem,
sc_smem_layout_staged,
alignment=1024,
)
copy_atom_t2r = sm100_utils.get_tmem_load_op(
cta_tile_shape_mnk,
d_layout,
d_dtype,
acc_dtype,
epi_tile,
use_2cta_instrs,
)
# Take only one stage of the TMEM buffer for the epilogue
accumulators = cute.zipped_divide(bufferAcc, ((epi_tile), 1))
acc_epi_div = accumulators[((None, None), 0), 0]
# Create the TMEM copy atom based on the size of transfer within one iteration of epilogue
tiled_copy_t2r = cute.nvgpu.tcgen05.make_tmem_copy(copy_atom_t2r, acc_epi_div)
# Calculate the per thread destination size per iteration for output of TMEM and input of SMEM
thr_copy_t2r = tiled_copy_t2r.get_slice(tid_x)
gC_mnl_epi = cute.flat_divide(tDgD, epi_tile)
tTR_gC = thr_copy_t2r.partition_D(gC_mnl_epi)
acc_d_rmem_layout = cute.make_fragment_like(
tTR_gC[(None, None, None, 0, 0)].layout
)
bufferRAcc = cute_ext.allocate(
acc_dtype,
cute.AddressSpace.rmem,
acc_d_rmem_layout,
alignment=32,
)
bufferRD = cute_ext.allocate(
d_dtype,
cute.AddressSpace.rmem,
acc_d_rmem_layout,
alignment=32,
)
tma_mcast_proj_A = 2
tma_mcast_proj_B = 1
mma_operation_type = tma_operation_type = None
acc_pipe = mainloop_pipe = None
if cutlass.const_expr(use_2cta_instrs):
mma_operation_type = cute_ext.OperationTypeEnum.SM100_MMA_2SM_SS
if cutlass.const_expr(use_tma_multicast):
tma_operation_type = (
cute_ext.OperationTypeEnum.SM100_TMA_LOAD_2SM_MULTICAST
)
else:
tma_operation_type = cute_ext.OperationTypeEnum.SM100_TMA_LOAD_2SM
else:
mma_operation_type = cute_ext.OperationTypeEnum.SM100_MMA_1SM_SS
if cutlass.const_expr(use_tma_multicast):
tma_operation_type = cute_ext.OperationTypeEnum.SM90_TMA_LOAD_MULTICAST
else:
tma_operation_type = cute_ext.OperationTypeEnum.SM90_TMA_LOAD
# MMA <-> TMEM load pipeline
# if 2CTA MMA, warpgroup from both peer and leader CTA consumer.release
acc_pipe_consumer_arv_count = 256 if use_2cta_instrs else 128
acc_pipe = cute_ext.UMMAtoAsyncPipeline.create(
num_stages=EPILOGUE_STAGE_DEPTH,
mma_operation_type=mma_operation_type,
consumer=cute_ext.OperationTypeEnum.SM100_COPY_T2R,
consumer_arv_count=acc_pipe_consumer_arv_count,
cluster_layout_vmnk=cluster_layout_vmnk,
)
if cutlass.const_expr(use_tma_multicast):
# TMA load <-> MMA pipeline
mainloop_pipe = cute_ext.TMAToUMMAPipeline.create_with_mask(
num_stages=MAINLOOP_STAGE_DEPTH,
tma_operation_type=tma_operation_type,
mma_operation_type=mma_operation_type,
cluster_layout_vmnk=cluster_layout_vmnk,
)
else:
mainloop_pipe = cute_ext.TMAToUMMAPipeline.create(
num_stages=MAINLOOP_STAGE_DEPTH,
mma_operation_type=mma_operation_type,
tma_operation_type=tma_operation_type,
cluster_layout_vmnk=cluster_layout_vmnk,
)
tma_store_warp_id = 0
mma_warp_id = 4
tma_load_warp_id = 5
is_tma_thr = warp_idx == tma_load_warp_id
is_mma_thr = warp_idx == mma_warp_id
is_epi_thr = warp_idx < 4
is_leader_cta = mma_coord_vmnk[0] == 0
# SMEM -> GMEM
tma_store_pipe = cute_ext.TMAStorePipeline(
stages=TMA_STORE_PIPE_DEPTH,
arv_count=128,
barrier_id=1,
tma_warp_id=tma_store_warp_id,
)
k_tile_count = cute.size(gA, mode=[1, 1])
if is_tma_thr:
for k_tile in cutlass.range(0, k_tile_count, 1, unroll=1):
gA_k = tAgA[None, None, k_tile]
gB_k = tBgB[None, None, k_tile]
producer_stage_token, idx = (
mainloop_pipe.producer_acquire_and_get_stage()
)
mbar = cute_ext.get_mbarrier(producer_stage_token)
bufferA_sliced = bufferA[None, None, None, idx]
bufferB_sliced = bufferB[None, None, None, idx]
a_cta_v_map = cute_ext.get_cta_v_map_ab(mA, mnk_tiler, tiled_mma, "A")
b_cta_v_map = cute_ext.get_cta_v_map_ab(mB, mnk_tiler, tiled_mma, "B")
if cutlass.const_expr(use_tma_multicast):
cute_ext.tma_load_multicast(
gA_k,
bufferA_sliced,
mbar,
vmnk_layout=cluster_layout_vmnk,
cta_v_map=a_cta_v_map,
tma_operation_type=tma_operation_type,
multicast_mode=tma_mcast_proj_A,
)
cute_ext.tma_load_multicast(
gB_k,
bufferB_sliced,
mbar,
vmnk_layout=cluster_layout_vmnk,
cta_v_map=b_cta_v_map,
tma_operation_type=tma_operation_type,
multicast_mode=tma_mcast_proj_B,
)
else:
cute_ext.tma_load(
gA_k,
bufferA_sliced,
mbar,
cta_v_map=a_cta_v_map,
tma_operation_type=tma_operation_type,
)
cute_ext.tma_load(
gB_k,
bufferB_sliced,
mbar,
cta_v_map=b_cta_v_map,
tma_operation_type=tma_operation_type,
)
if is_leader_cta:
mainloop_pipe.producer_commit()
mainloop_pipe.producer_state = cute_ext.pipeline_advance_iterator(
mainloop_pipe.raw_pipeline, mainloop_pipe.producer_state
)
if is_mma_thr and is_leader_cta:
producer_stage_token, idx = acc_pipe.producer_acquire_and_get_stage()
accumulators_sliced = bufferAcc[None, None, None, idx]
mma_atom = cute.make_mma_atom(tiled_mma.op)
mma_atom.set(cute.nvgpu.tcgen05.Field.ACCUMULATE, False)
for k_tile in cutlass.range(0, k_tile_count, 1, unroll=1):
_, mainloop_idx = mainloop_pipe.consumer_wait_and_get_stage()
bufferA_sliced_stage = cute.core.slice_(
bufferA, (None, None, None, mainloop_idx)
)
bufferB_sliced_stage = cute.core.slice_(
bufferB, (None, None, None, mainloop_idx)
)
for k_block in cutlass.range(mma_inst_tile_k, unroll_full=True):
cute_ext.dot(
mma_atom,
cute.append_ones(
bufferA_sliced_stage[None, None, k_block], up_to_rank=3
),
cute.append_ones(
bufferB_sliced_stage[None, None, k_block], up_to_rank=3
),
accumulators_sliced,
)
mma_atom.set(cute.nvgpu.tcgen05.Field.ACCUMULATE, True)
mainloop_pipe.consumer_release_and_advance()
acc_pipe.producer_commit_and_advance()
if is_epi_thr:
_, idx = acc_pipe.consumer_wait_and_get_stage()
accumulators_sliced = bufferAcc[(None, None), 0, 0, idx]
acc_epi_div_tiled = cute.flat_divide(accumulators_sliced, epi_tile)
tiled_copy_r2s = cute.make_tiled_copy_D(
cute.make_copy_atom(cute.nvgpu.CopyUniversalOp(), d_dtype),
tiled_copy_t2r,
)
c_cta_v_map = cute_ext.get_cta_v_map_c(mD, epi_tile)
subtile_cnt = cute.size(acc_epi_div_tiled.shape, mode=[3])
for mn in range(subtile_cnt):
# TMEM -> RMEM
cute_ext.partition_and_copy(
tiled_copy_t2r.get_slice(tid_x),
acc_epi_div_tiled[None, None, 0, mn],
bufferRAcc,
)
# RMEM -> RMEM
bufferRD.store(bufferRAcc.load().to(d_dtype))
tma_store_pipe.acquire_sync()
store_idx = tma_store_pipe.get_index()
# RMEM -> SMEM
cute_ext.partition_and_copy(
tiled_copy_r2s.get_slice(tid_x),
bufferRD,
bufferC[None, None, store_idx],
)
tma_store_pipe.commit_sync()
if warp_idx == tma_store_warp_id:
cute_ext.tma_store(
bufferC[None, None, store_idx],
gC_mnl_epi[None, None, 0, mn],
cta_v_map=c_cta_v_map,
)
tma_store_pipe.release_advance()
tma_store_pipe.tail()
acc_pipe.consumer_release_and_advance()
# Return a callable that launches the kernel with proper grid/block/cluster
@cute_ext.jit
def launch_kernel(mA: cute.Tensor, mB: cute.Tensor, mD: cute.Tensor):
kernel(mA, mB, mD).launch(
grid=GRID_SHAPE,
block=(32 * NUM_WARPS_PER_CTA, 1, 1),
cluster=CLUSTER_SHAPE,
smem=cute.Int64(utils.get_smem_capacity_in_bytes("sm_100")),
)
return launch_kernel
if __name__ == "__main__":
M = 256
N = 256
K = 64
use_tma_multicast = True
use_2cta_instrs = True
acc_dtype = cutlass.Float32
majors = (
cute.nvgpu.tcgen05.OperandMajorMode.K,
cute.nvgpu.tcgen05.OperandMajorMode.K,
cute.nvgpu.tcgen05.OperandMajorMode.K,
)
dtypes = (torch.float16, torch.float16, torch.float16)
A_torch, B_torch, D_torch, A_cute, B_cute, D_cute = get_gemm_tensors(
M, N, K, majors, dtypes
)
kernel_launcher = sm100_4x4x1_kernel_builder(
use_tma_multicast, use_2cta_instrs, acc_dtype, M, N
)
compiled_kernel = cute_ext.compile(kernel_launcher, A_cute, B_cute, D_cute)
compiled_kernel(A_cute, B_cute, D_cute)
# Reference check (may fail on simulator/unsupported GPU)
try:
ref = torch.mm(A_torch.float(), B_torch.float().T)
torch.testing.assert_close(D_torch.float(), ref, atol=1e-2, rtol=1e-2)
print("PASS")
except RuntimeError as e:
if "no kernel image is available" in str(e):
print("SKIP: Reference check skipped - GPU not supported by PyTorch")
else:
raise
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,825 @@
# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: BSD-3-Clause
# Redistribution and use in source and binary forms, with or without
# modification, are permitted provided that the following conditions are met:
# 1. Redistributions of source code must retain the above copyright notice, this
# list of conditions and the following disclaimer.
# 2. Redistributions in binary form must reproduce the above copyright notice,
# this list of conditions and the following disclaimer in the documentation
# and/or other materials provided with the distribution.
# 3. Neither the name of the copyright holder nor the names of its
# contributors may be used to endorse or promote products derived from
# this software without specific prior written permission.
# THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
# AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
# IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
# DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
# FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
# DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
# SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
# CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
# OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
# OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
import argparse
import torch
from typing import Type, Tuple, List
import cutlass
from cutlass.cute import experimental as cute_ext
from cutlass.base_dsl.typing import Numeric
from cutlass import cute as cute
from cutlass import utils
from cutlass import torch as cutlass_torch
import cutlass.utils.blackwell_helpers as sm100_utils
import cutlass.cute.testing as testing
class DenseGemmPtrArrayKernel:
def __init__(
self,
mn_tiler: tuple[int, int],
mma_dtype: tuple[Type[Numeric], Type[Numeric], Type[Numeric]],
tmem_output_dtype: Type[Numeric],
batch_count: int, # Number of batches, each batch will have its own pointer for the matrix
A_shape: tuple, # Shape of the matrix A
A_stride: tuple, # Stride of the matrix A
B_shape: tuple, # Shape of the matrix B
B_stride: tuple, # Stride of the matrix B
D_shape: tuple, # Shape of the matrix D
D_stride: tuple, # Stride of the matrix D
epilogue_op=lambda x: x,
):
self.mn_tiler = mn_tiler
self.ab_dtype, self.acc_dtype, self.d_dtype = mma_dtype
self.tmem_output_dtype = tmem_output_dtype
self.use_2cta_instrs = False
self.TMA_STORE_STAGE = 4
self.epilogue_op = epilogue_op
self.batch_count = batch_count
self.A_shape = A_shape
self.A_stride = A_stride
self.B_shape = B_shape
self.B_stride = B_stride
self.D_shape = D_shape
self.D_stride = D_stride
"""
Helper function to convert an int64 to a cute.ptr of a certain type.
The cute.ptr is always located in Gmem.
This is used to load the pointers for A/B/D from the Ptr array.
"""
@cute.experimental.jit
def _get_pointer(self, address_as_int, cute_type):
cute_ptr = cute.make_ptr(
cute_type,
address_as_int,
mem_space=cute.AddressSpace.gmem,
assumed_align=16,
)
return cute_ptr
@cute.experimental.jit
def __call__(
self, mA_tensor: cute.Tensor, mB_tensor: cute.Tensor, mD_tensor: cute.Tensor
):
# Get the pointer to the first batch of D
d_ptr = self._get_pointer(mD_tensor[0], self.d_dtype)
d_ptr_base_tensor = cute.make_tensor(
d_ptr, layout=cute.make_layout(self.D_shape, stride=self.D_stride)
)
tile_mn = cute.core._pack_shape((*self.mn_tiler, 1))
div = cute.tiled_divide(d_ptr_base_tensor, tile_mn)
grid = (div.shape[1], div.shape[2], div.shape[3])
self.kernel(mA_tensor, mB_tensor, mD_tensor).launch(
grid=grid,
block=(192, 1, 1),
cluster=(1, 1, 1),
smem=cute.Int64(utils.get_smem_capacity_in_bytes("sm_100")),
)
@cute.experimental.kernel
def kernel(
self,
mA_tensor: cute.Tensor,
mB_tensor: cute.Tensor,
mD_tensor: cute.Tensor,
):
# Get pointers for the first batch to perform shape and stage calculations
A_0_ptr = self._get_pointer(mA_tensor[0], self.ab_dtype)
B_0_ptr = self._get_pointer(mB_tensor[0], self.ab_dtype)
D_0_ptr = self._get_pointer(mD_tensor[0], self.d_dtype)
mA = cute.make_tensor(
A_0_ptr, layout=cute.make_layout(self.A_shape, stride=self.A_stride)
)
mB = cute.make_tensor(
B_0_ptr, layout=cute.make_layout(self.B_shape, stride=self.B_stride)
)
mD = cute.make_tensor(
D_0_ptr, layout=cute.make_layout(self.D_shape, stride=self.D_stride)
)
tiled_mma = sm100_utils.make_trivial_tiled_mma(
self.ab_dtype,
utils.LayoutEnum.from_tensor(mA).mma_major_mode(),
utils.LayoutEnum.from_tensor(mB).mma_major_mode(),
self.acc_dtype,
cute.nvgpu.tcgen05.CtaGroup.ONE,
self.mn_tiler,
)
mma_inst_shape_k = cute.size(tiled_mma.shape_mnk, mode=[2])
mma_inst_tile_k = 4
mnk_tiler = (
self.mn_tiler[0],
self.mn_tiler[1],
mma_inst_shape_k * mma_inst_tile_k,
)
d_layout = utils.LayoutEnum.from_tensor(mD)
d_dtype = mD.element_type
tiler_mk = (mnk_tiler[0], mnk_tiler[2])
tiler_nk = (mnk_tiler[1], mnk_tiler[2])
tiler_mn = (mnk_tiler[0], mnk_tiler[1])
gA = cute.zipped_divide(mA, tiler_mk)
gB = cute.zipped_divide(mB, tiler_nk)
gD = cute.zipped_divide(mD, tiler_mn)
mainloop_stage = 2
acc_stage = 2
cta_m, cta_n, cta_l = cute.arch.block_idx()
tid_x, _, _ = cute.arch.thread_idx()
gA_tile = gA[(None, None), (cta_m, None, cta_l)]
gB_tile = gB[(None, None), (cta_n, None, cta_l)]
gD_tile = gD[(None, None), (cta_m, cta_n, cta_l)]
# Compute A/B/C shared memory layout
a_smem_layout_staged = sm100_utils.make_smem_layout_a(
tiled_mma,
mnk_tiler,
self.ab_dtype,
mainloop_stage,
)
b_smem_layout_staged = sm100_utils.make_smem_layout_b(
tiled_mma,
mnk_tiler,
self.ab_dtype,
mainloop_stage,
)
cta_tile_shape_mnk = cute.shape_div(
mnk_tiler, (cute.size(tiled_mma.thr_id.shape), 1, 1)
)
epi_tile = sm100_utils.compute_epilogue_tile_shape(
cta_tile_shape_mnk,
self.use_2cta_instrs,
d_layout,
d_dtype,
)
sc_smem_layout_staged = sm100_utils.make_smem_layout_epi(
d_dtype,
d_layout,
epi_tile,
self.TMA_STORE_STAGE,
)
# UMMA ACC TMEM Layout
acc_shape = tiled_mma.partition_shape_C(mnk_tiler[:2])
tmem_layout = tiled_mma.make_fragment_C(
cute.append(acc_shape, acc_stage)
).layout
# Allocate UMMA Buffers
bufferA = cute_ext.allocate(
self.ab_dtype,
cute.AddressSpace.smem,
a_smem_layout_staged,
alignment=1024,
)
bufferB = cute_ext.allocate(
self.ab_dtype,
cute.AddressSpace.smem,
b_smem_layout_staged,
alignment=1024,
)
bufferAcc = cute_ext.allocate(
self.acc_dtype,
cute.AddressSpace.tmem,
tmem_layout,
alignment=16,
)
# Allocate SMEM buffer for C
bufferC = cute_ext.allocate(
d_dtype,
cute.AddressSpace.smem,
sc_smem_layout_staged,
alignment=1024,
)
# Create the TMEM load atom
copy_atom_t2r = sm100_utils.get_tmem_load_op(
cta_tile_shape_mnk,
d_layout,
self.tmem_output_dtype,
self.acc_dtype,
epi_tile,
self.use_2cta_instrs,
)
# Take only one stage of the TMEM buffer
accumulators = cute.zipped_divide(bufferAcc, ((epi_tile), 1))
acc_epi_div = accumulators[((None, None), 0), 0]
# Create the TMEM copy atom based on the size of transfer within one iteration of epilogue
tiled_copy_t2r = cute.nvgpu.tcgen05.make_tmem_copy(copy_atom_t2r, acc_epi_div)
# Calculate the per thread destination size per iteration for output of TMEM and input of SMEM
thr_copy_t2r = tiled_copy_t2r.get_slice(tid_x)
gC_mnl_epi = cute.flat_divide(gD_tile, epi_tile)
tTR_gC = thr_copy_t2r.partition_D(gC_mnl_epi)
acc_d_rmem_layout = cute.make_fragment_like(
tTR_gC[(None, None, None, 0, 0)].layout
)
# Allocate RMEM buffers
bufferRAcc = cute_ext.allocate(
self.acc_dtype,
cute.AddressSpace.rmem,
acc_d_rmem_layout,
alignment=32,
)
bufferRD = cute_ext.allocate(
d_dtype,
cute.AddressSpace.rmem,
acc_d_rmem_layout,
alignment=32,
)
# TMA -> UMMA
mainloop_pipe = cute_ext.TMAToUMMAPipeline.create(
num_stages=mainloop_stage,
mma_operation_type=cute_ext.OperationTypeEnum.SM100_MMA_1SM_SS,
)
# UMMA -> TMEM
acc_pipe = cute_ext.UMMAtoAsyncPipeline.create(
num_stages=acc_stage,
mma_operation_type=cute_ext.OperationTypeEnum.SM100_MMA_1SM_SS,
consumer=cute_ext.OperationTypeEnum.SM100_COPY_T2R,
consumer_arv_count=128,
)
warp_idx = cute.arch.warp_idx()
warp_idx = cute.arch.make_warp_uniform(warp_idx)
# warp assignment: [0]-tma_store, [0-3]-epi, [4]-mma, [5]-tma_load
tma_store_warp_id = 0
mma_warp_id = 4
tma_load_warp_id = 5
is_tma_thr = warp_idx == tma_load_warp_id
is_mma_thr = warp_idx == mma_warp_id
is_epi_thr = warp_idx < 4
# SMEM -> GMEM
tma_store_pipe = cute_ext.TMAStorePipeline(
stages=self.TMA_STORE_STAGE,
arv_count=128,
barrier_id=1,
tma_warp_id=tma_store_warp_id,
)
k_tile_size = cute.size(gA, mode=[1, 1])
# Outer loop over batches and perform GEMM for each batch as usual
# This is a dynamic for loop that lowers to an scf.for
# Note that the tensor loading is done in the if `thread` warp specialized
# sections. This is essential to ensure proper synchronization of tma loads
# and tma updates across batches.
for batch_idx in range(0, self.batch_count):
# Load pointers for the current batch
ptr_A = self._get_pointer(mA_tensor[batch_idx], self.ab_dtype)
ptr_B = self._get_pointer(mB_tensor[batch_idx], self.ab_dtype)
ptr_D = self._get_pointer(mD_tensor[batch_idx], self.d_dtype)
gALayout = cute.zipped_divide(mA, tiler_mk)
k_tile_size = cute.size(gALayout, mode=[1, 1])
if is_tma_thr:
mA = cute.make_tensor(
ptr_A, layout=cute.make_layout(self.A_shape, stride=self.A_stride)
)
mB = cute.make_tensor(
ptr_B, layout=cute.make_layout(self.B_shape, stride=self.B_stride)
)
gA = cute.zipped_divide(mA, tiler_mk)
gB = cute.zipped_divide(mB, tiler_nk)
gA_tile = gA[(None, None), (cta_m, None, cta_l)]
gB_tile = gB[(None, None), (cta_n, None, cta_l)]
for k in cutlass.range(0, k_tile_size, 1, unroll=1):
gA_k = gA_tile[None, None, k]
gB_k = gB_tile[None, None, k]
# Scoped state management - pipeline object manages state internally
(
producer_stage_token,
idx,
) = mainloop_pipe.producer_acquire_and_get_stage()
mbar = cute_ext.get_mbarrier(producer_stage_token)
## producer_body begin ##
bufferA_sliced = bufferA[None, None, None, idx]
bufferB_sliced = bufferB[None, None, None, idx]
a_cta_v_map = cute_ext.get_cta_v_map_ab(
mA, mnk_tiler, tiled_mma, "A"
)
b_cta_v_map = cute_ext.get_cta_v_map_ab(
mB, mnk_tiler, tiled_mma, "B"
)
cute_ext.tma_load(
gA_k,
bufferA_sliced,
mbar,
cta_v_map=a_cta_v_map,
)
cute_ext.tma_load(
gB_k,
bufferB_sliced,
mbar,
cta_v_map=b_cta_v_map,
)
## producer_body end ##
mainloop_pipe.producer_commit_and_advance()
# MMA section remains same as a regular GEMM
if is_mma_thr:
producer_stage_token, idx = acc_pipe.producer_acquire_and_get_stage()
## acc_producer_body begin ##
accumulators_sliced = bufferAcc[None, None, None, idx]
mma_atom = cute.make_mma_atom(tiled_mma.op)
mma_atom.set(cute.nvgpu.tcgen05.Field.ACCUMULATE, False)
for k_tile in cutlass.range(0, k_tile_size, 1, unroll=1):
# Scoped state management - pipeline object manages consumer state internally
(
_,
mainloop_idx,
) = mainloop_pipe.consumer_wait_and_get_stage()
## tma_consumer_body begin ##
bufferA_sliced_stage = cute.core.slice_(
bufferA, (None, None, None, mainloop_idx)
)
bufferB_sliced_stage = cute.core.slice_(
bufferB, (None, None, None, mainloop_idx)
)
for k_block in cutlass.range(mma_inst_tile_k, unroll_full=True):
bufferA_sliced = bufferA_sliced_stage[None, None, k_block]
bufferB_sliced = bufferB_sliced_stage[None, None, k_block]
cute_ext.dot(
mma_atom,
cute.append_ones(bufferA_sliced, up_to_rank=3),
cute.append_ones(bufferB_sliced, up_to_rank=3),
accumulators_sliced,
)
mma_atom.set(cute.nvgpu.tcgen05.Field.ACCUMULATE, True)
## tma_consumer_body end ##
mainloop_pipe.consumer_release_and_advance()
## acc_producer_body end ##
acc_pipe.producer_commit_and_advance()
if is_epi_thr:
# Load the D tensor in the warp specialized section
mD = cute.make_tensor(
ptr_D, layout=cute.make_layout(self.D_shape, stride=self.D_stride)
)
gD = cute.zipped_divide(mD, tiler_mn)
gD_tile = gD[(None, None), (cta_m, cta_n, cta_l)]
gC_mnl_epi = cute.flat_divide(gD_tile, epi_tile)
_, idx = acc_pipe.consumer_wait_and_get_stage()
## acc_consume_body begin ##
accumulators_sliced = bufferAcc[(None, None), 0, 0, idx]
acc_epi_div_tiled = cute.flat_divide(accumulators_sliced, epi_tile)
subtile_cnt = cute.size(acc_epi_div_tiled.shape, mode=[3])
for mn in range(subtile_cnt):
# TMEM -> RMEM
cute_ext.partition_and_copy(
tiled_copy_t2r.get_slice(tid_x),
acc_epi_div_tiled[None, None, 0, mn],
bufferRAcc,
)
# RMEM -> RMEM
bufferRD.store(self.epilogue_op(bufferRAcc.load().to(self.d_dtype)))
# Acquire pipeline stage and synchronize before RMEM->SMEM copy
tma_store_pipe.acquire_sync()
idx = tma_store_pipe.get_index()
# RMEM -> SMEM
tiled_copy_r2s = cute.make_tiled_copy_D(
cute.make_copy_atom(cute.nvgpu.CopyUniversalOp(), self.d_dtype),
tiled_copy_t2r,
)
cute_ext.partition_and_copy(
tiled_copy_r2s.get_slice(tid_x),
bufferRD,
bufferC[None, None, idx],
)
# Fence SMEM writes and synchronize before TMA store
tma_store_pipe.commit_sync()
# SMEM -> GMEM (only designated TMA store warp performs TMA store)
if warp_idx == tma_store_warp_id:
c_cta_v_map = cute_ext.get_cta_v_map_c(mD, epi_tile)
cute_ext.tma_store(
bufferC[None, None, idx],
gC_mnl_epi[None, None, 0, mn],
cta_v_map=c_cta_v_map,
)
# Release pipeline stage and advance
tma_store_pipe.release_advance()
tma_store_pipe.tail()
acc_pipe.consumer_release_and_advance()
def create_tensors(l, m, n, k, a_major, b_major, d_major, ab_dtype, d_dtype):
torch.manual_seed(1111)
a_torch_cpu = cutlass_torch.matrix(l, m, k, a_major == "m", ab_dtype)
b_torch_cpu = cutlass_torch.matrix(l, n, k, b_major == "n", ab_dtype)
d_torch_cpu = cutlass_torch.matrix(l, m, n, d_major == "m", d_dtype)
a_tensor, a_torch_gpu = cutlass_torch.cute_tensor_like(
a_torch_cpu, ab_dtype, is_dynamic_layout=True, assumed_align=16
)
b_tensor, b_torch_gpu = cutlass_torch.cute_tensor_like(
b_torch_cpu, ab_dtype, is_dynamic_layout=True, assumed_align=16
)
d_tensor, d_torch_gpu = cutlass_torch.cute_tensor_like(
d_torch_cpu, d_dtype, is_dynamic_layout=True, assumed_align=16
)
return (
a_tensor,
b_tensor,
d_tensor,
a_torch_cpu,
b_torch_cpu,
d_torch_cpu,
a_torch_gpu,
b_torch_gpu,
d_torch_gpu,
)
# Helper creates a cute.Tensor from a List of device pointers
def make_tensor_of_ptrs(torch_tensor_array: List):
tensor_of_ptrs_torch = torch.tensor(
[t.data_ptr() for t in torch_tensor_array],
dtype=torch.int64,
device="cuda",
requires_grad=False,
)
tensor_of_ptrs_cute, backing_torch_tensor = cutlass_torch.cute_tensor_like(
tensor_of_ptrs_torch,
cutlass.Int64,
is_dynamic_layout=False,
assumed_align=16,
)
return tensor_of_ptrs_cute, backing_torch_tensor
def create_tensors_for_ptr_array(
l, m, n, k, a_major, b_major, d_major, ab_dtype, d_dtype
):
# Store torch gpu pointers
As_torch_gpu = []
Bs_torch_gpu = []
Ds_torch_gpu = []
# Store cute tensors
A_cutes = []
B_cutes = []
D_cutes = []
for batch_idx in range(l):
torch.manual_seed(111 + batch_idx)
(
A_tensor,
B_tensor,
D_tensor,
A_torch_cpu,
B_torch_cpu,
D_torch_cpu,
A_torch_gpu,
B_torch_gpu,
D_torch_gpu,
) = create_tensors(
1, # outer loop creates a new tensor for each batch
m,
n,
k,
a_major,
b_major,
d_major,
ab_dtype,
d_dtype,
)
A_cutes.append(A_tensor)
B_cutes.append(B_tensor)
D_cutes.append(D_tensor)
As_torch_gpu.append(A_torch_gpu)
Bs_torch_gpu.append(B_torch_gpu)
Ds_torch_gpu.append(D_torch_gpu)
# Create cute tensors of pointers
a_tensor, a_backing_torch_tensor = make_tensor_of_ptrs(As_torch_gpu)
b_tensor, b_backing_torch_tensor = make_tensor_of_ptrs(Bs_torch_gpu)
d_tensor, d_backing_torch_tensor = make_tensor_of_ptrs(Ds_torch_gpu)
return (
a_tensor,
b_tensor,
d_tensor,
a_backing_torch_tensor,
b_backing_torch_tensor,
d_backing_torch_tensor,
A_cutes,
B_cutes,
D_cutes,
As_torch_gpu,
Bs_torch_gpu,
Ds_torch_gpu,
)
def compare(a_torch_cpu, b_torch_cpu, d_torch_gpu, d_dtype, tolerance):
ref = torch.einsum("mkl,nkl->mnl", a_torch_cpu, b_torch_cpu)
_, ref_torch_gpu = cutlass_torch.cute_tensor_like(
ref, d_dtype, is_dynamic_layout=True, assumed_align=16
)
ref_result = ref_torch_gpu.cpu()
torch.testing.assert_close(
d_torch_gpu.cpu(), ref_result, atol=tolerance, rtol=1e-05
)
def run(
mnkl: Tuple[int, int, int, int],
mma_tiler_mn: Tuple[int, int],
cluster_shape_mn: Tuple[int, int],
ab_dtype: Type[Numeric],
c_dtype: Type[Numeric],
acc_dtype: Type[Numeric],
a_major: str,
b_major: str,
c_major: str,
warmup_iterations: int = 0,
iterations: int = 1,
use_cold_l2: bool = False,
tolerance: float = 1e-02,
skip_ref_check: bool = False,
**kwargs,
):
"""Execute a Pointer array batched dense GEMM operation on Blackwell architecture with performance benchmarking.
The main difference between this and a regular bathced GEMM is that the inputs to the kernel are arrays of pointers.
Every batch of each operand (A/B/D) has its own pointer. Thus, the size of the array of pointers is the batch size.
These pointers NEED NOT be stored contiguously in memory.
This example also demonstrates how cute_ext.tma_load/cute_ext.tma_store performs automatic device side TMA updates.
Note that the dimensions of the operand for each batch are the same across all batches. That is, all batches of A have the same shape and stride, same for B and D.
This function prepares input tensors, configures and launches the GEMM kernel,
optionally performs reference validation, and benchmarks the execution performance.
:param mnkl: Problem size (M, N, K, L)
:type mnkl: Tuple[int, int, int, int]
:param mma_tiler_mn: MMA tiling size.
:type mma_tiler_mn: Tuple[int, int]
:param cluster_shape_mn: Cluster shape.
:type cluster_shape_mn: Tuple[int, int]
:param ab_dtype: Data type for input tensors A and B
:type ab_dtype: Type[Numeric]
:param d_dtype: Data type for output tensor D
:type d_dtype: Type[Numeric]
"""
print("Running Blackwell Dense GEMM test with:")
print(f"mnkl: {mnkl}")
print(f"AB dtype: {ab_dtype}, D dtype: {c_dtype}, Acc dtype: {acc_dtype}")
print(f"Matrix majors - A: {a_major}, B: {b_major}, D: {c_major}")
print(f"Mma Tiler (M, N): {mma_tiler_mn}, Cluster Shape (M, N): {cluster_shape_mn}")
print(f"Tolerance: {tolerance}")
print(f"Warmup iterations: {warmup_iterations}")
print(f"Iterations: {iterations}")
print(f"Skip reference checking: {skip_ref_check}")
print(f"Use cold L2: {'True' if use_cold_l2 else 'False'}")
m, n, k, l = mnkl
ab_dtype = ab_dtype
d_major = c_major
d_dtype = c_dtype
# a_tensor, b_tensor, d_tensor are cute Tensors where each element is an Int64 pointer to global memory
# A_cutes, B_cutes, D_cutes are lists of cute Tensors for each batch of A/B/D
(
a_tensor,
b_tensor,
d_tensor,
a_backing_torch_tensor,
b_backing_torch_tensor,
d_backing_torch_tensor,
A_cutes,
B_cutes,
D_cutes,
As_torch_gpu,
Bs_torch_gpu,
Ds_torch_gpu,
) = create_tensors_for_ptr_array(
l, m, n, k, a_major, b_major, d_major, ab_dtype, d_dtype
)
ptr_array_dense_gemm = DenseGemmPtrArrayKernel(
mn_tiler=mma_tiler_mn,
mma_dtype=(ab_dtype, acc_dtype, d_dtype),
tmem_output_dtype=d_dtype,
batch_count=l,
A_shape=A_cutes[0].shape,
A_stride=A_cutes[0].stride,
B_shape=B_cutes[0].shape,
B_stride=B_cutes[0].stride,
D_shape=D_cutes[0].shape,
D_stride=D_cutes[0].stride,
)
compiled_dense_gemm = cute_ext.compile(
ptr_array_dense_gemm, a_tensor, b_tensor, d_tensor
)
compiled_dense_gemm(a_tensor, b_tensor, d_tensor)
if not skip_ref_check:
for batch_idx in range(l):
compare(
As_torch_gpu[batch_idx].cpu(),
Bs_torch_gpu[batch_idx].cpu(),
Ds_torch_gpu[batch_idx],
d_dtype,
tolerance,
)
print("check reference: PASS")
def generate_tensors():
(
a_tensor,
b_tensor,
d_tensor,
a_backing_torch_tensor,
b_backing_torch_tensor,
d_backing_torch_tensor,
A_cutes,
B_cutes,
D_cutes,
As_torch_gpu,
Bs_torch_gpu,
Ds_torch_gpu,
) = create_tensors_for_ptr_array(
l, m, n, k, a_major, b_major, d_major, ab_dtype, d_dtype
)
args = testing.JitArguments(a_tensor, b_tensor, d_tensor)
args.add_to_scope([A_cutes, B_cutes, D_cutes])
return args
workspace_count = 1
if use_cold_l2:
one_workspace_bytes = (
sum(
As_torch_gpu[batch_idx].numel() * As_torch_gpu[batch_idx].element_size()
for batch_idx in range(l)
)
+ sum(
Bs_torch_gpu[batch_idx].numel() * Bs_torch_gpu[batch_idx].element_size()
for batch_idx in range(l)
)
+ sum(
Ds_torch_gpu[batch_idx].numel() * Ds_torch_gpu[batch_idx].element_size()
for batch_idx in range(l)
)
)
workspace_count = testing.get_workspace_count(
one_workspace_bytes, warmup_iterations, iterations
)
exec_time = testing.benchmark(
compiled_dense_gemm,
workspace_generator=generate_tensors,
workspace_count=workspace_count,
warmup_iterations=warmup_iterations,
iterations=iterations,
)
return exec_time
if __name__ == "__main__":
def parse_comma_separated_ints(s: str) -> Tuple[int, ...]:
try:
return tuple(int(x.strip()) for x in s.split(","))
except ValueError:
raise argparse.ArgumentTypeError(
"Invalid format. Expected comma-separated integers."
)
parser = argparse.ArgumentParser(description="Example of Dense GEMM on Blackwell.")
parser.add_argument(
"--mnkl",
type=parse_comma_separated_ints,
default=(256, 256, 512, 1),
help="mnkl dimensions (comma-separated)",
)
parser.add_argument(
"--mma_tiler_mn",
type=parse_comma_separated_ints,
default=(128, 128),
help="Mma tile shape (comma-separated)",
)
parser.add_argument(
"--cluster_shape_mn",
type=parse_comma_separated_ints,
default=(1, 1),
help="Cluster shape (comma-separated)",
)
parser.add_argument("--ab_dtype", type=cutlass.dtype, default=cutlass.Float32)
parser.add_argument("--d_dtype", type=cutlass.dtype, default=cutlass.Float32)
parser.add_argument("--acc_dtype", type=cutlass.dtype, default=cutlass.Float32)
parser.add_argument("--a_major", choices=["k", "m"], type=str, default="k")
parser.add_argument("--b_major", choices=["k", "n"], type=str, default="k")
parser.add_argument("--d_major", choices=["n", "m"], type=str, default="n")
parser.add_argument(
"--warmup_iterations", type=int, default=0, help="Warmup iterations"
)
parser.add_argument(
"--iterations", type=int, default=1, help="Number of iterations"
)
parser.add_argument("--use_cold_l2", action="store_true", help="Use cold L2")
parser.add_argument(
"--tolerance", type=float, default=1e-02, help="Tolerance for validation"
)
parser.add_argument(
"--skip_ref_check", action="store_true", help="Skip reference checking"
)
args = parser.parse_args()
if len(args.mnkl) != 4:
parser.error("--mnkl must contain exactly 4 values")
if len(args.mma_tiler_mn) != 2:
parser.error("--mma_tiler_mn must contain exactly 2 values")
exec_time = run(
args.mnkl,
args.mma_tiler_mn,
args.cluster_shape_mn,
args.ab_dtype,
args.d_dtype,
args.acc_dtype,
args.a_major,
args.b_major,
args.d_major,
args.warmup_iterations,
args.iterations,
args.use_cold_l2,
args.tolerance,
args.skip_ref_check,
)
print(f"Execution time: {exec_time} microseconds per iteration")
@@ -832,6 +832,12 @@ public:
mainloop_pipeline.init_masks(cluster_shape, block_id_in_cluster);
accumulator_pipeline.init_masks(cluster_shape, block_id_in_cluster);
// Ensure that the prefetched kernel does not touch
// unflushed global memory prior to this instruction.
// For the static grouped scheduler, the problem shapes
// might be produced by a previous kernel in global memory.
cutlass::arch::wait_on_dependent_grids();
// TileID scheduler
TileScheduler scheduler(
(!IsTensorMapUpdateAsync || is_participant.sched || is_participant.tensor_map_updater)
@@ -842,12 +848,6 @@ public:
);
auto work_tile_info = [&] () {
// Ensure that the prefetched kernel does not touch
// unflushed global memory prior to this instruction.
// For the static grouped scheduler, the problem shapes
// might be produced by a previous kernel in global memory.
cutlass::arch::wait_on_dependent_grids();
if constexpr (IsTensorMapUpdateAsync) {
return scheduler.initial_work_tile_info(cluster_shape, [] (typename TileScheduler::CLCResponse response) {
CLCResponseWithAdditionalInformation response_with_additional_info = response;
@@ -652,12 +652,12 @@ public:
// Allocate accumulators
auto acc_shape = collective_mainloop.partition_accumulator_shape();
// TileID scheduler
TileScheduler scheduler(&shared_storage.clc_response[0], params.scheduler, block_id_in_cluster);
// Ensure memory ops in this kernel are not done prior to completion of dependent grids.
cutlass::arch::wait_on_dependent_grids();
// TileID scheduler
TileScheduler scheduler(&shared_storage.clc_response[0], params.scheduler, block_id_in_cluster);
typename TileScheduler::WorkTileInfo work_tile_info = scheduler.initial_work_tile_info(cluster_shape);
auto cta_coord_mnkl = scheduler.work_tile_to_cta_coord(work_tile_info);
@@ -735,12 +735,12 @@ public:
Tensor accumulators = cutlass::detail::make_sm100_accumulator<AccumulatorPipelineStageCount, IsOverlappingAccum>(
tiled_mma, acc_shape, EpilogueTile{});
// TileID scheduler
TileScheduler scheduler(&shared_storage.clc_response[0], params.scheduler, block_id_in_cluster);
// Ensure memory ops in this kernel are not done prior to completion of dependent grids.
cutlass::arch::wait_on_dependent_grids();
// TileID scheduler
TileScheduler scheduler(&shared_storage.clc_response[0], params.scheduler, block_id_in_cluster);
typename TileScheduler::WorkTileInfo work_tile_info = scheduler.initial_work_tile_info(cluster_shape);
auto cta_coord_mnkl = scheduler.work_tile_to_cta_coord(work_tile_info);
@@ -116,12 +116,15 @@ public:
CUTLASS_DEVICE
PersistentTileSchedulerSm100Group() { }
// Note: constructing this tile scheduler can touch global memory that was
// written to by the prior kernel.
CUTLASS_DEVICE
PersistentTileSchedulerSm100Group(CLCResponse* clc_response_ptr, Params const& params)
: scheduler_params(params),
scheduler_sm90(params.params_sm90_, clc_response_ptr) { }
// Note: constructing this tile scheduler can touch global memory that was
// written to by the prior kernel.
CUTLASS_DEVICE
PersistentTileSchedulerSm100Group(CLCResponse* clc_response_ptr, Params const& params, dim3 /* block_id_in_cluster */)
: scheduler_params(params),
@@ -752,12 +752,12 @@ public:
Tensor accumulators = cutlass::detail::make_sm100_accumulator<AccumulatorPipelineStageCount, IsOverlappingAccum>(
tiled_mma, acc_shape, EpilogueTile{});
// TileID scheduler
TileScheduler scheduler(&shared_storage.clc_response[0], params.scheduler, block_id_in_cluster);
// Ensure memory ops in this kernel are not done prior to completion of dependent grids.
cutlass::arch::wait_on_dependent_grids();
// TileID scheduler
TileScheduler scheduler(&shared_storage.clc_response[0], params.scheduler, block_id_in_cluster);
typename TileScheduler::WorkTileInfo work_tile_info = scheduler.initial_work_tile_info(cluster_shape);
auto cta_coord_mnkl = scheduler.work_tile_to_cta_coord(work_tile_info);
@@ -454,16 +454,6 @@ public:
// Kernel level shared memory storage
SharedStorage& shared_storage = *reinterpret_cast<SharedStorage*>(smem_buf);
auto scheduler = [&] () {
// Group scheduler requires a different constructor that takes a response ptr
if constexpr (cute::is_same_v<SchedulerTag, GroupScheduler>) {
return TileScheduler{params.scheduler, shared_storage.scheduler_response};
}
else {
return TileScheduler{params.scheduler};
}
} ();
// In a warp specialized kernel, collectives expose data movement and compute operations separately
CollectiveMainloop collective_mainloop;
CollectiveEpilogue collective_epilogue(params.epilogue, shared_storage.tensors.epilogue);
@@ -585,6 +575,16 @@ public:
// Ensure memory ops in this kernel are not done prior to completion of dependent grids.
cutlass::arch::wait_on_dependent_grids();
auto scheduler = [&] () {
// Group scheduler requires a different constructor that takes a response ptr
if constexpr (cute::is_same_v<SchedulerTag, GroupScheduler>) {
return TileScheduler{params.scheduler, shared_storage.scheduler_response};
}
else {
return TileScheduler{params.scheduler};
}
} ();
auto work_tile_info = scheduler.initial_work_tile_info(ClusterShape{});
if (not work_tile_info.is_valid()) {
@@ -463,16 +463,6 @@ public:
// Kernel level shared memory storage
SharedStorage& shared_storage = *reinterpret_cast<SharedStorage*>(smem_buf);
auto scheduler = [&] () {
// Group scheduler requires a different constructor that takes a response ptr
if constexpr (cute::is_same_v<SchedulerTag, GroupScheduler>) {
return TileScheduler{params.scheduler, shared_storage.scheduler_response};
}
else {
return TileScheduler{params.scheduler};
}
} ();
// In a warp specialized kernel, collectives expose data movement and compute operations separately
CollectiveMainloop collective_mainloop;
CollectiveEpilogue collective_epilogue(params.epilogue, shared_storage.tensors.epilogue);
@@ -600,6 +590,16 @@ public:
// Ensure memory ops in this kernel are not done prior to completion of dependent grids.
cutlass::arch::wait_on_dependent_grids();
auto scheduler = [&] () {
// Group scheduler requires a different constructor that takes a response ptr
if constexpr (cute::is_same_v<SchedulerTag, GroupScheduler>) {
return TileScheduler{params.scheduler, shared_storage.scheduler_response};
}
else {
return TileScheduler{params.scheduler};
}
} ();
auto work_tile_info = scheduler.initial_work_tile_info(ClusterShape{});
if (not work_tile_info.is_valid()) {
@@ -246,6 +246,8 @@ public:
PersistentTileSchedulerSm90Group() = default;
// Note: constructing this tile scheduler can touch global memory that was
// written to by the prior kernel.
CUTLASS_DEVICE explicit PersistentTileSchedulerSm90Group(Params const& params_, SchedulerResponse* response_ptr) : scheduler_params(params_), response_ptr_(response_ptr) {
// MSVC requires protecting use of CUDA-specific nonstandard syntax,
// like blockIdx and gridDim, with __CUDA_ARCH__.
+8
View File
@@ -8,6 +8,14 @@ The CUTLASS DSL 4.4 release currently supports **Linux** and **Python 3.10 - 3.1
Installation
-----------------------
Before installing the latest version, you need to uninstall any previous CUTLASS DSL Installation.
.. code-block:: bash
pip uninstall nvidia-cutlass-dsl nvidia-cutlass-dsl-libs-base nvidia-cutlass-dsl-libs-cu13 -y
To ensure compatibility with the examples and code on `GitHub <https://github.com/NVIDIA/cutlass/tree/main>`_,
use the `setup.sh <https://github.com/NVIDIA/cutlass/blob/main/python/CuTeDSL/setup.sh>`_ file from the corresponding commit in the repository.
+11 -3
View File
@@ -14,6 +14,16 @@ from ._mlir._mlir_libs import _cutlass_ir
_cutlass_ir.populate(_cutlass_ir)
__version__ = "@CUTLASS_IR_WHEEL_RELEASE_VERSION@"
# Monkey patch CUDA version query function
from ._mlir._mlir_libs._cutlass_ir._base_dsl import (
get_cuda_version as _get_cuda_version,
)
from .base_dsl import common as _common
_common._get_cuda_version = _get_cuda_version
# Import CUDA version from base_dsl
from .base_dsl.version_info import CUDA_VERSION
from .cutlass_dsl import (
Constexpr,
@@ -47,6 +57,7 @@ from .cutlass_dsl import (
extract_mlir_values,
new_from_mlir_values,
DSLCudaVersion,
target_version,
)
from .cute.typing import *
@@ -54,7 +65,6 @@ from .cute.typing import *
# Utilities not belonging to CuTe
from . import utils as utils
from . import pipeline as pipeline
from .utils.version_info import CUDA_VERSION
# Used as internal symbol
from . import cutlass_dsl as _dsl
@@ -67,5 +77,3 @@ cuda = _dsl.cuda_helpers
# Jax Framework support
from . import jax as jax
CACHE_FILE = "compiled_cache.db"
+1 -1
View File
@@ -24,4 +24,4 @@ from .utils.tree_utils import (
DSLTreeFlattenError,
)
from .common import DSLCudaVersion
from .common import DSLCudaVersion, target_version
+34 -5
View File
@@ -52,6 +52,7 @@ class Executor:
self._any_executor = None
self._all_executor = None
self._builtin_redirector = None
self._ifexp_dynamic = None
def set_functions(
self,
@@ -64,6 +65,7 @@ class Executor:
any_executor: Callable = None,
all_executor: Callable = None,
builtin_redirector: Callable = None,
ifexp_dynamic: Callable = None,
):
self._is_dynamic_expression = is_dynamic_expression
self._loop_execute_range_dynamic = loop_execute_range_dynamic
@@ -73,6 +75,7 @@ class Executor:
self._any_executor = any_executor
self._all_executor = all_executor
self._builtin_redirector = builtin_redirector
self._ifexp_dynamic = ifexp_dynamic
@staticmethod
def convert_to_list(x):
@@ -173,6 +176,16 @@ class Executor:
write_args_names,
)
def ifexp_execute(
self,
pred,
generator_targets: tuple,
then_block: Callable,
else_block: Callable,
):
assert self._ifexp_dynamic, "Functions must be set before execution."
return self._ifexp_dynamic(pred, generator_targets, then_block, else_block)
# =============================================================================
# Decorator
@@ -293,6 +306,21 @@ def if_executor(
)
def ifExp_executor(
*,
pred,
generator_targets: tuple,
then_block: Callable,
else_block: Callable,
):
if not executor._is_dynamic_expression(pred):
return (
then_block(*generator_targets) if pred else else_block(*generator_targets)
)
else:
return executor.ifexp_execute(pred, generator_targets, then_block, else_block)
# =============================================================================
# Range
# =============================================================================
@@ -552,18 +580,19 @@ def cf_symbol_check(symbol):
name = symbol.__name__
self_module = _get_self_module()
if inspect.ismodule(symbol):
name = "range"
if not self_module.__name__.startswith(symbol.__name__):
if not self_module.__name__.startswith(name):
failed = True
else:
owning_module = inspect.getmodule(symbol)
if owning_module != self_module:
root_module = owning_module.__name__.split(".")[0]
self_root_module = self_module.__name__.split(".")[0]
if root_module != self_root_module:
failed = True
if failed:
raise DSLRuntimeError(
f"Incorrect {symbol.__name__} is used.",
suggestion=f"Please avoid overriding `{symbol.__name__}` from DSL package.",
f"Incorrect `{name}` is used.",
suggestion=f"Please avoid overriding `{name}` from DSL package.",
)
@@ -140,25 +140,150 @@ class ScopeManager:
"""
scopes: List[Set[str]]
callables: List[Set[str]]
@classmethod
def create(cls) -> "ScopeManager":
return cls([])
return cls([], [])
def add_to_scope(self, name: str) -> None:
if name == "_":
return
self.scopes[-1].add(name)
def add_to_callables(self, name: str) -> None:
if not self.callables:
return
self.callables[-1].add(name)
def get_active_symbols(self) -> List[Set[str]]:
return self.scopes.copy()
def __enter__(self) -> "ScopeManager":
def get_active_callables(self) -> List[Set[str]]:
return self.callables.copy()
@contextlib.contextmanager
def enter_local_scope(self):
"""
Context manager for entering a new local variable and callable scope.
This is conceptually Python's local scope, such as within a function or class definition.
Use this in a ``with`` statement to temporarily push a new, empty set for both variable and callable
tracking onto the respective ScopeManager stacks. These sets accumulate any new symbols
introduced within the local context. When the context manager exits, the local sets are popped,
restoring the previous scope state.
**Example**
.. code-block:: python
with scope_manager.enter_local_scope():
# Symbols defined here are local to this scope
...
:yields: None
"""
self.scopes.append(set())
self.callables.append(set())
yield
self.scopes.pop()
self.callables.pop()
@contextlib.contextmanager
def enter_control_flow_scope(self):
"""
Context manager for entering a new dynamic control-flow scope.
This scope rule diverge from Python's local scope, variables defined here are discarded after exiting the block, but callables are kept in parent scope.
This context manager pushes a new, empty variable scope onto the stack for the
duration of a control-flow block (such as within loops or if/else blocks). Variables
introduced inside this block are tracked separately and discarded after exiting the block.
Callable symbol scopes are not affected.
:yields: None
**Example**
.. code-block:: python
with scope_manager.enter_control_flow_scope():
# Variables defined here are local to this control-flow scope
...
"""
self.scopes.append(set())
yield
self.scopes.pop()
class Region:
"""
Context manager for handling regions during AST transformations.
This class is used to manage region-scoped state during DSL preprocessing.
It is responsible for tracking and collecting new statements generated while
visiting and transforming regions, such as the bodies of AST nodes representing
constructs like loops or conditional blocks.
Upon entering a region (using a ``with`` statement), the region is pushed onto
the session's ``region_stack``, and prepares a place for new statements to be collected.
On exit, the region is popped from the stack and any temporary state is cleaned up.
Parameters
----------
session_data : SessionData
The shared session context for the AST preprocessor, which holds the region stack.
owning_node : Optional[ast.stmt], default=None
If provided, the AST statement node that owns this region; new statements will be append to _new_value of this new node.
new_value : Optional[list[ast.stmt]], default=None
If provided, a list for collecting new statements for this region.
Methods
-------
__enter__()
Enter the region context, mutate state as needed.
__exit__(exc_type, exc_value, traceback)
Exit the context, clean up state.
append_new_stmts(stmts)
Append new AST statements to the region's collection.
"""
def __init__(
self,
session_data: "SessionData",
*,
owning_node: ast.stmt = None,
new_value: list[ast.stmt] = None,
):
self.session_data = session_data
self.owning_node = owning_node
self.new_value = new_value
def __enter__(self):
if self.new_value is not None or isinstance(self.owning_node, ast.stmt):
self.session_data.region_stack.append(self)
if self.owning_node is not None:
self.owning_node._new_value = []
return self
def __exit__(self, exc_type, exc_value, traceback) -> None:
self.scopes.pop()
def __exit__(self, exc_type, exc_value, traceback):
if self.new_value is not None or isinstance(self.owning_node, ast.stmt):
self.session_data.region_stack.pop()
if self.owning_node is not None:
delattr(self.owning_node, "_new_value")
def append_new_stmts(self, stmts: list[ast.stmt]):
"""
Append a list of statements to the region's collection.
Parameters
----------
stmts : list[ast.stmt]
The AST statements to append to this region.
"""
if self.owning_node is not None:
self.owning_node._new_value.extend(stmts)
else:
self.new_value.extend(stmts)
@dataclass
@@ -173,10 +298,24 @@ class SessionData:
function_name: str = "<unknown function>"
class_name: Optional[str] = None
file_name: str = "<unknown filename>"
function_depth: int = 0
local_closures: set[str] = field(default_factory=set)
function_globals: Optional[dict[str, Any]] = None
import_top_module: bool = False
region_stack: list[Region] = field(default_factory=list)
generator_targets: list[str] = field(default_factory=list)
@contextlib.contextmanager
def set_current_class_name(self, class_name: str):
old_class_name = self.class_name
self.class_name = class_name
yield
self.class_name = old_class_name
@contextlib.contextmanager
def set_current_function_name(self, function_name: str):
old_function_name = self.function_name
self.function_name = function_name
yield
self.function_name = old_function_name
def _create_module_attribute(
@@ -225,54 +364,6 @@ def _create_module_attribute(
set_location(node, lineno, col_offset)
return node
class DSLPreprocessorSession:
"""Context manager for managing a DSL preprocessor session.
This context manager is used to ensure that each preprocessing operation
(typically a transformation of a Python AST via the DSL preprocessor)
is performed within a well-defined session. When entering the context,
it initializes session-specific resources or state by calling
`_start_session()` on the provided DSL object. Upon exit, it performs
appropriate cleanup by calling `_end_session()`.
Example usage::
with DSLPreprocessorSession(dsl_object):
# perform AST transformations or other preprocessing actions
:param dsl_object: An instance of a DSL object that implements
`_start_session()` and `_end_session()` methods to manage the
session state
:type dsl_object: DSLPreprocessor
"""
def __init__(self, dsl_object):
self.dsl_object = dsl_object
def __enter__(self):
"""Starts the DSL preprocessor session.
:return: The DSL object for use within the context
:rtype: DSLPreprocessor
"""
self.dsl_object._start_session()
# Let `with preprocessor.get_session() as p:` keep using `p` as the preprocessor.
return self.dsl_object
def __exit__(self, exc_type, exc_value, traceback):
"""Ends the DSL preprocessor session.
:param exc_type: The exception type if an exception was raised in the context
:type exc_type: type, optional
:param exc_value: The exception value if an exception was raised in the context
:type exc_value: Exception, optional
:param traceback: The traceback if an exception was raised in the context
:type traceback: traceback, optional
"""
self.dsl_object._end_session()
class DSLPreprocessor(ast.NodeTransformer):
"""
A preprocessor for transforming Python ASTs. It supports:
@@ -286,15 +377,52 @@ class DSLPreprocessor(ast.NodeTransformer):
DECORATOR_IF_STATEMENT = "if_selector"
DECORATOR_WHILE_STATEMENT = "while_selector"
IF_EXECUTOR = "if_executor"
IFEXP_EXECUTOR = "ifExp_executor"
WHILE_EXECUTOR = "while_executor"
ASSERT_EXECUTOR = "assert_executor"
BOOL_CAST = "bool_cast"
IMPLICIT_DOWNCAST_NUMERIC_TYPE = "implicitDowncastNumericType"
SUPPORTED_FOR_RANGE_STATEMENTS = {"range", "range_dynamic", "range_constexpr"}
CONST_EXPR_NAME = {"const_expr", "target_version"}
COMPARE_EXECUTOR = "compare_executor"
ANY_EXECUTOR = "any_executor"
ALL_EXECUTOR = "all_executor"
def generic_visit(self, node):
"""
Copy of :meth:`ast.NodeTransformer.generic_visit` with support for inserting statements during expression visits.
This version provides the same recursive traversal and transformation as the standard
``generic_visit``, but extends it to allow statement insertion when visiting expressions.
This is particularly useful for DSL AST processing that needs to emit new statements within
regions associated with expression nodes (e.g., using the ``Region`` context manager).
:param node: The AST node to process.
:type node: ast.AST
:return: The transformed AST node.
:rtype: ast.AST
"""
for field, old_value in ast.iter_fields(node):
if isinstance(old_value, list):
with Region(self.session_data, owning_node=node):
for value in old_value:
if isinstance(value, ast.AST):
value = self.visit(value)
if value is None:
continue
elif not isinstance(value, ast.AST):
node._new_value.extend(value)
continue
node._new_value.append(value)
old_value[:] = node._new_value
elif isinstance(old_value, ast.AST):
new_node = self.visit(old_value)
if new_node is None:
delattr(node, field)
else:
setattr(node, field, new_node)
return node
def __init__(self, client_module_name):
super().__init__()
# Persistent state
@@ -303,29 +431,14 @@ class DSLPreprocessor(ast.NodeTransformer):
self.module_cache = {}
self._session_data = None
def _start_session(self):
"""
Starts a new preprocessing session by initializing session data.
This method sets up a fresh SessionData instance for use during
AST transformations. It must be called before performing any
preprocessing actions that require access to context-specific
information during the transformation of a function's AST.
"""
self._session_data = SessionData()
def _end_session(self):
"""
Ends the current preprocessing session and clears session data.
This method resets the session-specific data, marking the end of
a preprocessing context. It should be called after all necessary
AST processing is complete to ensure no stale context remains.
"""
self._session_data = None
@contextlib.contextmanager
def get_session(self):
return DSLPreprocessorSession(dsl_object=self)
try:
self._session_data = SessionData()
yield self
finally:
self._session_data = None
@property
def session_data(self):
@@ -733,10 +846,13 @@ class DSLPreprocessor(ast.NodeTransformer):
if isinstance(node.test, ast.Call):
func = node.test.func
if isinstance(func, ast.Attribute) and func.attr == "const_expr":
if (
isinstance(func, ast.Attribute)
and func.attr in self.CONST_EXPR_NAME
):
return True
elif isinstance(func, ast.Name) and func.id == "const_expr":
elif isinstance(func, ast.Name) and func.id in self.CONST_EXPR_NAME:
return True
return False
@@ -775,7 +891,10 @@ class DSLPreprocessor(ast.NodeTransformer):
return unified_tree
def analyze_region_variables(
self, node: Union[ast.For, ast.If, ast.While], active_symbols: List[Set[str]]
self,
node: Union[ast.For, ast.If, ast.While],
active_symbols: List[Set[str]],
active_callables: List[Set[str]],
):
"""
Analyze variables in different code regions to identify read-only, write-only,
@@ -785,8 +904,7 @@ class DSLPreprocessor(ast.NodeTransformer):
# we need orderedset to keep the insertion order the same. otherwise generated IR is different each time
write_args = OrderedSet()
invoked_args = OrderedSet()
local_closure = self.session_data.local_closures
called_closures = OrderedSet()
called_functions = OrderedSet()
class RegionAnalyzer(ast.NodeVisitor):
force_store = False
@@ -853,8 +971,7 @@ class DSLPreprocessor(ast.NodeTransformer):
if isinstance(node.func, ast.Name):
func_name = node.func.id
if func_name in local_closure:
called_closures.add(func_name)
called_functions.add(func_name)
# Classes are mutable by default. Mark them as write. If they are
# dataclass(frozen=True), treat them as read in runtime.
@@ -878,8 +995,8 @@ class DSLPreprocessor(ast.NodeTransformer):
write_args = list(write_args.intersections(active_symbols))
invoked_args = list(invoked_args.intersections(active_symbols))
return write_args + invoked_args, len(write_args), called_closures
called_functions = list(called_functions.intersections(active_callables))
return write_args + invoked_args, len(write_args), called_functions
def extract_range_args(self, iter_node):
args = iter_node.args
@@ -958,12 +1075,15 @@ class DSLPreprocessor(ast.NodeTransformer):
# Create the loop body
transformed_body = []
for stmt in node.body:
transformed_stmt = self.visit(stmt) # Recursively visit inner statements
if isinstance(transformed_stmt, list):
transformed_body.extend(transformed_stmt)
else:
transformed_body.append(transformed_stmt)
with Region(self.session_data, new_value=transformed_body):
for stmt in node.body:
transformed_stmt = self.visit(
stmt
) # Recursively visit inner statements
if isinstance(transformed_stmt, list):
transformed_body.extend(transformed_stmt)
else:
transformed_body.append(transformed_stmt)
# Handle the return for a single iterated argument correctly
if len(write_args) == 0:
@@ -1186,8 +1306,9 @@ class DSLPreprocessor(ast.NodeTransformer):
return node
active_symbols = self.session_data.scope_manager.get_active_symbols()
active_callables = self.session_data.scope_manager.get_active_callables()
with self.session_data.scope_manager:
with self.session_data.scope_manager.enter_control_flow_scope():
if isinstance(node.target, ast.Name):
self.session_data.scope_manager.add_to_scope(node.target.id)
@@ -1210,7 +1331,9 @@ class DSLPreprocessor(ast.NodeTransformer):
# Get toplevel module
check_call = self._insert_cf_symbol_check(node.iter.func.value)
new_for_node = self.transform_for_loop(node, active_symbols)
new_for_node = self.transform_for_loop(
node, active_symbols, active_callables
)
if check_call is not None:
new_for_node = [check_call] + new_for_node
@@ -1234,7 +1357,6 @@ class DSLPreprocessor(ast.NodeTransformer):
),
location,
)
self.generic_visit(node)
return node
def _handle_negative_step(self, node, start_expr, stop_expr, step_expr):
@@ -1309,11 +1431,12 @@ class DSLPreprocessor(ast.NodeTransformer):
location=node,
)
extra_exprs.append(isNegative)
extra_exprs.append(start)
extra_exprs.append(stop)
extra_exprs.append(step)
extra_exprs.append(offset)
with Region(self.session_data, new_value=extra_exprs):
extra_exprs.append(self.generic_visit(isNegative))
extra_exprs.append(self.generic_visit(start))
extra_exprs.append(self.generic_visit(stop))
extra_exprs.append(self.generic_visit(step))
extra_exprs.append(self.generic_visit(offset))
# Add this to begining of loop body
# for i in range(start, stop, step):
@@ -1360,7 +1483,7 @@ class DSLPreprocessor(ast.NodeTransformer):
)
)
def transform_for_loop(self, node, active_symbols):
def transform_for_loop(self, node, active_symbols, active_callables):
# Check for early exit and raise exception
self.check_early_exit(node, "for")
if node.orelse:
@@ -1409,7 +1532,7 @@ class DSLPreprocessor(ast.NodeTransformer):
prefetch_stages = self.extract_prefetch_stages_args(node.iter)
vectorize = self.extract_vectorize_args(node.iter)
write_args, full_write_args_count, called_closures = (
self.analyze_region_variables(node, active_symbols)
self.analyze_region_variables(node, active_symbols, active_callables)
)
if has_step and self.client_module_name[0] == "cutlass":
@@ -1678,10 +1801,8 @@ class DSLPreprocessor(ast.NodeTransformer):
return node
def visit_ClassDef(self, node):
self.session_data.class_name = node.name
self.generic_visit(node)
self.session_data.class_name = None
return node
with self.session_data.set_current_class_name(node.name):
return self.generic_visit(node)
def _visit_target(self, target):
if isinstance(target, ast.Name):
@@ -1798,13 +1919,13 @@ class DSLPreprocessor(ast.NodeTransformer):
return new_decorator_list
def visit_FunctionDef(self, node):
with self.session_data.scope_manager:
self.session_data.function_counter += 1
self.session_data.function_name = node.name
if self.session_data.function_depth > 0:
self.session_data.local_closures.add(node.name)
# Add self to active symbols of parent scope
self.session_data.scope_manager.add_to_callables(node.name)
self.session_data.function_depth += 1
with self.session_data.scope_manager.enter_local_scope(), self.session_data.set_current_function_name(
node.name
):
self.session_data.function_counter += 1
# Add function name and arguments
self.session_data.scope_manager.add_to_scope(node.name)
@@ -1822,7 +1943,6 @@ class DSLPreprocessor(ast.NodeTransformer):
self.generic_visit(node)
self.session_data.function_depth -= 1
# Remove .jit and .kernel decorators
node.decorator_list = self.remove_dsl_decorator(node.decorator_list)
@@ -1832,13 +1952,10 @@ class DSLPreprocessor(ast.NodeTransformer):
return node
def visit_With(self, node):
with self.session_data.scope_manager:
for item in node.items:
if isinstance(item.optional_vars, ast.Name):
self.session_data.scope_manager.add_to_scope(item.optional_vars.id)
self.generic_visit(node)
return node
for item in node.items:
if isinstance(item.optional_vars, ast.Name):
self.session_data.scope_manager.add_to_scope(item.optional_vars.id)
return self.generic_visit(node)
def visit_While(self, node):
# Constexpr doesn't get preprocessed
@@ -1848,12 +1965,14 @@ class DSLPreprocessor(ast.NodeTransformer):
return [check, node]
active_symbols = self.session_data.scope_manager.get_active_symbols()
with self.session_data.scope_manager:
active_callables = self.session_data.scope_manager.get_active_callables()
with self.session_data.scope_manager.enter_control_flow_scope():
# Check for early exit and raise exception
self.check_early_exit(node, "while")
write_args, full_write_args_count, called_closures = (
self.analyze_region_variables(node, active_symbols)
self.analyze_region_variables(node, active_symbols, active_callables)
)
exprs = []
if called_closures:
@@ -1869,18 +1988,6 @@ class DSLPreprocessor(ast.NodeTransformer):
return exprs + [func_def] + assign
def visit_Try(self, node):
with self.session_data.scope_manager:
self.generic_visit(node)
return node
def visit_ExceptHandler(self, node):
with self.session_data.scope_manager:
if node.name: # Exception variable
self.session_data.scope_manager.add_to_scope(node.name)
self.generic_visit(node)
return node
def create_cf_call(self, func_name, yield_args, node):
"""Creates the assignment statement for the if function call"""
if not yield_args:
@@ -1928,42 +2035,142 @@ class DSLPreprocessor(ast.NodeTransformer):
else:
return [ast.copy_location(assign, node)]
def _visit_Comprehension(self, node, ele_visitor):
node.generators = [self.visit(generator) for generator in node.generators]
targets = []
class NameCollector(ast.NodeVisitor):
def visit_Name(self, node):
if isinstance(node.ctx, ast.Store):
targets.append(node.id)
# Collect generator targets
collector = NameCollector()
[collector.visit(generator) for generator in node.generators]
self.session_data.generator_targets = targets
ele_visitor(node)
self.session_data.generator_targets = []
return node
def visit_DictComp(self, node):
def key_value_visitor(n):
n.key = self.visit(n.key)
n.value = self.visit(n.value)
return self._visit_Comprehension(node, key_value_visitor)
def visit_ListComp(self, node):
return self._visit_Comprehension(
node, lambda n: setattr(n, "elt", self.visit(n.elt))
)
def visit_GeneratorExp(self, node):
return self._visit_Comprehension(
node, lambda n: setattr(n, "elt", self.visit(n.elt))
)
def visit_SetComp(self, node):
return self._visit_Comprehension(
node, lambda n: setattr(n, "elt", self.visit(n.elt))
)
def visit_IfExp(self, node):
"""
Visits an inline if-else expression (ternary operator).
This is the Python equivalent of `x if condition else y`.
Transforms an inline if-else (ternary) expression into runtime-dispatched
control flow using synthesized function definitions for each branch.
This converts an expression of the form ``x if cond else y`` into two local
function blocks (for the ``then`` and ``else`` branches), inserts those blocks
just before the current statement, and produces a call to the conditional executor.
This lets the DSL infrastructure analyze and dispatch dynamic inline conditionals
in a uniform way at runtime.
Parameters
----------
node : ast.IfExp
The AST node representing the inline if-else expression.
Returns
-------
ast.Call
An AST node that calls the conditional expression executor, referencing
the synthesized blocks and the predicate.
"""
self.generic_visit(node)
# Emit
# node if type(pred) == bool else select_(pred, body, orelse)
# so if pred is a python bool, use python to short-circuit and avoid emit arith.select
self.session_data.import_top_module = True
return ast.copy_location(
ast.IfExp(
test=ast.Compare(
left=ast.Call(
func=ast.Name(id="type", ctx=ast.Load()),
args=[node.test],
keywords=[],
),
ops=[ast.Eq()],
comparators=[ast.Name(id="bool", ctx=ast.Load())],
),
body=node, # Original ternary expression
orelse=ast.Call(
func=_create_module_attribute(
"select_", use_base_dsl=False, submodule_name=None
),
args=[
node.test,
node.body,
node.orelse,
],
keywords=[],
),
# Create unique names for the then and else branch function blocks
then_block_name = f"ifexp_then_block_{self.session_data.counter}"
else_block_name = f"ifexp_else_block_{self.session_data.counter}"
self.session_data.counter += 1
# Define the then-block function, with no arguments and returning the visited body
then_block_def = ast.FunctionDef(
name=then_block_name,
args=ast.arguments(
posonlyargs=[],
args=[
ast.arg(arg=target, annotation=None)
for target in self.session_data.generator_targets
],
kwonlyargs=[],
kw_defaults=[],
defaults=[],
),
node,
body=[ast.Return(value=self.visit(node.body))],
decorator_list=[],
)
# Define the else-block function, with no arguments and returning the visited orelse
else_block_def = ast.FunctionDef(
name=else_block_name,
args=ast.arguments(
posonlyargs=[],
args=[
ast.arg(arg=target, annotation=None)
for target in self.session_data.generator_targets
],
kwonlyargs=[],
kw_defaults=[],
defaults=[],
),
body=[ast.Return(value=self.visit(node.orelse))],
decorator_list=[],
)
# Insert the block definitions into the most recent (innermost) region before the statement
self.session_data.region_stack[-1].append_new_stmts(
[then_block_def, else_block_def]
)
# Create the executor call node, wiring up the predicate and newly synthesized blocks
executor_call = ast.Call(
func=_create_module_attribute(self.IFEXP_EXECUTOR),
args=[],
keywords=[
ast.keyword(arg="pred", value=self.visit(node.test)),
ast.keyword(
arg="generator_targets",
value=ast.Tuple(
elts=[
ast.Name(id=name, ctx=ast.Load())
for name in self.session_data.generator_targets
],
ctx=ast.Load(),
),
),
ast.keyword(
arg="then_block", value=ast.Name(id=then_block_name, ctx=ast.Load())
),
ast.keyword(
arg="else_block", value=ast.Name(id=else_block_name, ctx=ast.Load())
),
],
)
# Return the transformed executor call node at the original location in the AST
return ast.copy_location(executor_call, node)
cmpops = {
"Eq": "==",
@@ -2016,12 +2223,14 @@ class DSLPreprocessor(ast.NodeTransformer):
return [check, node]
active_symbols = self.session_data.scope_manager.get_active_symbols()
with self.session_data.scope_manager:
active_callables = self.session_data.scope_manager.get_active_callables()
with self.session_data.scope_manager.enter_control_flow_scope():
# Check for early exit and raise exception
self.check_early_exit(node, "if")
yield_args, full_write_args_count, called_closures = (
self.analyze_region_variables(node, active_symbols)
self.analyze_region_variables(node, active_symbols, active_callables)
)
exprs = []
if called_closures:
@@ -2060,12 +2269,18 @@ class DSLPreprocessor(ast.NodeTransformer):
func_args_then_else = [ast.arg(arg=var, annotation=None) for var in write_args]
then_body = []
for stmt in node.body:
transformed_stmt = self.visit(stmt) # Recursively visit inner statements
if isinstance(transformed_stmt, list):
then_body.extend(transformed_stmt)
else:
then_body.append(transformed_stmt)
with (
Region(self.session_data, new_value=then_body),
self.session_data.scope_manager.enter_control_flow_scope(),
):
for stmt in node.body:
transformed_stmt = self.visit(
stmt
) # Recursively visit inner statements
if isinstance(transformed_stmt, list):
then_body.extend(transformed_stmt)
else:
then_body.append(transformed_stmt)
# Create common return list for all blocks
return_list = ast.List(
@@ -2210,14 +2425,18 @@ class DSLPreprocessor(ast.NodeTransformer):
)
else:
else_body = []
for stmt in node.orelse:
transformed_stmt = self.visit(
stmt
) # Recursively visit inner statements
if isinstance(transformed_stmt, list):
else_body.extend(transformed_stmt)
else:
else_body.append(transformed_stmt)
with (
Region(self.session_data, new_value=else_body),
self.session_data.scope_manager.enter_control_flow_scope(),
):
for stmt in node.orelse:
transformed_stmt = self.visit(
stmt
) # Recursively visit inner statements
if isinstance(transformed_stmt, list):
else_body.extend(transformed_stmt)
else:
else_body.append(transformed_stmt)
# Regular else block
else_block = ast.FunctionDef(
@@ -2302,7 +2521,6 @@ class DSLPreprocessor(ast.NodeTransformer):
cond, write_args = while_before_block(write_args)
return write_args
"""
test_expr = self.visit(node.test)
# Section: decorator construction
decorator_keywords = [
@@ -2343,11 +2561,15 @@ class DSLPreprocessor(ast.NodeTransformer):
)
# Section: while_before_block FunctionDef, which contains condition
while_before_stmts = []
with Region(self.session_data, new_value=while_before_stmts):
test_expr = ast.copy_location(self.visit(node.test), node.test)
while_before_return_list = ast.List(
elts=[test_expr, yield_args_ast_name_list],
ctx=ast.Load(),
)
while_before_stmts = [ast.Return(value=while_before_return_list)]
while_before_stmts.append(ast.Return(value=while_before_return_list))
while_before_block = ast.copy_location(
ast.FunctionDef(
name=while_before_block_name,
@@ -2360,12 +2582,15 @@ class DSLPreprocessor(ast.NodeTransformer):
# Section: while_after_block FunctionDef, which contains loop body
while_after_stmts = []
for stmt in node.body:
transformed_stmt = self.visit(stmt) # Recursively visit inner statements
if isinstance(transformed_stmt, list):
while_after_stmts.extend(transformed_stmt)
else:
while_after_stmts.append(transformed_stmt)
with Region(self.session_data, new_value=while_after_stmts):
for stmt in node.body:
transformed_stmt = self.visit(
stmt
) # Recursively visit inner statements
if isinstance(transformed_stmt, list):
while_after_stmts.extend(transformed_stmt)
else:
while_after_stmts.append(transformed_stmt)
while_after_stmts.append(ast.Return(value=yield_args_ast_name_list))
while_after_block = ast.copy_location(
+133 -17
View File
@@ -10,7 +10,9 @@
# is strictly prohibited.
import os
from typing import Any, Dict, Iterable, Optional, Union, Sequence
from typing import Any, Dict, Optional, Union
from functools import total_ordering
from dataclasses import dataclass
"""
This module provides a Exception classes DSL class for any Dialect.
@@ -325,29 +327,143 @@ This error typically occurs when:
)
def _get_cuda_version() -> str:
# Client of this module should implement this function
"""
Placeholder for CUDA version query.
This function should be implemented by the client of this module.
When implemented, it must return the CUDA version as a string, e.g. "12.2".
Raises:
NotImplementedError: Always, unless overridden by the package initializer or client.
"""
raise NotImplementedError("_get_cuda_version is not implemented")
@total_ordering
@dataclass(frozen=True)
class DSLCudaVersion:
"""
Class to represent the CUDA version used to build the DSL.
"""
def __init__(self, version: str):
self.version_tuple = tuple(int(part) for part in version.split("."))
major: int
minor: int
def __str__(self):
return f"{self.major}.{self.minor}"
def __init__(self, version: str):
parts = version.split(".")
object.__setattr__(self, "major", int(parts[0]))
object.__setattr__(self, "minor", int(parts[1]))
def __eq__(self, other):
if isinstance(other, DSLCudaVersion):
return self.version_tuple == other.version_tuple
elif isinstance(other, str):
return self == DSLCudaVersion(other)
else:
return False
return self.major == other.major and self.minor == other.minor
@property
def major(self):
return self.version_tuple[0]
def __lt__(self, other):
return [self.major, self.minor] < [other.major, other.minor]
@property
def minor(self):
return self.version_tuple[1]
def _coerce_to_cuda_version(
value: Optional[Union[DSLCudaVersion, str]], param_name: str
) -> Optional[DSLCudaVersion]:
"""
Coerce a value to DSLCudaVersion.
:param value: The value to coerce (DSLCudaVersion, str, or None).
:param param_name: The parameter name for error messages.
:returns: DSLCudaVersion or None if value is None.
:raises DSLRuntimeError: If value is not a supported type.
"""
if value is None:
return None
if isinstance(value, DSLCudaVersion):
return value
if isinstance(value, str):
return DSLCudaVersion(value)
raise DSLRuntimeError(
f"{param_name} must be a DSLCudaVersion or str, got {type(value).__name__}"
)
def target_version(
exact_version: Optional[Union[DSLCudaVersion, str]] = None,
*,
min_version: Optional[Union[DSLCudaVersion, str]] = None,
max_version: Optional[Union[DSLCudaVersion, str]] = None,
) -> bool:
"""
Check if the current CUDA version used to build the DSL matches an exact version
or falls within specified bounds at compile-time.
Only one of ``exact_version`` *or* ``min_version``/``max_version`` may be specified.
At least one must be provided.
:param exact_version: The required CUDA version (e.g., "12.3").
:type exact_version: Optional[Union[DSLCudaVersion, str]]
:param min_version: The minimum CUDA version required (inclusive, e.g., "12.0").
:type min_version: Optional[Union[DSLCudaVersion, str]]
:param max_version: The maximum CUDA version allowed (inclusive, e.g., "13.2").
:type max_version: Optional[Union[DSLCudaVersion, str]]
:returns: ``True`` if the CUDA version matches the requirement(s) specified.
:rtype: bool
:raises DSLRuntimeError:
- If neither an ``exact_version`` nor version range is given.
- If both an exact version and a range are provided.
- If ``min_version`` > ``max_version``.
- If any version parameter is not a DSLCudaVersion or str.
**Examples**
.. code-block:: python
target_version(exact_version="12.3")
# True if CUDA_VERSION == 12.3
target_version(min_version="12.0")
# True if CUDA_VERSION >= 12.0
target_version(max_version="13.2")
# True if CUDA_VERSION <= 13.2
target_version(min_version="12.0", max_version="13.2")
# True if 12.0 <= CUDA_VERSION <= 13.2
"""
# Avoid circular dependency
from .version_info import CUDA_VERSION
# Coerce all version parameters to DSLCudaVersion at the start
exact_v = _coerce_to_cuda_version(exact_version, "exact_version")
min_v = _coerce_to_cuda_version(min_version, "min_version")
max_v = _coerce_to_cuda_version(max_version, "max_version")
# Sanity check
is_range_check = min_v is not None or max_v is not None
is_exact_version_check = exact_v is not None
if is_range_check and is_exact_version_check:
raise DSLRuntimeError(
"Cannot use exact_version and [min_version, max_version] check at the same time"
)
if is_range_check:
if min_v is None and max_v is None:
raise DSLRuntimeError(
"min_version and max_version cannot be None at the same time"
)
if min_v is not None and max_v is not None:
if min_v > max_v:
raise DSLRuntimeError("min_version must be less than max_version")
result = True
if min_v is not None:
result = result and CUDA_VERSION >= min_v
if max_v is not None:
result = result and CUDA_VERSION <= max_v
return result
elif is_exact_version_check:
return CUDA_VERSION == exact_v
else:
raise DSLRuntimeError(
"either exact_version, min_version, or max_version must be provided"
)
@@ -340,21 +340,6 @@ class LinkLibraries(StringCompileOption):
class GPUArch(StringCompileOption):
option_name = "cubin-chip"
def __init__(self, val):
if isinstance(val, str) and val.startswith("sm_110"):
val = val.replace("sm_110", "sm_101")
super().__init__(val)
@property
def value(self) -> bool:
return self._value
@value.setter
def value(self, value: bool):
if isinstance(value, str) and value.startswith("sm_110"):
value = value.replace("sm_110", "sm_101")
self._value = value
class EnableTVMFFI(EmptyCompileOption):
pass
+2 -6
View File
@@ -846,6 +846,7 @@ class BaseDSL(metaclass=DSLSingletonMeta):
use_pdl: bool = False
auto_smem: bool = False
cooperative: bool = False
@staticmethod
def _check_and_canonicalize_dim(dim, name):
if not isinstance(dim, (list, tuple)):
@@ -967,18 +968,13 @@ class BaseDSL(metaclass=DSLSingletonMeta):
sys.stderr = redirect_stderr = io.StringIO()
sys.stdout = redirect_stdout = io.StringIO()
compile_gpu_arch = (
self.envar.arch
if not self.compile_options.gpu_arch
else self.compile_options.gpu_arch
)
try:
kernel = self.compiler_provider.compile_and_jit(
module,
pipeline,
shared_libs=shared_libs,
cuda_toolkit=self.envar.cuda_toolkit,
arch=compile_gpu_arch,
arch=self.envar.arch,
)
finally:
@@ -125,8 +125,6 @@ def detect_gpu_arch(prefix):
suffix = ""
if major >= 9:
suffix = "a"
if major == 11 and minor == 0:
major, minor = 10, 1
return f"sm_{major}{minor}{suffix}"
@@ -367,8 +365,6 @@ class EnvironmentVarManager(LogEnvironmentManager):
# Other options
self.dryrun = get_bool_env_var(f"{prefix}_DRYRUN", False)
self.arch = get_str_env_var(f"{prefix}_ARCH", detect_gpu_arch(prefix))
if self.arch.startswith("sm_110"):
self.arch = self.arch.replace("sm_110", "sm_101")
self.warnings_as_errors = get_bool_env_var(
f"{prefix}_WARNINGS_AS_ERRORS", False
)
@@ -477,9 +477,7 @@ class MLIRBuilder(MLIRTypeBuilder):
)
func_op.attributes["llvm.linkage"] = ir.StringAttr.get("external")
def create_alloca(
self, entry_block: ir.Block, alloca_type: ir.Type, array_size: int
) -> ir.Value:
def create_alloca(self, entry_block: ir.Block, alloca_type: ir.Type, array_size: int) -> ir.Value:
"""Create an alloca operation."""
with ir.InsertionPoint(entry_block.operations[0]):
# declare the struct type
@@ -1277,7 +1277,10 @@ class TVMFFIFunctionBuilder(TVMFFIBuilder):
return cond
return self.check_condition(
current_block, check_value_mismatch, error_kind, error_msg_mismatch
current_block,
check_value_mismatch,
error_kind,
error_msg_mismatch,
)
def set_or_check_matched_var_binding_from_shape(
@@ -290,7 +290,7 @@ def set_dataclass_attributes(
for field, value in zip(fields, values):
setattr(instance, field, value)
return instance
return instance
def default_dataclass_from_iterable(
@@ -1,4 +1,4 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-FileCopyrightText: Copyright (c) 2025 - 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: LicenseRef-NvidiaProprietary
#
# Use of this software is governed by the terms and conditions of the
@@ -9,14 +9,15 @@
# and related documentation outside the scope permitted by the EULA
# is strictly prohibited.
from ..cutlass_dsl import DSLCudaVersion, DSLRuntimeError
from typing import Callable
from .common import DSLCudaVersion, DSLRuntimeError, _get_cuda_version
try:
from .._mlir._mlir_libs._cutlass_ir._base_dsl import get_cuda_version
CUDA_VERSION = DSLCudaVersion(get_cuda_version())
CUDA_VERSION = DSLCudaVersion(_get_cuda_version())
except Exception as e:
raise DSLRuntimeError(
"💥💥💥 Failed to get CUDA version 💥💥💥",
cause=e,
suggestion="Consider re-installing the package."
suggestion="Consider re-installing the package.",
) from e
+53 -131
View File
@@ -181,19 +181,21 @@ from .atom import (
make_tiled_copy_C_atom,
make_cotiled_copy,
copy_atom_call,
mma_atom_call,
)
from .algorithm import gemm, copy, basic_copy, basic_copy_if, autovec_copy, prefetch
from . import typing as typing_module
from . import core
from . import arch
from . import export
from . import nvgpu
from . import testing
from . import runtime
from . import math
# Export all math ops without "math."
from .math import *
@@ -212,7 +214,6 @@ GenerateLineInfo = _dsl.GenerateLineInfo
KeepCUBIN = _dsl.KeepCUBIN
KeepPTX = _dsl.KeepPTX
GPUArch = _dsl.GPUArch
LinkLibraries = _dsl.LinkLibraries
EnableTVMFFI = _dsl.EnableTVMFFI
# attach the TVM FFI ABI interface postprocessor to the DSL
@@ -222,16 +223,52 @@ _tvm_ffi_args_spec_converter.attach_args_spec_converter(_dsl.CuTeDSL._get_dsl())
# Explicitly export all symbols for documentation generation
__all__ = [
# Core types
*core.__all__,
# ==================== cutlass._mlir.dialects.cute ====================
"AddressSpace",
"CacheEvictionPriority",
# ==================== .typing ====================
"Tensor",
"Layout",
"ComposedLayout",
"Swizzle",
"E",
"ScaledBasis",
"SymInt",
"is_integer",
"is_int_tuple",
# ==================== .core ====================
*core.__all__,
# ==================== .tuple ====================
"transform_leaf",
"find_if",
"find",
"flatten_to_tuple",
"unflatten",
"product",
"product_like",
"product_each",
"elem_less",
"tuple_cat",
"transform_apply",
"filter_tuple",
# ==================== .tensor ====================
"TensorSSA",
"ReductionOp",
"make_tensor",
"make_identity_tensor",
"make_fragment",
"make_fragment_like",
"make_rmem_tensor_like",
"make_rmem_tensor",
"recast_tensor",
"domain_offset",
"print_tensor",
"full",
"full_like",
"empty_like",
"ones_like",
"zeros_like",
"where",
"any_",
"all_",
# ==================== .atom ====================
"Atom",
"MmaAtom",
"CopyAtom",
@@ -239,106 +276,6 @@ __all__ = [
"TiledMma",
"ThrMma",
"ThrCopy",
"TensorSSA",
"ReductionOp",
"SymInt",
# Basic utility functions
"assume",
"is_integer",
"is_int_tuple",
"is_static",
"has_underscore",
"shape",
"printf",
"print_tensor",
"pretty_str",
# Layout functions
"make_layout",
"recast_layout",
"make_identity_layout",
"make_ordered_layout",
"make_layout_like",
"make_composed_layout",
"make_layout_tv",
"make_layout_image_mask",
"get_nonswizzle_portion",
"get_swizzle_portion",
# Tensor functions
"make_ptr",
"make_tensor",
"make_identity_tensor",
"make_fragment",
"make_fragment_like",
"make_rmem_tensor",
"make_rmem_tensor_like",
"recast_ptr",
"recast_tensor",
# Tensor manipulation
"get",
"select",
"front",
"is_major",
"leading_dim",
"find",
"find_if",
"transform_leaf",
"basis_value",
"basis_get",
"coalesce",
"group_modes",
"cosize",
"size_in_bytes",
# Tuple operations
"flatten_to_tuple",
"flatten",
"unflatten",
"product",
"product_like",
"product_each",
"prepend",
"append",
"prepend_ones",
"append_ones",
"elem_less",
"tuple_cat",
"transform_apply",
"filter_tuple",
# Math operations
"ceil_div",
"round_up",
# Layout operations
"slice_and_offset",
"crd2idx",
"domain_offset",
"filter_zeros",
"filter",
"tile_to_shape",
"shape_div",
"dice",
# Layout algebra
"composition",
"complement",
"right_inverse",
"left_inverse",
"max_common_layout",
"max_common_vector",
"is_congruent",
"is_weakly_congruent",
# Product operations
"logical_product",
"zipped_product",
"tiled_product",
"flat_product",
"raked_product",
"blocked_product",
# Division operations
"flat_divide",
"logical_divide",
"zipped_divide",
"tiled_divide",
"local_partition",
"local_tile",
# MMA and Copy atom operations
"make_atom",
"make_mma_atom",
"make_tiled_mma",
@@ -353,39 +290,24 @@ __all__ = [
"make_tiled_copy_C_atom",
"make_cotiled_copy",
"copy_atom_call",
# Algorithm operations
"mma_atom_call",
# ==================== .algorithm ====================
"gemm",
"copy",
"basic_copy",
"basic_copy_if",
"autovec_copy",
"copy",
"prefetch",
"gemm",
# Tensor creation
"full",
"full_like",
"empty_like",
"ones_like",
"zeros_like",
"where",
"any_",
"all_",
"repeat_as_tuple",
"repeat",
"repeat_like",
# User defined struct
"struct",
# FastDivmod operations
"FastDivmodDivisor",
"fast_divmod_create_divisor",
# Modules
# ==================== .extension ====================
# ==================== .math ====================
*math.__all__,
# ==================== submodules ====================
"arch",
"export",
"nvgpu",
"testing",
"runtime",
# Math utils
*math.__all__,
# Decorators and code generation
# ==================== DSL (cutlass_dsl) ====================
"jit",
"kernel",
"register_jit_arg_adapter",
+40 -12
View File
@@ -10,7 +10,7 @@
# is strictly prohibited.
import math
from typing import Optional, Dict, Any, List, Tuple
from typing import Optional, Dict, Any, List, Tuple, Union
from cutlass._mlir import ir
from cutlass.cutlass_dsl import for_generate, yield_out, if_generate, dsl_user_op
@@ -32,12 +32,26 @@ from .core import (
from .atom import MmaAtom, CopyAtom, make_atom
def _normalize_gemm_operand_list(
x: Union["Tensor", List["Tensor"], Tuple["Tensor", ...]], name: str
) -> List["Tensor"]:
if isinstance(x, Tensor):
return [x]
if isinstance(x, (list, tuple)):
if len(x) == 0:
raise ValueError(f"`{name}` must contain at least one Tensor")
if not all(isinstance(t, Tensor) for t in x):
raise TypeError(f"All elements of `{name}` must be Tensor")
return list(x) # type: ignore
raise TypeError(f"`{name}` must be a Tensor or a sequence of Tensors")
@dsl_user_op
def gemm(
atom: MmaAtom,
d: Tensor,
a: Tensor,
b: Tensor,
a: Union[Tensor, List[Tensor], Tuple[Tensor, ...]],
b: Union[Tensor, List[Tensor], Tuple[Tensor, ...]],
c: Tensor,
*,
loc=None,
@@ -62,14 +76,17 @@ def gemm(
- Dispatch [4]: (V,M) x (V,N) => (V,M,N) => (V,M,1) x (V,N,1) => (V,M,N)
- Dispatch [5]: (V,M,K) x (V,N,K) => (V,M,N)
Operand flexibility:
- `a` and `b` can be a single Tensor (regular GEMM) or a sequence `[operand, scale_factor]` for block-scaled GEMM.
:param atom: MMA atom
:type atom: MmaAtom
:param d: Destination tensor
:type d: Tensor
:param a: First source tensor
:type a: Tensor
:param b: Second source tensor
:type b: Tensor
:param a: First source tensor or sequence for advanced modes (e.g., `[a, sfa]`)
:type a: Union[Tensor, List[Tensor], Tuple[Tensor, ...]]
:param b: Second source tensor or sequence for advanced modes (e.g., `[b, sfb]`)
:type b: Union[Tensor, List[Tensor], Tuple[Tensor, ...]]
:param c: Third source tensor
:type c: Tensor
:param loc: Source location for MLIR, defaults to None
@@ -82,8 +99,13 @@ def gemm(
:rtype: None
"""
a_rank = rank(a.shape)
b_rank = rank(b.shape)
# Normalize A/B to lists for variadic IR operands, while keeping old API working.
a_list = _normalize_gemm_operand_list(a, "a")
b_list = _normalize_gemm_operand_list(b, "b")
# Rank validations based on the primary A/B tensors (guaranteed non-empty)
a_rank = rank(a_list[0].shape)
b_rank = rank(b_list[0].shape)
c_rank = rank(c.shape)
d_rank = rank(d.shape)
@@ -104,7 +126,9 @@ def gemm(
raise ValueError("`c` must have rank 3 when `a` has rank 3")
value = atom._unpack(loc=loc, ip=ip, **kwargs)
return _cute_ir.gemm(value, d.value, a.value, b.value, c.value, loc=loc, ip=ip)
a_vals = [t.value for t in a_list]
b_vals = [t.value for t in b_list]
return _cute_ir.gemm(value, d.value, a_vals, b_vals, c.value, loc=loc, ip=ip)
@dsl_user_op
@@ -258,19 +282,21 @@ def _parse_auto_multicast_args(
This function consumes the following key from kwargs if present:
- 'auto_multicast': dict
dict: { 'multicast_layout': str, 'use_2cta': bool }
dict: { 'multicast_layout': str, 'use_2cta': bool, 'from_block_api': bool }
Returns:
List of (attr_name, ir.Attribute) pairs to be attached to the op.
Recognized attributes:
- ('multicast_layout', #cute.layout<...>) when a layout string is provided
- ('use_2cta', unit) when use_2cta is True
- ('from_block_api', unit) when from_block_api is True
"""
attr_pairs: List[Tuple[str, ir.Attribute]] = []
# Pop known keys to avoid leaking to trait unpack
auto_multicast = kwargs.pop("auto_multicast", None)
from_block_api: bool = False
use_2cta: bool = False
layout_str: Optional[str] = None
@@ -281,6 +307,7 @@ def _parse_auto_multicast_args(
)
layout_str = auto_multicast.get("multicast_layout", None)
use_2cta = bool(auto_multicast.get("use_2cta", False))
from_block_api = bool(auto_multicast.get("from_block_api", False))
if layout_str is not None:
if not isinstance(layout_str, str):
@@ -293,7 +320,8 @@ def _parse_auto_multicast_args(
ir.Attribute.parse(f'#cute.layout<"{layout_str}">'),
)
)
if from_block_api:
attr_pairs.append(("from_block_api", ir.UnitAttr.get()))
if use_2cta:
attr_pairs.append(("use_2cta", ir.UnitAttr.get()))
+9 -5
View File
@@ -11,7 +11,6 @@
from .elect import *
from .mbar import *
from .numeric_conversion import *
from .nvvm_wrappers import *
from .smem import *
from .tmem import *
@@ -74,6 +73,8 @@ __all__ = [
"vote_any_sync",
"vote_all_sync",
"vote_uni_sync",
"warp_redux_sync",
"atomic_max_float32",
"atomic_add",
"atomic_and",
"atomic_or",
@@ -98,12 +99,15 @@ __all__ = [
"fmax",
"rcp_approx",
"exp2",
"cvt_i8x4_to_f32x4",
"cvt_i8x2_to_f32x2",
"cvt_i8_bf16",
"cvt_i8x2_to_bf16x2",
"cvt_i8x4_to_bf16x4",
"cvt_f32x2_bf16x2",
"warp_redux_sync",
# Constants
"WARP_SIZE",
# Forward from auto-generated nvvm python
"ProxyKind",
"SharedSpace",
"RoundingModeKind",
#
# smem.py
#
+4 -5
View File
@@ -10,8 +10,11 @@
# is strictly prohibited.
from typing import Tuple
from cutlass.cutlass_dsl import T, dsl_user_op
from cutlass._mlir.dialects import nvvm, vector
from cutlass._mlir import ir
from cutlass._mlir.dialects import nvvm, llvm, vector, arith
from ..typing import Int32, Pointer, Int128
@@ -78,7 +81,6 @@ def clc_response(
)
# Query if the cluster was canceled
pred = nvvm.clusterlaunchcontrol_query_cancel_is_canceled(
T.bool(),
clc_result_i128,
loc=loc,
ip=ip,
@@ -87,7 +89,6 @@ def clc_response(
# Get first CTA ID x component
m_idx_i32 = nvvm.clusterlaunchcontrol_query_cancel_get_first_ctaid_x(
T.i32(),
clc_result_i128,
loc=loc,
ip=ip,
@@ -95,7 +96,6 @@ def clc_response(
# Get first CTA ID y component
n_idx_i32 = nvvm.clusterlaunchcontrol_query_cancel_get_first_ctaid_y(
T.i32(),
clc_result_i128,
loc=loc,
ip=ip,
@@ -103,7 +103,6 @@ def clc_response(
# Get first CTA ID z component
l_idx_i32 = nvvm.clusterlaunchcontrol_query_cancel_get_first_ctaid_z(
T.i32(),
clc_result_i128,
loc=loc,
ip=ip,
+1 -2
View File
@@ -9,7 +9,6 @@
# and related documentation outside the scope permitted by the EULA
# is strictly prohibited.
from cutlass.base_dsl.arch import Arch
from cutlass.cutlass_dsl import BaseDSL, T, dsl_user_op
import cutlass._mlir.dialects.cute_nvgpu as _cute_nvgpu_ir
@@ -72,6 +71,6 @@ def elect_one(*, loc=None, ip=None) -> IfOpRegion:
from cutlass.base_dsl.arch import Arch
BaseDSL._get_dsl().check_arch(lambda arch: arch >= Arch.sm_90)
is_thread_leader = nvvm.elect_sync(T.bool())
is_thread_leader = nvvm.elect_sync()
if_op = scf.IfOp(is_thread_leader, loc=loc, ip=ip)
return IfOpRegion(if_op.then_block, loc=loc, ip=ip)
+28 -18
View File
@@ -13,7 +13,7 @@ from typing import Optional
from cutlass.base_dsl.arch import Arch
from cutlass.cutlass_dsl import BaseDSL, T, if_generate, dsl_user_op
from cutlass._mlir.dialects import nvvm
from cutlass._mlir.dialects import nvvm, llvm
from ..typing import Pointer, Int, Boolean, Int32, AddressSpace
@@ -35,10 +35,7 @@ def mbarrier_init(mbar_ptr: Pointer, cnt: Int, *, loc=None, ip=None) -> None:
:type cnt: Int
"""
nvvm.mbarrier_init_shared(
mbar_ptr.to_llvm_ptr(loc=loc, ip=ip),
Int32(cnt).ir_value(loc=loc, ip=ip),
loc=loc,
ip=ip,
mbar_ptr.llvm_ptr, Int32(cnt).ir_value(loc=loc, ip=ip), loc=loc, ip=ip
)
@@ -68,15 +65,18 @@ def mbarrier_arrive_and_expect_tx(
"""
BaseDSL._get_dsl().check_arch(lambda arch: arch >= Arch.sm_90)
mbar_llvm_ptr = mbar_ptr.to_llvm_ptr(loc=loc, ip=ip)
mbar_llvm_ptr = mbar_ptr.llvm_ptr
if peer_cta_rank_in_cluster is not None:
mbar_llvm_ptr = nvvm.mapa_shared_cluster(
mbar_llvm_ptr.type,
mbar_cluster_type = llvm.PointerType.get(AddressSpace.dsmem)
mbar_llvm_ptr = nvvm.mapa(
mbar_cluster_type,
mbar_llvm_ptr,
Int32(peer_cta_rank_in_cluster).ir_value(loc=loc, ip=ip),
loc=loc,
ip=ip,
)
mbar_shared_type = llvm.PointerType.get(AddressSpace.smem)
mbar_llvm_ptr = llvm.addrspacecast(mbar_shared_type, mbar_llvm_ptr)
space = nvvm.MBarrierSpaceKind.CLUSTER
else:
space = nvvm.MBarrierSpaceKind.CTA
@@ -108,15 +108,18 @@ def mbarrier_expect_tx(
"""
BaseDSL._get_dsl().check_arch(lambda arch: arch >= Arch.sm_90)
mbar_llvm_ptr = mbar_ptr.to_llvm_ptr(loc=loc, ip=ip)
mbar_llvm_ptr = mbar_ptr.llvm_ptr
if peer_cta_rank_in_cluster is not None:
mbar_cluster_type = llvm.PointerType.get(AddressSpace.dsmem)
mbar_llvm_ptr = nvvm.mapa(
mbar_llvm_ptr.type,
mbar_cluster_type,
mbar_llvm_ptr,
Int32(peer_cta_rank_in_cluster).ir_value(loc=loc, ip=ip),
loc=loc,
ip=ip,
)
mbar_shared_type = llvm.PointerType.get(AddressSpace.smem)
mbar_llvm_ptr = llvm.addrspacecast(mbar_shared_type, mbar_llvm_ptr)
space = nvvm.MBarrierSpaceKind.CLUSTER
else:
space = nvvm.MBarrierSpaceKind.CTA
@@ -147,7 +150,7 @@ def mbarrier_wait(mbar_ptr: Pointer, phase: Int, *, loc=None, ip=None) -> None:
# This NVVM Op is a spin-loop wrapping the mbarrier.try_wait.parity.shared.b64 PTX
# The timeout in ns only applies to the latter and this call is truly blocking
nvvm.mbarrier_try_wait_parity_shared(
mbar_ptr.to_llvm_ptr(loc=loc, ip=ip),
mbar_ptr.llvm_ptr,
Int32(phase).ir_value(loc=loc, ip=ip),
Int32(timeout_ns).ir_value(loc=loc, ip=ip),
loc=loc,
@@ -171,8 +174,7 @@ def mbarrier_try_wait(mbar_ptr: Pointer, phase: Int, *, loc=None, ip=None) -> Bo
return Boolean(
nvvm.mbarrier_wait_parity(
T.bool(),
mbar_ptr.to_llvm_ptr(loc=loc, ip=ip),
mbar_ptr.llvm_ptr,
Int32(phase).ir_value(loc=loc, ip=ip),
nvvm.MBarrierWaitKind.TRY,
loc=loc,
@@ -226,17 +228,20 @@ def mbarrier_arrive(
the mbarrier is converted to a remote address in the peer CTA's
SMEM.
"""
mbar_llvm_ptr = mbar_ptr.to_llvm_ptr(loc=loc, ip=ip)
mbar_llvm_ptr = mbar_ptr.llvm_ptr
if peer_cta_rank_in_cluster is not None:
BaseDSL._get_dsl().check_arch(lambda arch: arch >= Arch.sm_90)
mbar_llvm_ptr = nvvm.mapa_shared_cluster(
mbar_llvm_ptr.type,
mbar_cluster_type = llvm.PointerType.get(AddressSpace.dsmem)
mbar_llvm_ptr = nvvm.mapa(
mbar_cluster_type,
mbar_llvm_ptr,
Int32(peer_cta_rank_in_cluster).ir_value(loc=loc, ip=ip),
loc=loc,
ip=ip,
)
mbar_shared_type = llvm.PointerType.get(AddressSpace.smem)
mbar_llvm_ptr = llvm.addrspacecast(mbar_shared_type, mbar_llvm_ptr)
space = nvvm.MBarrierSpaceKind.CLUSTER
else:
space = nvvm.MBarrierSpaceKind.CTA
@@ -264,5 +269,10 @@ def cp_async_mbarrier_arrive_noinc(mbar_ptr: Pointer, *, loc=None, ip=None) -> N
"""
BaseDSL._get_dsl().check_arch(lambda arch: arch >= Arch.sm_90)
mbar_llvm_ptr = mbar_ptr.to_llvm_ptr(loc=loc, ip=ip)
nvvm.cp_async_mbarrier_arrive_shared(mbar_llvm_ptr, noinc=True, loc=loc, ip=ip)
mbar_llvm_ptr = mbar_ptr.llvm_ptr
nvvm.cp_async_mbarrier_arrive_shared(
mbar_llvm_ptr,
noinc=True,
loc=loc,
ip=ip,
)
@@ -9,16 +9,17 @@
# and related documentation outside the scope permitted by the EULA
# is strictly prohibited.
from cutlass.base_dsl.arch import Arch
from cutlass.base_dsl.common import DSLRuntimeError
from cutlass.cutlass_dsl import BaseDSL, dsl_user_op
from cutlass._mlir import ir
from cutlass._mlir.dialects import builtin, arith, llvm, vector
from cutlass._mlir.dialects import arith, llvm, vector
from .nvvm_wrappers import (
cvt_i8_bf16,
cvt_i8x2_to_bf16x2,
cvt_i8x4_to_bf16x4,
cvt_f32x2_bf16x2,
cvt_i8x4_to_f32x4,
cvt_i8x2_to_f32x2,
@@ -26,22 +27,11 @@ from .nvvm_wrappers import (
cvt_i4x4_to_bf16x4,
cvt_i4x2_to_bf16x2,
cvt_i4_bf16,
cvt_f4e2m1x8_to_f16x8,
cvt_f4e2m1x4_to_f16x4,
cvt_f4e2m1x2_to_f16x2,
cvt_f4e2m1_f16,
cvt_f32_bf16,
sext_unpacked_i4x4_to_i8x4,
)
from ..typing import Int4, Int8, Float32, BFloat16, Int32
from ..typing import (
Int4,
Int8,
Int32,
Float16,
Float32,
BFloat16,
Float32,
)
@dsl_user_op
def cvt_i8_bf16_intrinsic(vec_i8, length, *, loc=None, ip=None):
@@ -64,6 +54,7 @@ def cvt_i8_bf16_intrinsic(vec_i8, length, *, loc=None, ip=None):
vec_f32x2_type = ir.VectorType.get([2], Float32.mlir_type, loc=loc)
vec_dst_type = ir.VectorType.get([length], BFloat16.mlir_type, loc=loc)
vec_dst = llvm.mlir_zero(vec_dst_type, loc=loc, ip=ip)
arch = BaseDSL._get_dsl().get_arch_enum()
# try to use vectorized version
if length >= 4:
num_vec4 = length // 4
@@ -71,45 +62,66 @@ def cvt_i8_bf16_intrinsic(vec_i8, length, *, loc=None, ip=None):
vec_i8x4 = vector.extract_strided_slice(
vec_i8x4_type, vec_i8, [src_pos], [4], [1], loc=loc, ip=ip
)
vec_f32x4 = cvt_i8x4_to_f32x4(vec_i8x4, loc=loc, ip=ip)
vec_f32x2_lo = vector.extract_strided_slice(
vec_f32x2_type, vec_f32x4, [0], [2], [1], loc=loc, ip=ip
)
vec_f32x2_hi = vector.extract_strided_slice(
vec_f32x2_type, vec_f32x4, [2], [2], [1], loc=loc, ip=ip
)
vec_bf16x2_lo = cvt_f32x2_bf16x2(vec_f32x2_lo, loc=loc, ip=ip)
vec_bf16x2_hi = cvt_f32x2_bf16x2(vec_f32x2_hi, loc=loc, ip=ip)
vec_dst = vector.insert_strided_slice(
vec_bf16x2_lo, vec_dst, [src_pos], [1], loc=loc, ip=ip
)
vec_dst = vector.insert_strided_slice(
vec_bf16x2_hi, vec_dst, [src_pos + 2], [1], loc=loc, ip=ip
)
if arch in cvt_i8_bf16_intrinsic.s26_bf16_supported_archs:
vec_bf16x4 = cvt_i8x4_to_bf16x4(vec_i8x4, loc=loc, ip=ip)
vec_dst = vector.insert_strided_slice(
vec_bf16x4, vec_dst, [src_pos], [1], loc=loc, ip=ip
)
else:
vec_f32x4 = cvt_i8x4_to_f32x4(vec_i8x4, loc=loc, ip=ip)
vec_f32x2_lo = vector.extract_strided_slice(
vec_f32x2_type, vec_f32x4, [0], [2], [1], loc=loc, ip=ip
)
vec_f32x2_hi = vector.extract_strided_slice(
vec_f32x2_type, vec_f32x4, [2], [2], [1], loc=loc, ip=ip
)
vec_bf16x2_lo = cvt_f32x2_bf16x2(vec_f32x2_lo, loc=loc, ip=ip)
vec_bf16x2_hi = cvt_f32x2_bf16x2(vec_f32x2_hi, loc=loc, ip=ip)
vec_dst = vector.insert_strided_slice(
vec_bf16x2_lo, vec_dst, [src_pos], [1], loc=loc, ip=ip
)
vec_dst = vector.insert_strided_slice(
vec_bf16x2_hi, vec_dst, [src_pos + 2], [1], loc=loc, ip=ip
)
src_pos += 4
length -= 4
if length >= 2:
vec_i8x2 = vector.extract_strided_slice(
vec_i8x2_type, vec_i8, [src_pos], [2], [1], loc=loc, ip=ip
)
vec_f32x2 = cvt_i8x2_to_f32x2(vec_i8x2, loc=loc, ip=ip)
vec_bf16x2 = cvt_f32x2_bf16x2(vec_f32x2, loc=loc, ip=ip)
if arch in cvt_i8_bf16_intrinsic.s26_bf16_supported_archs:
vec_bf16x2 = cvt_i8x2_to_bf16x2(vec_i8x2, loc=loc, ip=ip)
else:
vec_f32x2 = cvt_i8x2_to_f32x2(vec_i8x2, loc=loc, ip=ip)
vec_bf16x2 = cvt_f32x2_bf16x2(vec_f32x2, loc=loc, ip=ip)
vec_dst = vector.insert_strided_slice(
vec_bf16x2, vec_dst, [src_pos], [1], loc=loc, ip=ip
)
src_pos += 2
length -= 2
if length >= 1:
val_bf16 = cvt_i8_bf16(
vector.extractelement(
if arch in cvt_i8_bf16_intrinsic.s26_bf16_supported_archs:
val_bf16 = cvt_i8_bf16(
vector.extractelement(
vec_i8,
position=arith.constant(Int32.mlir_type, src_pos),
loc=loc,
ip=ip,
),
loc=loc,
ip=ip,
)
else:
src_i8 = vector.extractelement(
vec_i8,
position=arith.constant(Int32.mlir_type, src_pos),
loc=loc,
ip=ip,
),
loc=loc,
ip=ip,
)
)
src_i32 = llvm.sext(Int32.mlir_type, src_i8, loc=loc, ip=ip)
src_f32 = llvm.sitofp(Float32.mlir_type, src_i32, loc=loc, ip=ip)
val_bf16 = cvt_f32_bf16(src_f32, loc=loc, ip=ip)
vec_dst = vector.insertelement(
val_bf16,
vec_dst,
@@ -121,7 +133,7 @@ def cvt_i8_bf16_intrinsic(vec_i8, length, *, loc=None, ip=None):
@dsl_user_op
def cvt_i4_bf16_intrinsic(vec_i4, length, *, loc=None, ip=None):
def cvt_i4_bf16_intrinsic(vec_i4, length, *, with_shuffle=False, loc=None, ip=None):
"""
Fast conversion from int4 to bfloat16. It converts a vector of int4 to a vector of bfloat16.
@@ -129,6 +141,13 @@ def cvt_i4_bf16_intrinsic(vec_i4, length, *, loc=None, ip=None):
:type vec_i4: 1D vector of int4
:param length: The length of the input vector.
:type length: int
:param with_shuffle: Whether the input vec_i4 follows a specific shuffle pattern.
If True, for consecutive 8 int4 values with indices of (0, 1, 2, 3, 4, 5, 6, 7),
the input elements are shuffled to (0, 2, 1, 3, 4, 6, 5, 7). For tailing elements less than 8,
the shuffle pattern is (0, 2, 1, 3) for 4 elements. No shuffle is needed for less than 4 elements.
Shuffle could help to produce converted bf16 values in the natural order of (0, 1, 2 ,3 ,4 ,5 ,6 ,7)
without extra prmt instructions and thus better performance.
:type with_shuffle: bool
:return: The output 1D vector of bfloat16 with the same length as the input vector.
:rtype: 1D vector of bfloat16
"""
@@ -141,6 +160,7 @@ def cvt_i4_bf16_intrinsic(vec_i4, length, *, loc=None, ip=None):
vec_i4x2_type = ir.VectorType.get([2], Int4.mlir_type, loc=loc)
vec_dst_type = ir.VectorType.get([length], BFloat16.mlir_type, loc=loc)
vec_dst = llvm.mlir_zero(vec_dst_type, loc=loc, ip=ip)
# try to use vectorized version
if length >= 8:
num_vec8 = length // 8
@@ -148,7 +168,9 @@ def cvt_i4_bf16_intrinsic(vec_i4, length, *, loc=None, ip=None):
vec_i4x8 = vector.extract_strided_slice(
vec_i4x8_type, vec_i4, [src_pos], [8], [1], loc=loc, ip=ip
)
vec_bf16x8 = cvt_i4x8_to_bf16x8(vec_i4x8, loc=loc, ip=ip)
vec_bf16x8 = cvt_i4x8_to_bf16x8(
vec_i4x8, with_shuffle=with_shuffle, loc=loc, ip=ip
)
vec_dst = vector.insert_strided_slice(
vec_bf16x8, vec_dst, [src_pos], [1], loc=loc, ip=ip
)
@@ -158,7 +180,9 @@ def cvt_i4_bf16_intrinsic(vec_i4, length, *, loc=None, ip=None):
vec_i4x4 = vector.extract_strided_slice(
vec_i4x4_type, vec_i4, [src_pos], [4], [1], loc=loc, ip=ip
)
vec_bf16x4 = cvt_i4x4_to_bf16x4(vec_i4x4, loc=loc, ip=ip)
vec_bf16x4 = cvt_i4x4_to_bf16x4(
vec_i4x4, with_shuffle=with_shuffle, loc=loc, ip=ip
)
vec_dst = vector.insert_strided_slice(
vec_bf16x4, vec_dst, [src_pos], [1], loc=loc, ip=ip
)
@@ -168,7 +192,9 @@ def cvt_i4_bf16_intrinsic(vec_i4, length, *, loc=None, ip=None):
vec_i4x2 = vector.extract_strided_slice(
vec_i4x2_type, vec_i4, [src_pos], [2], [1], loc=loc, ip=ip
)
vec_bf16x2 = cvt_i4x2_to_bf16x2(vec_i4x2, loc=loc, ip=ip)
vec_bf16x2 = cvt_i4x2_to_bf16x2(
vec_i4x2, with_shuffle=with_shuffle, loc=loc, ip=ip
)
vec_dst = vector.insert_strided_slice(
vec_bf16x2, vec_dst, [src_pos], [1], loc=loc, ip=ip
)
@@ -195,84 +221,6 @@ def cvt_i4_bf16_intrinsic(vec_i4, length, *, loc=None, ip=None):
return vec_dst
@dsl_user_op
def cvt_f4e2m1_f16_intrinsic(vec_f4e2m1, length, *, loc=None, ip=None):
"""
Convert a vector of float4e2m1 to a vector of float16.
:param vec_f4e2m1: The input vector of float4e2m1.
:type vec_f4e2m1: 1D vector of float4e2m1
:param length: The length of the input vector.
:type length: int
:return: The output 1D vector of float16 with the same length as the input vector.
:rtype: 1D vector of float16
"""
src_pos = 0
vec_src_i4 = builtin.unrealized_conversion_cast(
[ir.VectorType.get([length], Int4.mlir_type, loc=loc)],
[vec_f4e2m1],
loc=loc,
ip=ip,
)
vec_i4x8_type = ir.VectorType.get([8], Int4.mlir_type, loc=loc)
vec_i4x4_type = ir.VectorType.get([4], Int4.mlir_type, loc=loc)
vec_i4x2_type = ir.VectorType.get([2], Int4.mlir_type, loc=loc)
vec_dst_type = ir.VectorType.get([length], Float16.mlir_type, loc=loc)
vec_dst = llvm.mlir_zero(vec_dst_type, loc=loc, ip=ip)
# try to use vectorized version
if length >= 8:
num_vec8 = length // 8
for _ in range(num_vec8):
vec_f4e2m1x8 = vector.extract_strided_slice(
vec_i4x8_type, vec_src_i4, [src_pos], [8], [1], loc=loc, ip=ip
)
vec_f16x8 = cvt_f4e2m1x8_to_f16x8(vec_f4e2m1x8, loc=loc, ip=ip)
vec_dst = vector.insert_strided_slice(
vec_f16x8, vec_dst, [src_pos], [1], loc=loc, ip=ip
)
src_pos += 8
length -= 8
if length >= 4:
vec_f4e2m1x4 = vector.extract_strided_slice(
vec_i4x4_type, vec_src_i4, [src_pos], [4], [1], loc=loc, ip=ip
)
vec_f16x4 = cvt_f4e2m1x4_to_f16x4(vec_f4e2m1x4, loc=loc, ip=ip)
vec_dst = vector.insert_strided_slice(
vec_f16x4, vec_dst, [src_pos], [1], loc=loc, ip=ip
)
src_pos += 4
length -= 4
if length >= 2:
vec_f4e2m1x2 = vector.extract_strided_slice(
vec_i4x2_type, vec_src_i4, [src_pos], [2], [1], loc=loc, ip=ip
)
vec_f16x2 = cvt_f4e2m1x2_to_f16x2(vec_f4e2m1x2, loc=loc, ip=ip)
vec_dst = vector.insert_strided_slice(
vec_f16x2, vec_dst, [src_pos], [1], loc=loc, ip=ip
)
src_pos += 2
length -= 2
if length >= 1:
val_f16 = cvt_f4e2m1_f16(
vector.extractelement(
vec_src_i4,
position=arith.constant(Int32.mlir_type, src_pos),
loc=loc,
ip=ip,
),
loc=loc,
ip=ip,
)
vec_dst = vector.insertelement(
val_f16,
vec_dst,
position=arith.constant(Int32.mlir_type, src_pos),
loc=loc,
ip=ip,
)
return vec_dst
@dsl_user_op
def sext_unpacked_i4_i8_intrinsic(vec_unpacked_i4, length, *, loc=None, ip=None):
"""
@@ -295,9 +243,7 @@ def sext_unpacked_i4_i8_intrinsic(vec_unpacked_i4, length, *, loc=None, ip=None)
vec_unpacked_i4x4 = vector.extract_strided_slice(
vec_i8x4_type, vec_unpacked_i4, [pos], [4], [1], loc=loc, ip=ip
)
vec_i8x4 = sext_unpacked_i4x4_to_i8x4(
vec_unpacked_i4x4, loc=loc, ip=ip
)
vec_i8x4 = sext_unpacked_i4x4_to_i8x4(vec_unpacked_i4x4, loc=loc, ip=ip)
vec_i8 = vector.insert_strided_slice(
vec_i8x4, vec_i8, [pos], [1], loc=loc, ip=ip
)
@@ -312,6 +258,12 @@ cvt_i8_bf16_intrinsic.supported_archs = (
*Arch.HopperArchs(),
*Arch.BlackwellArchs(),
)
cvt_i8_bf16_intrinsic.s26_bf16_supported_archs = (
Arch.sm_100a,
Arch.sm_110a,
Arch.sm_120a,
Arch.sm_121a,
)
cvt_i4_bf16_intrinsic.supported_archs = (
Arch.sm_100a,
Arch.sm_110a,
+474 -124
View File
@@ -10,7 +10,7 @@
# is strictly prohibited.
from functools import partial
from typing import Optional, Tuple, Union, Callable, Literal
from typing import Any, Optional, Tuple, Union, Callable, Literal
from typing_extensions import deprecated
from cutlass.cutlass_dsl import T, dsl_user_op
@@ -20,15 +20,6 @@ import cutlass.cutlass_dsl as cutlass_dsl
from cutlass._mlir import ir
from cutlass._mlir.dialects import arith, llvm, nvvm, vector
# Forward nvvm enums
from cutlass._mlir.dialects.nvvm import (
ProxyKind,
SharedSpace,
Tcgen05WaitKind,
SetMaxRegisterAction,
RoundingModeKind,
)
from ..core import size
from ..typing import (
@@ -95,25 +86,23 @@ def _enhance_enum_with_str_mapping(enum_class):
"""
Convert a string literal to the corresponding enum member.
:param s: String representation of the enum member, or an enum member itself (deprecated)
:param s: String representation of the enum member
:return: The enum member (or None if s is None)
:raises ValueError: If the string is not a valid enum member
:raises TypeError: If an enum is passed instead of a string
"""
import warnings
if s is None:
return None
# Check if user passed an enum (should be a string literal instead)
# This catches cases where user passes e.g., RoundingModeKind.RN instead of "rn"
from enum import Enum
# Check if s is already an enum member of the correct type
if isinstance(s, cls):
warnings.warn(
f"Passing enum member directly to {cls.__name__}.from_str() is deprecated. "
f"Please use string literals instead (e.g., '{str(s)}' instead of {cls.__name__}.{s.name}).",
DeprecationWarning,
stacklevel=2,
if isinstance(s, Enum):
raise TypeError(
f"Expected a string literal for {cls.__name__}, but got enum '{type(s).__name__}.{s.name}'. "
f"Please pass a string instead (e.g., '{str(s)}' instead of {type(s).__name__}.{s.name}). "
f"Valid string options are: {sorted(str_to_enum_map.keys())}"
)
return s
if s not in str_to_enum_map:
valid_options = sorted(str_to_enum_map.keys())
raise ValueError(
@@ -446,6 +435,7 @@ def warp_reduction(
offset = offset // 2
return val
warp_reduction_max = partial(
warp_reduction,
op=lambda x, y: fmax(x, y) if isinstance(x, Float32) else cutlass_dsl.max(x, y),
@@ -460,34 +450,13 @@ def barrier(*, barrier_id=None, number_of_threads=None, loc=None, ip=None) -> No
"""
if barrier_id is not None:
barrier_id = Int32(barrier_id).ir_value(loc=loc, ip=ip)
else:
barrier_id = Int32(0).ir_value(loc=loc, ip=ip)
if number_of_threads is not None:
number_of_threads = Int32(number_of_threads).ir_value(loc=loc, ip=ip)
llvm.inline_asm(
None,
[barrier_id, number_of_threads],
"bar.sync $0, $1;",
"r,r",
has_side_effects=True,
is_align_stack=False,
asm_dialect=llvm.AsmDialect.AD_ATT,
loc=loc,
ip=ip,
)
else:
llvm.inline_asm(
None,
[barrier_id],
"bar.sync $0;",
"r",
has_side_effects=True,
is_align_stack=False,
asm_dialect=llvm.AsmDialect.AD_ATT,
loc=loc,
ip=ip,
)
nvvm.barrier(
barrier_id=barrier_id, number_of_threads=number_of_threads, loc=loc, ip=ip
)
@dsl_user_op
@@ -496,8 +465,6 @@ def barrier_arrive(
) -> None:
if barrier_id is not None:
barrier_id = Int32(barrier_id).ir_value(loc=loc, ip=ip)
else:
barrier_id = Int32(0).ir_value(loc=loc, ip=ip)
if number_of_threads is None:
raise ValueError(
@@ -505,14 +472,8 @@ def barrier_arrive(
)
number_of_threads = Int32(number_of_threads).ir_value(loc=loc, ip=ip)
llvm.inline_asm(
None,
[barrier_id, number_of_threads],
"bar.arrive $0, $1;",
"r,r",
has_side_effects=True,
is_align_stack=False,
asm_dialect=llvm.AsmDialect.AD_ATT,
nvvm.barrier_arrive(
barrier_id=barrier_id, number_of_threads=number_of_threads, loc=loc, ip=ip
)
@@ -638,19 +599,73 @@ def cluster_arrive_relaxed(*, aligned=None, loc=None, ip=None) -> None:
@dsl_user_op
def fence_proxy(
kind: ProxyKind,
kind: Literal[
"alias", "async", "async.global", "async.shared", "tensormap", "generic"
],
*,
space: Optional[SharedSpace] = None,
space: Optional[Literal["cta", "cluster"]] = None,
use_intrinsic=None,
loc=None,
ip=None,
) -> None:
"""
Fence operation to ensure memory consistency between proxies.
:param kind: Proxy kind string literal:
- "alias" : Alias proxy
- "async" : Async proxy
- "async.global" : Async global proxy
- "async.shared" : Async shared proxy
- "tensormap" : Tensormap proxy
- "generic" : Generic proxy
:type kind: Literal["alias", "async", "async.global", "async.shared", "tensormap", "generic"]
:param space: Shared memory space scope string literal (optional):
- "cta" : CTA (Cooperative Thread Array) scope
- "cluster" : Cluster scope
:type space: Optional[Literal["cta", "cluster"]]
:param use_intrinsic: Whether to use intrinsic version
"""
from cutlass._mlir.dialects.nvvm import (
SharedSpace,
ProxyKind,
)
# Enhance enum with str mapping
SharedSpace = _enhance_enum_with_str_mapping(SharedSpace)
ProxyKind = _enhance_enum_with_str_mapping(ProxyKind)
kind = ProxyKind.from_str(kind)
space = SharedSpace.from_str(space)
nvvm.fence_proxy(
kind=kind, space=space, use_intrinsic=use_intrinsic, loc=loc, ip=ip
kind=kind,
space=space,
use_intrinsic=use_intrinsic,
loc=loc,
ip=ip,
)
@dsl_user_op
def vote_sync_op(
pred: Boolean, kind: nvvm.VoteSyncKind, mask: Int = FULL_MASK, *, loc=None, ip=None
) -> Union[Int32, Boolean]:
"""
Performs a vote operation across the warp.
"""
return_type = Int32 if kind == nvvm.VoteSyncKind.ballot else Boolean
return return_type(
nvvm.vote_sync(
T.i32() if kind == nvvm.VoteSyncKind.ballot else T.bool(),
Int32(mask).ir_value(loc=loc, ip=ip),
Boolean(pred).ir_value(loc=loc, ip=ip),
kind,
loc=loc,
ip=ip,
)
)
def vote_ballot_sync(
pred: Boolean, mask: Int = FULL_MASK, *, loc=None, ip=None
) -> Int32:
@@ -668,45 +683,7 @@ def vote_ballot_sync(
See the `PTX documentation <https://docs.nvidia.com/cuda/parallel-thread-execution/#parallel-synchronization-and-communication-instructions-vote-sync>`__.
"""
return Int32(
nvvm.vote_ballot_sync(
T.i32(),
Int32(mask).ir_value(loc=loc, ip=ip),
Boolean(pred).ir_value(loc=loc, ip=ip),
loc=loc,
ip=ip,
)
)
@dsl_user_op
def vote_sync_op(
pred: Boolean, kind: str, mask: Int = FULL_MASK, *, loc=None, ip=None
) -> Union[Int32, Boolean]:
return_type = Boolean
return_type_str = "pred"
return return_type(
llvm.inline_asm(
T.bool(),
[
Boolean(pred).ir_value(loc=loc, ip=ip),
Int32(mask).ir_value(loc=loc, ip=ip),
],
f"""{{\n\t
.reg .pred ps;\n\t
.reg .pred pd;\n\t
setp.ne.b32 ps, $1, 0;\n\t
vote.sync.{kind}.{return_type_str} pd, ps, $2;\n\t
selp.b32 $0, 1, 0, pd;\n\t
}}""",
"=r,r,i",
has_side_effects=True,
is_align_stack=False,
asm_dialect=llvm.AsmDialect.AD_ATT,
loc=loc,
ip=ip,
)
)
return vote_sync_op(pred, nvvm.VoteSyncKind.ballot, mask, loc=loc, ip=ip)
@dsl_user_op
@@ -714,7 +691,7 @@ def vote_any_sync(
pred: Boolean, mask: Int = FULL_MASK, *, loc=None, ip=None
) -> Boolean:
"""True if source predicate is True for any non-exited threads in mask. Negate the source
predicate to compute .not_all.
predicate to compute .none.
:param pred: The predicate value for the current thread
:type pred: Boolean
@@ -727,7 +704,7 @@ def vote_any_sync(
See the `PTX documentation <https://docs.nvidia.com/cuda/parallel-thread-execution/#parallel-synchronization-and-communication-instructions-vote-sync>`__.
"""
return vote_sync_op(pred, "any", mask, loc=loc, ip=ip)
return vote_sync_op(pred, nvvm.VoteSyncKind.any, mask, loc=loc, ip=ip)
@dsl_user_op
@@ -748,7 +725,7 @@ def vote_all_sync(
See the `PTX documentation <https://docs.nvidia.com/cuda/parallel-thread-execution/#parallel-synchronization-and-communication-instructions-vote-sync>`__.
"""
return vote_sync_op(pred, "all", mask, loc=loc, ip=ip)
return vote_sync_op(pred, nvvm.VoteSyncKind.all, mask, loc=loc, ip=ip)
@dsl_user_op
@@ -767,7 +744,7 @@ def vote_uni_sync(
threads in mask
:rtype: Boolean
"""
return vote_sync_op(pred, "uni", mask, loc=loc, ip=ip)
return vote_sync_op(pred, nvvm.VoteSyncKind.uni, mask, loc=loc, ip=ip)
@dsl_user_op
@@ -821,8 +798,8 @@ def fence_view_async_tmem_op(
from cutlass._mlir.dialects.nvvm import Tcgen05WaitKind
# Enhance enum and convert string literal to enum type
Tcgen05WaitKind_enhanced = _enhance_enum_with_str_mapping(Tcgen05WaitKind)
kind = Tcgen05WaitKind_enhanced.from_str(kind)
Tcgen05WaitKind = _enhance_enum_with_str_mapping(Tcgen05WaitKind)
kind = Tcgen05WaitKind.from_str(kind)
nvvm.tcgen05_wait(kind=kind, loc=loc, ip=ip)
@@ -847,9 +824,8 @@ def fence_view_async_shared(
This function is usually used for async execution unit (like TMA, UMMA) after the load/store operations.
"""
nvvm.fence_proxy(
nvvm.ProxyKind.async_shared, space=nvvm.SharedSpace.shared_cta, loc=loc, ip=ip
)
# Use the fence_proxy wrapper function with string literals
fence_proxy(kind="async.shared", space="cta", loc=loc, ip=ip)
@dsl_user_op
@@ -859,6 +835,7 @@ def setmaxregister_increase(
loc=None,
ip=None,
):
from cutlass._mlir.dialects.nvvm import SetMaxRegisterAction
return nvvm.setmaxregister(reg_count, SetMaxRegisterAction.increase, loc=loc, ip=ip)
@@ -869,6 +846,7 @@ def setmaxregister_decrease(
loc=None,
ip=None,
):
from cutlass._mlir.dialects.nvvm import SetMaxRegisterAction
return nvvm.setmaxregister(reg_count, SetMaxRegisterAction.decrease, loc=loc, ip=ip)
@@ -880,6 +858,7 @@ def warpgroup_reg_alloc(
loc=None,
ip=None,
) -> None:
from cutlass._mlir.dialects.nvvm import SetMaxRegisterAction
nvvm.setmaxregister(reg_count, SetMaxRegisterAction.increase, loc=loc, ip=ip)
@@ -891,8 +870,10 @@ def warpgroup_reg_dealloc(
loc=None,
ip=None,
) -> None:
from cutlass._mlir.dialects.nvvm import SetMaxRegisterAction
nvvm.setmaxregister(reg_count, SetMaxRegisterAction.decrease, loc=loc, ip=ip)
@dsl_user_op
def calc_packed_f32x2_op(
src_a: Tuple[Float32, Float32],
@@ -908,8 +889,8 @@ def calc_packed_f32x2_op(
from cutlass._mlir.dialects.nvvm import RoundingModeKind
# Enhance enum and convert string literal to enum type
RoundingModeKind_enhanced = _enhance_enum_with_str_mapping(RoundingModeKind)
rnd = RoundingModeKind_enhanced.from_str(rnd)
RoundingModeKind = _enhance_enum_with_str_mapping(RoundingModeKind)
rnd = RoundingModeKind.from_str(rnd)
vec_type = ir.VectorType.get([2], Float32.mlir_type, loc=loc)
vec_src_a = vector.from_elements(
@@ -960,14 +941,13 @@ add_packed_f32x2 = partial(
calc_packed_f32x2_op, src_c=None, calc_func=nvvm.add_packed_f32x2
)
@dsl_user_op
def fmax(
a: Union[float, Float32], b: Union[float, Float32], *, loc=None, ip=None
) -> Float32:
return Float32(
nvvm.fmax(
T.f32(),
Float32(a).ir_value(loc=loc, ip=ip),
Float32(b).ir_value(loc=loc, ip=ip),
loc=loc,
@@ -975,12 +955,11 @@ def fmax(
)
)
@dsl_user_op
def rcp_approx(a: Union[float, Float32], *, loc=None, ip=None):
return Float32(
nvvm.rcp_approx_ftz_f(
T.f32(), Float32(a).ir_value(loc=loc, ip=ip), loc=loc, ip=ip
)
nvvm.rcp_approx_ftz_f(Float32(a).ir_value(loc=loc, ip=ip), loc=loc, ip=ip)
)
@@ -1024,6 +1003,68 @@ def cvt_i8_bf16(src_i8, *, loc=None, ip=None):
return val_bf16
@dsl_user_op
def cvt_i8x2_to_bf16x2(src_vec2, *, loc=None, ip=None):
# pack 2 int8 into 1 int16 value
src_i16 = llvm.bitcast(Int16.mlir_type, src_vec2, loc=loc, ip=ip)
val_i32 = llvm.inline_asm(
Int32.mlir_type,
[
src_i16,
],
"""{\n\t
.reg .b16 scale;\n\t
mov.b16 scale, 0x8585;\n\t
cvt.rn.satfinite.scaled::n2::ue8m0.bf16x2.s2f6x2 $0, $1, scale;\n\t
}""",
"=r,h",
)
vec_bf16x2_type = ir.VectorType.get([2], BFloat16.mlir_type, loc=loc)
vec_bf16x2 = llvm.bitcast(vec_bf16x2_type, val_i32, loc=loc, ip=ip)
return vec_bf16x2
@dsl_user_op
def cvt_i8x4_to_bf16x4(src_vec4, *, loc=None, ip=None):
# pack 4 int8 into 1 int32 value
src_i32 = llvm.bitcast(Int32.mlir_type, src_vec4, loc=loc, ip=ip)
rst01 = llvm.inline_asm(
Int32.mlir_type,
[
src_i32,
],
"""{\n\t
.reg .b16 pair<2>;\n\t
.reg .b16 scale;\n\t
mov.b32 {pair0, pair1}, $1;\n\t
mov.b16 scale, 0x8585;\n\t
cvt.rn.satfinite.scaled::n2::ue8m0.bf16x2.s2f6x2 $0, pair0, scale;\n\t
}""",
"=r,r",
)
rst23 = llvm.inline_asm(
Int32.mlir_type,
[
src_i32,
],
"""{\n\t
.reg .b16 pair<2>;\n\t
.reg .b16 scale;\n\t
mov.b32 {pair0, pair1}, $1;\n\t
mov.b16 scale, 0x8585;\n\t
cvt.rn.satfinite.scaled::n2::ue8m0.bf16x2.s2f6x2 $0, pair1, scale;\n\t
}""",
"=r,r",
)
vec_type = ir.VectorType.get([2], Int32.mlir_type, loc=loc)
rst_i32 = vector.from_elements(vec_type, [rst01, rst23], loc=loc, ip=ip)
vec_bf16x4_type = ir.VectorType.get([4], BFloat16.mlir_type, loc=loc)
vec_bf16x4 = llvm.bitcast(vec_bf16x4_type, rst_i32, loc=loc, ip=ip)
return vec_bf16x4
# Convert vector of 2 float values to vector of 2 bfloat16 values with satfinite rounding
@dsl_user_op
def cvt_f32x2_bf16x2(src_vec2, *, loc=None, ip=None):
@@ -1263,12 +1304,116 @@ def prmt(src, src_reg_shifted, prmt_indices, *, loc=None, ip=None):
@dsl_user_op
def cvt_i4_bf16(src_i4, *, loc=None, ip=None):
# i4 -> i32 -> f32 -> bf
src_i32 = llvm.zext(Int32.mlir_type, src_i4, loc=loc, ip=ip)
src_i32 = llvm.sext(Int32.mlir_type, src_i4, loc=loc, ip=ip)
src_f32 = llvm.sitofp(Float32.mlir_type, src_i32, loc=loc, ip=ip)
bf16_val = cvt_f32_bf16(src_f32, loc=loc, ip=ip)
return bf16_val
# Convert multiple shuffled int4 values to bfloat16 values.
# The input elements are assumed to be already shuffled following a specific shuffle pattern.
# Specifically, for consecutive 8 int4 values with indices of (0, 1, 2, 3, 4, 5, 6, 7),
# they are shuffled to (0, 2, 1, 3, 4, 6, 5, 7). For tailing elements less than 8, the
# shuffle pattern is (0, 2, 1, 3) for 4 elements. No shuffle is needed for less than 4 elements.
# Shuffle could help to produce converted bf16 values in the natural order of (0, 1, 2 ,3 ,4 ,5 ,6 ,7)
# without extra prmt instructions and thus better performance.
# The number of elements to be converted must be be even as specified by num_elts.
# Int4 values are packed into int32 values with upper bits filled with 0 if there are less than 4 int4 values.
# Results bfloat16 values are also packed into int32 values.
@dsl_user_op
def cvt_i4_to_bf16_with_shuffle_impl(src_i32, num_elts, *, loc=None, ip=None):
from cutlass import CUDA_VERSION
if CUDA_VERSION.major < 13:
raise cutlass_dsl.DSLCudaVerNotImplemented(
feature="cvt_i4_to_bf16_with_shuffle_impl", required_version="13.1"
)
num_i32_elts = num_elts // 2
mask_odd = arith.constant(Int32.mlir_type, 0xF0F0F0F0, loc=loc, ip=ip)
mask_even = arith.constant(Int32.mlir_type, 0x0F0F0F0F, loc=loc, ip=ip)
src_odd = arith.andi(src_i32, mask_odd, loc=loc, ip=ip)
src_even = arith.andi(src_i32, mask_even, loc=loc, ip=ip)
c4 = arith.constant(Int32.mlir_type, 4, loc=loc, ip=ip)
src_even = arith.shli(src_even, c4, loc=loc, ip=ip)
rst13 = llvm.inline_asm(
Int32.mlir_type,
[
src_odd,
],
"""{\n\t
.reg .b16 pair<2>;\n\t
.reg .b16 scale;\n\t
mov.b32 {pair0, pair1}, $1;\n\t
mov.b16 scale, 0x8181;\n\t
cvt.rn.satfinite.scaled::n2::ue8m0.bf16x2.s2f6x2 $0, pair0, scale;\n\t
}""",
"=r,r",
)
rst57 = llvm.inline_asm(
Int32.mlir_type,
[
src_odd,
],
"""{\n\t
.reg .b16 pair<2>;\n\t
.reg .b16 scale;\n\t
mov.b32 {pair0, pair1}, $1;\n\t
mov.b16 scale, 0x8181;\n\t
cvt.rn.satfinite.scaled::n2::ue8m0.bf16x2.s2f6x2 $0, pair1, scale;\n\t
}""",
"=r,r",
)
rst02 = llvm.inline_asm(
Int32.mlir_type,
[
src_even,
],
"""{\n\t
.reg .b16 pair<2>;\n\t
.reg .b16 scale;\n\t
mov.b16 scale, 0x8181;\n\t
mov.b32 {pair0, pair1}, $1;\n\t
cvt.rn.satfinite.scaled::n2::ue8m0.bf16x2.s2f6x2 $0, pair0, scale;\n\t
}""",
"=r,r",
)
rst46 = llvm.inline_asm(
Int32.mlir_type,
[
src_even,
],
"""{\n\t
.reg .b16 pair<2>;\n\t
.reg .b16 scale;\n\t
mov.b16 scale, 0x8181;\n\t
mov.b32 {pair0, pair1}, $1;\n\t
cvt.rn.satfinite.scaled::n2::ue8m0.bf16x2.s2f6x2 $0, pair1, scale;\n\t
}""",
"=r,r",
)
vec_type = ir.VectorType.get([num_i32_elts], Int32.mlir_type, loc=loc)
if num_elts == 2:
prmt_index = arith.constant(Int32.mlir_type, 0x00005410, loc=loc, ip=ip)
rst = llvm.inline_asm(
Int32.mlir_type,
[
rst02,
rst13,
prmt_index,
],
"prmt.b32 $0, $1, $2, $3;",
"=r,r,r,r",
)
vec_rsts = vector.from_elements(vec_type, [rst], loc=loc, ip=ip)
elif num_elts == 4:
vec_rsts = vector.from_elements(vec_type, [rst02, rst13], loc=loc, ip=ip)
else:
vec_rsts = vector.from_elements(
vec_type, [rst02, rst13, rst46, rst57], loc=loc, ip=ip
)
return vec_rsts
# Convert multiple int4 values to bfloat16 values.
# The number of elements to be converted must be be even as specified by num_elts.
# Int4 values are packed into int32 values with upper bits filled with 0 if there are less than 4 int4 values.
@@ -1357,11 +1502,12 @@ def cvt_i4_to_bf16_impl(src_i32, num_elts, *, loc=None, ip=None):
# Convert 2 int4 values to 2 bfloat16 values
@dsl_user_op
def cvt_i4x2_to_bf16x2(src_vec2, *, loc=None, ip=None):
def cvt_i4x2_to_bf16x2(src_vec2, *, with_shuffle=False, loc=None, ip=None):
cvt_func = cvt_i4_to_bf16_with_shuffle_impl if with_shuffle else cvt_i4_to_bf16_impl
# pack 2 int4 into 1 int32 value and fill upper bits with 0
src_i8 = llvm.bitcast(Int8.mlir_type, src_vec2, loc=loc, ip=ip)
src_i32 = llvm.zext(Int32.mlir_type, src_i8, loc=loc, ip=ip)
rst_i32 = cvt_i4_to_bf16_impl(src_i32, 2, loc=loc, ip=ip)
rst_i32 = cvt_func(src_i32, 2, loc=loc, ip=ip)
vec_bf16x2_type = ir.VectorType.get([2], BFloat16.mlir_type, loc=loc)
vec_bf16x2 = llvm.bitcast(vec_bf16x2_type, rst_i32, loc=loc, ip=ip)
return vec_bf16x2
@@ -1369,11 +1515,12 @@ def cvt_i4x2_to_bf16x2(src_vec2, *, loc=None, ip=None):
# Convert 4 int4 values to 4 bfloat16 values
@dsl_user_op
def cvt_i4x4_to_bf16x4(src_vec4, *, loc=None, ip=None):
def cvt_i4x4_to_bf16x4(src_vec4, *, with_shuffle=False, loc=None, ip=None):
cvt_func = cvt_i4_to_bf16_with_shuffle_impl if with_shuffle else cvt_i4_to_bf16_impl
# pack 4 int4 into 1 int32 value and fill upper bits with 0
src_i16 = llvm.bitcast(Int16.mlir_type, src_vec4, loc=loc, ip=ip)
src_i32 = llvm.zext(Int32.mlir_type, src_i16, loc=loc, ip=ip)
rst_i32 = cvt_i4_to_bf16_impl(src_i32, 4, loc=loc, ip=ip)
rst_i32 = cvt_func(src_i32, 4, loc=loc, ip=ip)
vec_bf16x4_type = ir.VectorType.get([4], BFloat16.mlir_type, loc=loc)
vec_bf16x4 = llvm.bitcast(vec_bf16x4_type, rst_i32, loc=loc, ip=ip)
return vec_bf16x4
@@ -1381,14 +1528,16 @@ def cvt_i4x4_to_bf16x4(src_vec4, *, loc=None, ip=None):
# Convert 8 int4 values to 8 bfloat16 values
@dsl_user_op
def cvt_i4x8_to_bf16x8(src_vec8, *, loc=None, ip=None):
def cvt_i4x8_to_bf16x8(src_vec8, *, with_shuffle=False, loc=None, ip=None):
cvt_func = cvt_i4_to_bf16_with_shuffle_impl if with_shuffle else cvt_i4_to_bf16_impl
# pack 8 int4 into 1 int32 value and fill upper bits with 0
src_i32 = llvm.bitcast(Int32.mlir_type, src_vec8, loc=loc, ip=ip)
rst_i32 = cvt_i4_to_bf16_impl(src_i32, 8, loc=loc, ip=ip)
rst_i32 = cvt_func(src_i32, 8, loc=loc, ip=ip)
vec_bf16x8_type = ir.VectorType.get([8], BFloat16.mlir_type, loc=loc)
vec_bf16x8 = llvm.bitcast(vec_bf16x8_type, rst_i32, loc=loc, ip=ip)
return vec_bf16x8
# Sign extend 4 int4 unpacked in 8b containers
@dsl_user_op
def sext_unpacked_i4x4_to_i8x4(src_vec4, *, loc=None, ip=None):
@@ -1485,6 +1634,196 @@ def griddepcontrol_launch_dependents(*, loc=None, ip=None) -> None:
@dsl_user_op
def _warp_redux_sync_nvvm(
value: Numeric,
kind: Literal[
"fmax",
"fmin",
"max",
"min",
"add",
"xor",
"or",
"and",
],
mask_and_clamp: Int = FULL_MASK,
abs: bool = False,
nan: bool = None,
*,
loc=None,
ip=None,
) -> Numeric:
from cutlass._mlir.dialects.nvvm import ReduxKind
# Enhance enum and convert string literal to enum type
ReduxKind = _enhance_enum_with_str_mapping(ReduxKind)
kind = ReduxKind.from_str(kind)
value_type = type(value)
value_ir = value.ir_value(loc=loc, ip=ip)
return value_type(
nvvm.redux_sync(
res=value_ir.type,
val=value_ir,
kind=kind,
mask_and_clamp=Int32(mask_and_clamp).ir_value(loc=loc, ip=ip),
abs=abs,
nan=nan,
loc=loc,
ip=ip,
)
)
@dsl_user_op
def _warp_redux_sync_ptx(
value: Numeric,
kind: Literal[
"fmax",
"fmin",
"max",
"min",
],
mask_and_clamp: Int = FULL_MASK,
abs: bool = None,
nan: bool = None,
*,
loc=None,
ip=None,
) -> Numeric:
value_type = type(value)
value_ir = value.ir_value(loc=loc, ip=ip)
mlir_type = value_type.mlir_type
mask_ir = Int32(mask_and_clamp).ir_value(loc=loc, ip=ip)
kind_ptx_str = kind
if kind == "fmax":
kind_ptx_str = "max"
elif kind == "fmin":
kind_ptx_str = "min"
modifiers = []
if nan is True:
modifiers.append("NaN")
if abs is True:
modifiers.append("abs")
modifier_str = "." + ".".join(modifiers) if modifiers else ""
ptx_instr = f"redux.sync.{kind_ptx_str}{modifier_str}.f32 $0, $1, $2;"
return value_type(
llvm.inline_asm(
mlir_type,
[value_ir, mask_ir],
f"{ptx_instr}",
f"=f,f,i",
has_side_effects=True,
is_align_stack=False,
asm_dialect=llvm.AsmDialect.AD_ATT,
)
)
@dsl_user_op
def warp_redux_sync(
value: Numeric,
kind: Literal[
"fmax",
"fmin",
"max",
"min",
"add",
"xor",
"or",
"and",
],
mask_and_clamp: Int = FULL_MASK,
*,
abs: bool = None,
nan: bool = None,
loc=None,
ip=None,
) -> Numeric:
"""
Perform warp-level reduction operation across threads.
Reduces values from participating threads in a warp according to the specified operation.
All threads in the mask receive the same result.
:param value: Input value to reduce
:type value: Numeric
:param kind: Reduction operation. Supported operations:
- Integer types (Int32/Uint32): "add", "and", "max", "min", "or", "xor"
- Float types (Float32): "fmax", "fmin" (or "max"/"min" which auto-convert to "fmax"/"fmin")
:type kind: Literal["add", "and", "max", "min", "or", "xor", "fmin", "fmax"]
:param mask_and_clamp: Warp participation mask (default: FULL_MASK = 0xFFFFFFFF)
:type mask_and_clamp: Int
:param abs: Apply absolute value before reduction (float types only)
:type abs: bool
:param nan: Enable NaN propagation for fmax/fmin operations (float types only)
:type nan: Optional[bool]
:return: Reduced value (same for all participating threads)
:rtype: Numeric
"""
# Convert value to Numeric type if needed
if not isinstance(value, Numeric):
value = as_numeric(value)
# Determine value type and choose appropriate implementation
value_type = type(value)
mlir_type = value_type.mlir_type
# Use inline PTX for float types, NVVM for integer types
if mlir_type == T.f32():
return _warp_redux_sync_ptx(
value, kind, mask_and_clamp, abs, nan, loc=loc, ip=ip
)
else:
return _warp_redux_sync_nvvm(
value, kind, mask_and_clamp, abs, nan, loc=loc, ip=ip
)
@dsl_user_op
def atomic_max_float32(
ptr,
value: Float32,
*,
positive_only: bool = True,
loc=None,
ip=None,
) -> Float32:
"""
Performs an atomic max operation on a float32 value in global memory.
This implementation works correctly for non-negative values (>= 0) using direct bitcast.
:param ptr: Pointer to the memory location
:param value: The float32 value to compare and potentially store (should be >= 0 for correct results)
:type value: Float32
:param positive_only: If True (default), assumes input values are non-negative.
This parameter is provided for API compatibility and future extensions.
:type positive_only: bool
:return: The old value at the memory location
:rtype: Float32
"""
from cutlass._mlir.dialects.nvvm import AtomicOpKind
value_int = llvm.bitcast(T.i32(), value.ir_value(loc=loc, ip=ip), loc=loc, ip=ip)
old_value_int = nvvm.atomicrmw(
AtomicOpKind.MAX,
ptr,
value_int,
loc=loc,
ip=ip,
)
return Float32(llvm.bitcast(T.f32(), old_value_int, loc=loc, ip=ip))
def _normalize_ptr(addr, *, loc=None, ip=None) -> ir.Value:
"""
Helper function to normalize pointer types to MLIR ir.Value.
@@ -1549,7 +1888,7 @@ def _atomic(
:rtype: Union[Numeric, ir.Value]
"""
from cutlass._mlir.dialects.nvvm import AtomicOpKind, MemOrderKind, MemScopeKind
from cutlass.utils.version_info import CUDA_VERSION
from cutlass import CUDA_VERSION
# Enhance enums and convert string literals to enum types
AtomicOpKind = _enhance_enum_with_str_mapping(AtomicOpKind)
@@ -1848,7 +2187,7 @@ def atomic_cas(
:rtype: Numeric
"""
from cutlass._mlir.dialects.nvvm import AtomicOpKind, MemOrderKind, MemScopeKind
from cutlass.utils.version_info import CUDA_VERSION
from cutlass import CUDA_VERSION
# Enhance enums and convert string literals to enum types
MemOrderKind = _enhance_enum_with_str_mapping(MemOrderKind)
@@ -1882,6 +2221,17 @@ def atomic_cas(
loc=loc,
ip=ip,
)
elif CUDA_VERSION.major == 13 and CUDA_VERSION.minor == 1:
result = nvvm.atomicrmw(
op=AtomicOpKind.CAS,
ptr=ptr,
a=val_ir,
b=cmp_ir,
mem_order=sem,
syncscope=scope,
loc=loc,
ip=ip,
)
else:
result = nvvm.atomicrmw(
op=AtomicOpKind.CAS,
+1 -1
View File
@@ -17,7 +17,7 @@ import cutlass._mlir.dialects.cute as _cute_ir
import cutlass._mlir.dialects.cute_nvgpu as _cute_nvgpu_ir
from cutlass._mlir import ir
from ..typing import Pointer, Numeric, NumericMeta
from ..typing import Pointer, Numeric, NumericMeta, Layout
@dsl_user_op
+1 -4
View File
@@ -55,7 +55,6 @@ def get_max_tmem_alloc_cols(compute_capability: str) -> int:
return TMEM_MAX_ALLOC_COLUMNS_MAP[compute_capability]
def get_min_tmem_alloc_cols(compute_capability: str) -> int:
"""Get the minimum TMEM allocation columns for a given compute capability.
@@ -179,11 +178,9 @@ def dealloc_tmem(
:param num_columns: The number of columns in the TMEM allocation
:type num_columns: Int
:param is_two_cta: Optional boolean parameter for 2-CTA MMAs
:param arch: The architecture of the GPU.
:type arch: str
"""
tmem_max_alloc_cols = get_max_tmem_alloc_cols(arch)
tmem_min_alloc_cols = get_min_tmem_alloc_cols(arch)
tmem_max_alloc_cols = get_max_tmem_alloc_cols(arch)
if isinstance(num_columns, int):
if (
num_columns < tmem_min_alloc_cols
+51
View File
@@ -285,6 +285,8 @@ class MmaAtom(Atom):
if self.op is not None:
self.op._verify_fragment_B(input, loc=loc, ip=ip)
input = input.value
if isinstance(input, tuple):
input = _pack_shape(input, loc=loc, ip=ip)
return _cute_ir.mma_make_fragment(
_cute_ir.MmaOperand.B, self._trait.value, input, loc=loc, ip=ip
)
@@ -1190,3 +1192,52 @@ def copy_atom_call(
return _cute_ir.copy_atom_call(
value, src.value, dst.value, pred=pred, loc=loc, ip=ip
)
@dsl_user_op
def mma_atom_call(
atom: MmaAtom,
d: Tensor,
a: Tensor,
b: Tensor,
c: Tensor,
*,
loc=None,
ip=None,
**kwargs,
) -> None:
"""
Execute a single MMA atom operation.
The mma_atom_call operation executes an MMA atom with the given operands.
This performs a matrix multiplication and accumulation operation:
D = A * B + C
Note: The tensors 'd', 'a', 'b', and 'c' must only have a single fragment.
:param atom: The MMA atom to execute
:type atom: MmaAtom
:param d: Destination tensor (output accumulator)
:type d: Tensor
:param a: First source tensor (matrix A)
:type a: Tensor
:param b: Second source tensor (matrix B)
:type b: Tensor
:param c: Third source tensor (input accumulator C)
:type c: Tensor
:param loc: Source location for MLIR, defaults to None
:type loc: Optional[Location], optional
:param ip: Insertion point, defaults to None
:type ip: Optional[InsertionPoint], optional
Examples:
.. code-block:: python
# Call an MMA atom operation
cute.mma_atom_call(mma_atom, d_tensor, a_tensor, b_tensor, c_tensor)
"""
value = atom._unpack(loc=loc, ip=ip, **kwargs)
return _cute_ir.mma_atom_call(
value, d.value, a.value, b.value, c.value, loc=loc, ip=ip
)
+99 -73
View File
@@ -10,14 +10,14 @@
# is strictly prohibited.
from functools import partial, reduce
import inspect
from inspect import isclass
from typing import Any, Dict, List, Optional, Tuple, Type, Union, overload
from cutlass import const_expr
from typing_extensions import deprecated
from cutlass._mlir import ir
from cutlass._mlir.dialects import builtin, llvm, vector
from cutlass._mlir.dialects import builtin, llvm, vector, arith, nvvm
from cutlass._mlir.dialects import cute as _cute_ir
from cutlass._mlir.dialects.cute import (
Ratio as _Ratio,
@@ -125,6 +125,7 @@ __all__ = [
"shape",
"recast_ptr",
"make_ptr",
"get_remote_smem_ptr_in_cluster",
"composition",
"complement",
"right_inverse",
@@ -247,7 +248,7 @@ def _unpack_x_tuple(t: Union[ir.Type, ir.Value], *, loc=None, ip=None) -> XTuple
vals = []
else:
vals = get_leaves(t, loc=loc, ip=ip)
if not isinstance(vals, list):
if not isinstance(vals, ir.OpResultList):
vals = [vals]
else:
raise TypeError(f"expects static type or value, but got {t}")
@@ -383,9 +384,9 @@ class IntValue(cutlass_arith.ArithValue):
@property
def divisibility(self):
assert isinstance(
self.get_typed_value().type, _cute_ir.IntTupleType
), f"expected self.get_typed_value() to be int_tuple type, but got {self.get_typed_value().type}"
assert isinstance(self.get_typed_value().type, _cute_ir.IntTupleType), (
f"expected self.get_typed_value() to be int_tuple type, but got {self.get_typed_value().type}"
)
return self.get_typed_value().type.get_divisibility([0])
def __str__(self):
@@ -429,7 +430,9 @@ class IntValue(cutlass_arith.ArithValue):
@dsl_user_op
@_binary_op
def __add__(self, other, *, loc=None, ip=None):
return _cute_ir.tuple_add(self.get_typed_value(), other, loc=loc, ip=ip)
return _cute_ir.tuple_add(
self.get_typed_value(loc=loc, ip=ip), other, loc=loc, ip=ip
)
@dsl_user_op
@_binary_op
@@ -461,8 +464,10 @@ class IntValue(cutlass_arith.ArithValue):
@dsl_user_op
@_binary_op
def __radd__(self, other, *, loc=None, ip=None):
return _cute_ir.tuple_add(other, self.get_typed_value(), loc=loc, ip=ip)
def __radd__(self, other, *, loc=None, ip=None) -> "IntValue":
return _cute_ir.tuple_add(
other, self.get_typed_value(loc=loc, ip=ip), loc=loc, ip=ip
)
@dsl_user_op
@_binary_op
@@ -1207,10 +1212,6 @@ class _ComposedLayout(ComposedLayout):
@property
@dsl_user_op
def shape(self, *, loc=None, ip=None) -> Shape:
return self.shape_method(loc=loc, ip=ip)
@dsl_user_op
def shape_method(self, *, loc=None, ip=None) -> Shape:
return _unpack_x_tuple(
_cute_ir.get_shape(self.value, loc=loc, ip=ip), loc=loc, ip=ip
)
@@ -1262,9 +1263,9 @@ class _ComposedLayout(ComposedLayout):
# In this context, a _ComposedLayout instance is an encapsulated ir.Value which is automatically created
# by value caster for ComposedLayout typed values
assert len(values) == 1, f"Expected 1 value, but got {len(values)}"
assert isinstance(
values[0], (_ComposedLayout, ir.Value)
), f"Expected _ComposedLayout or ir.Value, but got {type(values[0])}"
assert isinstance(values[0], (_ComposedLayout, ir.Value)), (
f"Expected _ComposedLayout or ir.Value, but got {type(values[0])}"
)
return _ComposedLayout(
values[0] if isinstance(values[0], ir.Value) else values[0].value,
)
@@ -1313,9 +1314,9 @@ class _Pointer(Pointer):
# In this context, a _Pointer instance is an encapsulated ir.Value which is automatically created
# by value caster for cute.ptr typed values
assert len(values) == 1, f"Expected 1 value, but got {len(values)}"
assert isinstance(
values[0], (_Pointer, ir.Value)
), f"Expected _Pointer or ir.Value, but got {type(values[0])}"
assert isinstance(values[0], (_Pointer, ir.Value)), (
f"Expected _Pointer or ir.Value, but got {type(values[0])}"
)
return _Pointer(
values[0] if isinstance(values[0], ir.Value) else values[0].value
)
@@ -1359,29 +1360,12 @@ class _Pointer(Pointer):
"""
Get the LLVM pointer representation of this pointer.
:param loc: Source location for MLIR, defaults to None
:type loc: Optional[Location]
:param ip: Insertion point for MLIR, defaults to None
:type ip: Optional[InsertionPoint]
:return: The LLVM pointer representation
:rtype: ir.Value
"""
return self.to_llvm_ptr(loc=loc, ip=ip)
@dsl_user_op
@lru_cache_ir()
def to_llvm_ptr(self, *, loc=None, ip=None) -> ir.Value:
"""
Get the LLVM pointer representation of this pointer. (Used by internal API to propagate loc and ip)
:param loc: Source location for MLIR, defaults to None
:type loc: Optional[Location]
:param ip: Insertion point for MLIR, defaults to None
:type ip: Optional[InsertionPoint]
:return: The LLVM pointer representation
:rtype: ir.Value
"""
llvm_ptr_ty = llvm.PointerType.get(self.memspace.value)
llvm_ptr_ty = llvm.PointerType.get(
self.memspace.value if self.memspace != AddressSpace.rmem else 0
)
return builtin.unrealized_conversion_cast(
[llvm_ptr_ty], [self.value], loc=loc, ip=ip
)
@@ -1679,7 +1663,16 @@ def printf(*args, loc=None, ip=None) -> None:
elif isinstance(arg0, tuple):
# Assume it's a tile
return _pack_tile(arg0)
elif isinstance(arg0, (_Tensor, _Pointer, _ComposedLayout)):
elif isinstance(arg0, _Tensor):
arg0._check_can_load_store()
if isinstance(arg0.layout, ComposedLayout) and isinstance(
arg0.layout.inner, Swizzle
):
raise NotImplementedError(
"tensor with swizzled layout (PISL) is not supported in printf, please use swizzled pointer (PDSL) instead"
)
return arg0.value
elif isinstance(arg0, (_Pointer, _ComposedLayout)):
return arg0.value
else:
raise TypeError(f"unsupported argument type in printf, got {type(arg)}")
@@ -1751,6 +1744,7 @@ def make_swizzle(b, m, s, *, loc=None, ip=None):
return Swizzle(static(ty, loc=loc, ip=ip))
@dsl_user_op
def static(value, *, loc=None, ip=None):
return _cute_ir.static(value, loc=loc, ip=ip)
@@ -3409,39 +3403,90 @@ def make_ptr(
loc=None,
ip=None,
) -> Pointer:
# Perform checks
if dtype is None or not isinstance(dtype, NumericMeta):
raise TypeError(f"expects dtype to be a type of Numeric, but got {dtype}")
if not isinstance(mem_space, AddressSpace):
raise TypeError(f"expects mem_space to be an AddressSpace, but got {mem_space}")
if isinstance(value, ir.Value) and llvm.PointerType.isinstance(value.type):
value = llvm.ptrtoint(T.i64(), value)
if not is_integer(value):
raise TypeError(f"expects integer value, but got {type(value)}")
value = Int32(value) if mem_space == AddressSpace.tmem else Int64(value)
# TMEM addresses are 32b wide
is_tmem = mem_space == AddressSpace.tmem
value = Int32(value) if mem_space == AddressSpace.tmem else Int64(value)
# Set the alignment of the pointer
bytes_per_elt = max(1, dtype.width // 8)
if assumed_align is None:
assumed_align = bytes_per_elt
if bytes_per_elt % assumed_align != 0 and assumed_align % bytes_per_elt != 0:
raise ValueError(
f"{bytes_per_elt=} is not a multiple of {assumed_align=} and vice versa."
)
aligned_ty = _cute_ir.ConstrainedIntType.get(assumed_align, type(value).width)
aligned_intptr = _cute_ir.assume(
aligned_ty, value.ir_value(loc=loc, ip=ip), loc=loc, ip=ip
)
# Construct the pointer Type
data_ty = T.i8() if dtype is None else dtype.mlir_type
ptr_ty = _cute_ir.PtrType.get(data_ty, mem_space, assumed_align)
return _cute_ir.inttoptr(ptr_ty, aligned_intptr, loc=loc, ip=ip)
@dsl_user_op
def get_remote_smem_ptr_in_cluster(
smem_ptr: Pointer,
cta_rank_in_cluster: Int,
*,
loc=None,
ip=None,
) -> Pointer:
"""
Get the remote shared memory CuTe pointer in a cluster.
:param smem_ptr: The current shared memory pointer
:type smem_ptr: Pointer
:param cta_rank_in_cluster: The peer CTA rank in cluster to get the remote pointer for
:type cta_rank_in_cluster: Int
:param loc: Source location for MLIR, defaults to None
:type loc: Optional[Location]
:param ip: Insertion point, defaults to None
:type ip: Optional[InsertionPoint]
:return: The remote shared memory CuTe pointer
:rtype: Pointer
"""
cur_llvm_ptr = smem_ptr.llvm_ptr
remote_llvm_ptr = nvvm.mapa(
llvm.PointerType.get(7), # LLVM dsmem address space
cur_llvm_ptr,
Int32(cta_rank_in_cluster).ir_value(loc=loc, ip=ip),
loc=loc,
ip=ip,
)
remote_llvm_ptr_cast = llvm.addrspacecast(
llvm.PointerType.get(AddressSpace.smem), remote_llvm_ptr, loc=loc, ip=ip
)
remote_ptr = make_ptr(
smem_ptr.dtype,
remote_llvm_ptr_cast,
AddressSpace.smem,
assumed_align=smem_ptr.alignment,
loc=loc,
ip=ip,
)
if const_expr(smem_ptr.value.type.is_swizzled):
sw = Swizzle(static(smem_ptr.value.type.swizzle_type))
remote_ptr = recast_ptr(
remote_ptr, swizzle_=sw, dtype=smem_ptr.dtype, loc=loc, ip=ip
)
return remote_ptr
#
# Layout algebra
#
@@ -3868,9 +3913,7 @@ def local_tile(
return _cute_ir.local_tile(
input=input.value,
tile=tiler_val,
static_tile=None,
coord=coord_val,
static_coord=None,
proj=proj,
loc=loc,
ip=ip,
@@ -3907,9 +3950,9 @@ def make_layout_image_mask(
sliced_lay, offset = slice_and_offset(slicer, lay, loc=loc, ip=ip)
# Given that we replace only one mode with _, the rank of the slice should be 1
assert rank(sliced_lay) == 1
assert is_static(
sliced_lay
), "make_layout_image_mask requires the layout to be static"
assert is_static(sliced_lay), (
"make_layout_image_mask requires the layout to be static"
)
# Create the mask of the image
mcast_mask = Int16(0)
@@ -3952,6 +3995,7 @@ def leading_dim(shape: Shape, stride: Stride) -> Union[int, Tuple[int, ...], Non
return find_if(stride, pred_fn=pred_fn)
@dsl_user_op
def make_layout_tv(
thr_layout: Layout, val_layout: Layout, *, loc=None, ip=None
@@ -4468,9 +4512,9 @@ class struct:
"""
Return the round-up offset up to the next multiple of align.
"""
assert align > 0 and not (
align & (align - 1)
), "align should be a strictly positive power of 2."
assert align > 0 and not (align & (align - 1)), (
"align should be a strictly positive power of 2."
)
return (offset + (align - 1)) & ~(align - 1)
@@ -4607,29 +4651,11 @@ class FastDivmodDivisor:
new_obj = object.__new__(FastDivmodDivisor)
new_obj._divisor = values[0]
return new_obj
def __repr__(self):
return f"FastDivmodDivisor({self._divisor.type})"
# Set explicit signature for Sphinx documentation to avoid issues with @dsl_user_op decorator
FastDivmodDivisor.__init__.__signature__ = inspect.Signature(
[
inspect.Parameter("self", inspect.Parameter.POSITIONAL_OR_KEYWORD),
inspect.Parameter(
"divisor",
inspect.Parameter.POSITIONAL_OR_KEYWORD,
annotation=Integer,
),
inspect.Parameter(
"is_power_of_2",
inspect.Parameter.POSITIONAL_OR_KEYWORD,
default=None,
annotation=bool,
),
]
)
@dsl_user_op
def fast_divmod_create_divisor(
divisor: Integer, *, loc=None, ip=None
@@ -0,0 +1,54 @@
# CuTe Experimental APIs
> **Note:** APIs in this module are experimental and subject to change.
>
> This module serves as a staging area for new CuTe functionality that is still under active development. Performance, compile time, and interoperability with CuTe are works in progress. API signatures, behavior, and naming conventions may change without notice between releases.
>
> Once these APIs are stabilized, they will be migrated to the main `cute` submodules.
>
> Users are encouraged to experiment with these APIs but should be prepared to update their code as the interfaces evolve.
## Core APIs (`core.py`)
- `elect_sync` — Elects one thread within a warp
- `get_mbarrier` — Returns the mbarrier pointer for a given stage token
- `create_pipeline` — Creates a circular buffer of synchronization primitives indexed by stage count
- `create_pipeline_with_mask` — Creates a pipeline with an arrival mask for cluster-scoped synchronization
- `pipeline_advance_iterator` — Advances a pipeline iterator to the next stage
- `producer_acquire` / `producer_commit` — Producer-side pipeline synchronization
- `consumer_wait` / `consumer_release` / `consumer_tail` — Consumer-side pipeline synchronization
- `get_pipeline_produce_stage` / `get_pipeline_consume_stage` — Gets pipeline stage tokens
## Memory APIs (`memory.py`)
- `allocate` — Allocate a buffer with given type, layout, and address space
- `tma_load` — Copy tensor from global memory to shared memory using TMA
- `tma_load_multicast` — Copy tensor from global memory to shared memory using TMA with multicast
- `tma_store` — Copy tensor from shared memory to global memory using TMA
- `copy` — Copy tensor from src to dst using a given copy atom
## Algorithm APIs (`algorithm.py`)
- `simt_auto_vec_copy` — Copies a tensor between buffers with single thread (auto-vectorized)
- `partition` — Partition a buffer into a given layout and tiler
- `partition_and_copy` — Combines partitioning and copying in a single operation
## Math APIs (`math.py`)
- `dot` — Computes a dot product of two tensors using an MMA atom
- `dot_block_scaled` — Computes a block-scaled dot product with scale factors
## Pipeline Classes (`pipeline.py`)
- `GenericPipeline` — Generic pipeline for any producer/consumer combination
- `TMAToUMMAPipeline` — Pipeline for TMA load to UMMA consumption
- `TMAToAsyncPipeline` — Pipeline for TMA load to async consumer
- `AsyncToUMMAPipeline` — Pipeline for async producer to UMMA consumption
- `UMMAtoAsyncPipeline` — Pipeline for UMMA producer to async consumer
- `TMAStorePipeline` — Pipeline for SMEM producer to TMA store consumer
## Utilities (`utils.py`)
- `get_cta_v_map_ab` — Compute CTA-V map for A/B operands
- `get_cta_v_map_c` — Compute CTA-V map for C operand
+12 -3
View File
@@ -9,6 +9,15 @@
# and related documentation outside the scope permitted by the EULA
# is strictly prohibited.
raise NotImplementedError(
"CuTe Experimental module is only supported on Cuda toolkit 13.1 and above!"
)
from ... import cutlass_dsl as _dsl
jit = _dsl.CuteExperimentalDSL.jit
kernel = _dsl.CuteExperimentalDSL.kernel
compile = _dsl.CompileCallable()
from .algorithm import *
from .core import *
from .math import *
from .memory import *
from .pipeline import *
from .utils import *
@@ -0,0 +1,150 @@
# SPDX-FileCopyrightText: Copyright (c) 2025 - 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: LicenseRef-NvidiaProprietary
#
# Use of this software is governed by the terms and conditions of the
# NVIDIA End User License Agreement (EULA), available at:
# https://docs.nvidia.com/cutlass/media/docs/pythonDSL/license.html
#
# Any use, reproduction, disclosure, or distribution of this software
# and related documentation outside the scope permitted by the EULA
# is strictly prohibited.
from cutlass import cute
from cutlass.cutlass_dsl import dsl_user_op
from cutlass._mlir.dialects import lir as cutlass_lir
from .memory import copy
@dsl_user_op
def simt_auto_vec_copy(
src: cute.Tensor, dst: cute.Tensor, async_op=False, loc=None, ip=None
):
"""
Copies a tensor between two cute.memref buffers with single thread.
:param src: Source tensor
:type src: cute.Tensor
:param dst: Destination tensor
:type dst: cute.Tensor
:param async_op: Whether to use asynchronous operation, defaults to False
:type async_op: bool, optional
"""
if async_op:
cutlass_lir.SimtAutoVecCopyOp(
src.value, dst.value, async_=True, cache="always", loc=loc, ip=ip
)
else:
cutlass_lir.SimtAutoVecCopyOp(src.value, dst.value, loc=loc, ip=ip)
@dsl_user_op
def partition(
buffer: cute.Tensor, agent_id: cute.Int32, *, layout_tv, tiler, loc=None, ip=None
) -> cute.Tensor:
"""
Partition a buffer into a given layout and tiler.
:param buffer: Buffer to partition
:type buffer: cute.Tensor
:param agent_id: Agent ID
:type agent_id: cute.Int32
:param layout_tv: Layout tensor
:type layout_tv: cute.Tensor
:param tiler: Tiler
:type tiler: cute.Tensor
"""
assert isinstance(agent_id, cute.Int32), (
f"Expected agent_id to be cute.Int32, got {type(agent_id)}"
)
partition_op = cutlass_lir.PartitionOp(
buffer.value,
agent_id.ir_value(),
layout_tv=layout_tv.type.attribute,
tiler=tiler.type.attribute,
loc=loc,
ip=ip,
)
return partition_op.result
@dsl_user_op
def partition_and_copy(
tiled_copy: cute.core.ThrCopy,
src: cute.Tensor,
dst: cute.Tensor,
*,
loc=None,
ip=None,
):
"""
Copies a tensor between two cute.memref buffer
:param tiled_copy: Tiled copy
:type tiled_copy: cute.core.ThrCopy
:param src: Source tensor
:type src: cute.Tensor
:param dst: Destination tensor
:type dst: cute.Tensor
"""
src_partitioned = src
dst_partitioned = dst
tid_x = tiled_copy.thr_idx
if src.memspace != cute.AddressSpace.rmem:
src_partitioned = partition(
src,
tid_x,
layout_tv=tiled_copy.layout_src_tv_tiled,
tiler=cute.core._pack_tile(tiled_copy.tiler_mn),
)
if dst.memspace != cute.AddressSpace.rmem:
dst_partitioned = partition(
dst,
tid_x,
layout_tv=tiled_copy.layout_dst_tv_tiled,
tiler=cute.core._pack_tile(tiled_copy.tiler_mn),
)
# Handle copy where copy atom is used for both partition and copy during smem to rmem and rmem to smem copies
if type(tiled_copy.op) in [
cute.nvgpu.warp.LdMatrix8x8x16bOp,
cute.nvgpu.warp.LdMatrix16x16x8bOp,
cute.nvgpu.warp.StMatrix8x8x16bOp,
cute.nvgpu.warp.StMatrix16x8x8bOp,
]:
copy(
src_partitioned,
dst_partitioned,
copy_atom=tiled_copy,
loc=loc,
ip=ip,
)
# The rest handles copy where copy atom is used for partition
elif (
src.memspace,
dst.memspace,
) in [
(cute.AddressSpace.rmem, cute.AddressSpace.smem),
(cute.AddressSpace.smem, cute.AddressSpace.rmem),
(cute.AddressSpace.rmem, cute.AddressSpace.gmem),
(cute.AddressSpace.gmem, cute.AddressSpace.rmem),
]:
simt_auto_vec_copy(src_partitioned, dst_partitioned, loc=loc, ip=ip)
elif (
src.memspace == cute.AddressSpace.gmem
and dst.memspace == cute.AddressSpace.smem
):
simt_auto_vec_copy(
src_partitioned, dst_partitioned, async_op=True, loc=loc, ip=ip
)
# Handle copy where copy atom is used for partition and copy
else:
copy(
src_partitioned,
dst_partitioned,
copy_atom=tiled_copy,
loc=loc,
ip=ip,
)
@@ -0,0 +1,245 @@
# SPDX-FileCopyrightText: Copyright (c) 2025 - 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: LicenseRef-NvidiaProprietary
#
# NVIDIA CORPORATION, its affiliates and licensors retain all intellectual
# property and proprietary rights in and to this material, related
# documentation and any modifications thereto. Any use, reproduction,
# disclosure or distribution of this material and related documentation
# without an express license agreement from NVIDIA CORPORATION or
# its affiliates is strictly prohibited.
from cutlass.cutlass_dsl import dsl_user_op
from cutlass._mlir.dialects import lir as cutlass_lir_ir, nvvm as _nvvm
from cutlass._mlir import ir
from cutlass.cutlass_dsl import lru_cache_ir
from cutlass._mlir.dialects.core import OperationTypeEnum
from cutlass import cute
@dsl_user_op
def elect_sync(loc=None, ip=None):
"""
Elects one predicated thread within a warp.
"""
return _nvvm.elect_sync(loc=loc, ip=ip)
@dsl_user_op
def get_mbarrier(stage_token, loc=None, ip=None):
"""
Returns the mbarrier pointer for a given stage token.
"""
return cutlass_lir_ir.GetMbarrierOp(stage_token, loc=loc, ip=ip)
@ir.register_value_caster(cutlass_lir_ir.PipelineStateType.get_static_typeid())
class PipelineState(ir.Value):
def __init__(self, value):
if isinstance(value, ir.Value):
self.value = value
else:
raise TypeError(f"Expected ir.Value, got {type(value)}")
super().__init__(value)
@property
@lru_cache_ir()
def type(self) -> ir.Type:
return self.value.type
@classmethod
def __new_from_mlir_values__(cls, values):
assert len(values) == 1, f"Expected 1 value, but got {len(values)}"
return PipelineState(values[0])
@dsl_user_op
def create_pipeline(
stage: cute.Int32,
producer: OperationTypeEnum,
consumer: OperationTypeEnum,
producer_arv_count: cute.Int32,
consumer_arv_count: cute.Int32,
loc=None,
ip=None,
) -> tuple[PipelineState, PipelineState, PipelineState]:
"""
Creates an abstraction for a circular buffer of synchronizatoin primitives
indexed by stage count.
:param stage: Stage count
:type stage: cute.Int32
:param producer: Producer operation type
:type producer: OperationTypeEnum
:param consumer: Consumer operation type
:type consumer: OperationTypeEnum
:param producer_arv_count: Producer arrival count
:type producer_arv_count: cute.Int32
:param consumer_arv_count: Consumer arrival count
:type consumer_arv_count: cute.Int32
"""
if isinstance(producer_arv_count, int):
producer_arv_count = cute.Int32(producer_arv_count)
if isinstance(consumer_arv_count, int):
consumer_arv_count = cute.Int32(consumer_arv_count)
result = ir.Type.parse(f"!lir.pipeline<{stage}, {producer} -> {consumer}>")
op = cutlass_lir_ir.CreatePipelineOp(
result,
producer_arv_count.ir_value(),
consumer_arv_count.ir_value(),
loc=loc,
ip=ip,
)
pipeline = op.result
result = ir.Type.parse(f"!lir.pipeline_state<{stage}>")
op = cutlass_lir_ir.CreatePipelineStateOp(result, pipeline, loc=loc, ip=ip)
producer_state = op.result
result = ir.Type.parse(f"!lir.pipeline_state<{stage}>")
op = cutlass_lir_ir.CreatePipelineStateOp(result, pipeline, loc=loc, ip=ip)
consumer_state = op.result
return pipeline, producer_state, consumer_state
@dsl_user_op
def create_pipeline_with_mask(
stage: cute.Int32,
producer: OperationTypeEnum,
consumer: OperationTypeEnum,
producer_arv_count: cute.Int32,
consumer_arv_count: cute.Int32,
arrival_mask: cute.Int16,
loc=None,
ip=None,
) -> tuple[PipelineState, PipelineState, PipelineState]:
"""
Creates a pipeline with an arrival mask for cluster-scoped synchronization.
:param stage: Pipeline stage count.
:param producer: Producer operation type (e.g. SM90_TMA_LOAD_MULTICAST).
:param consumer: Consumer operation type (e.g. SM100_MMA_2SM_SS).
:param producer_arv_count: Producer arrival count for the pipeline barriers.
:param consumer_arv_count: Consumer arrival count for the pipeline barriers.
:param arrival_mask: Bitmask that selects participating peers (e.g. CTAs in a
cluster). This is attached to the pipeline value and is consulted by some
pipeline lowerings to generate cluster-scoped synchronization
"""
if isinstance(producer_arv_count, int):
producer_arv_count = cute.Int32(producer_arv_count)
if isinstance(consumer_arv_count, int):
consumer_arv_count = cute.Int32(consumer_arv_count)
if isinstance(arrival_mask, int):
arrival_mask = cute.Int16(arrival_mask)
result = ir.Type.parse(f"!lir.pipeline<{stage}, {producer} -> {consumer}>")
op = cutlass_lir_ir.CreatePipelineWithMaskOp(
result,
producer_arv_count.ir_value(),
consumer_arv_count.ir_value(),
arrival_mask.ir_value(),
loc=loc,
ip=ip,
)
pipeline = op.result
result = ir.Type.parse(f"!lir.pipeline_state<{stage}>")
op = cutlass_lir_ir.CreatePipelineStateOp(result, pipeline, loc=loc, ip=ip)
producer_state = op.result
result = ir.Type.parse(f"!lir.pipeline_state<{stage}>")
op = cutlass_lir_ir.CreatePipelineStateOp(result, pipeline, loc=loc, ip=ip)
consumer_state = op.result
return pipeline, producer_state, consumer_state
@dsl_user_op
def pipeline_advance_iterator(pipe, state, loc=None, ip=None):
"""
Advances a pipeline iterator to the next stage.
"""
op = cutlass_lir_ir.PipelineAdvanceIteratorOp(pipe, state, loc=loc, ip=ip)
return op.result
@dsl_user_op
def producer_acquire(pipe, state, loc=None, ip=None):
"""
Acquires exclusive access to a pipeline.
"""
op = cutlass_lir_ir.ProducerAcquireOp(pipe, state, loc=loc, ip=ip)
return op.result
@dsl_user_op
def producer_commit(pipe, state, loc=None, ip=None):
"""
Commits results to a pipeline.
"""
op = cutlass_lir_ir.ProducerCommitOp(pipe, state, loc=loc, ip=ip)
return op.result
@dsl_user_op
def consumer_wait(pipe, state, loc=None, ip=None):
"""
Waits for a pipeline to transition to `full`.
"""
op = cutlass_lir_ir.ConsumerWaitOp(pipe, state, loc=loc, ip=ip)
return op.result
@dsl_user_op
def consumer_release(pipe, state, loc=None, ip=None):
"""
Releases a pipeline that has been consumed.
"""
op = cutlass_lir_ir.ConsumerReleaseOp(pipe, state, loc=loc, ip=ip)
return op.result
@dsl_user_op
def consumer_tail(pipe, state, loc=None, ip=None):
"""
Called by the consumer to block until asynchronous tasks have completed.
"""
op = cutlass_lir_ir.ConsumerTailOp(pipe, state, loc=loc, ip=ip)
return op.result
@dsl_user_op
def get_pipeline_produce_stage(pipeline, state, loc=None, ip=None):
"""
Gets a pipeline produce stage.
"""
stage_token_type = ir.Type.parse(f"!lir.stage_token<{pipeline.type}>")
stage_idx = ir.IntegerType.get_signless(32)
op = cutlass_lir_ir.GetPipelineProduceStageOp(
stage_token=stage_token_type,
stage_index=stage_idx,
pipeline=pipeline,
pipelineState=state,
loc=loc,
ip=ip,
)
return op.stage_token, op.stage_index
@dsl_user_op
def get_pipeline_consume_stage(pipeline, state, loc=None, ip=None):
"""
Creates a pipeline consume stage.
"""
stage_token_type = ir.Type.parse(f"!lir.stage_token<{pipeline.type}>")
stage_idx = ir.IntegerType.get_signless(32)
op = cutlass_lir_ir.GetPipelineConsumeStageOp(
stage_token=stage_token_type,
stage_index=stage_idx,
pipeline=pipeline,
pipelineState=state,
loc=loc,
ip=ip,
)
return op.stage_token, op.stage_index
@@ -0,0 +1,84 @@
# SPDX-FileCopyrightText: Copyright (c) 2025 - 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: LicenseRef-NvidiaProprietary
#
# Use of this software is governed by the terms and conditions of the
# NVIDIA End User License Agreement (EULA), available at:
# https://docs.nvidia.com/cutlass/media/docs/pythonDSL/license.html
#
# Any use, reproduction, disclosure, or distribution of this software
# and related documentation outside the scope permitted by the EULA
# is strictly prohibited.
from cutlass import cute
from cutlass.cutlass_dsl import dsl_user_op
from cutlass._mlir.dialects import lir as cutlass_lir
@dsl_user_op
def dot_block_scaled(
mma_atom: cute.MmaAtom,
a: cute.Tensor,
sfa: cute.Tensor,
b: cute.Tensor,
sfb: cute.Tensor,
c: cute.Tensor,
loc=None,
ip=None,
):
"""
Computes the dot product of two tensors with block scaling and accumulates the result into a third tensor.
:param mma_atom: MMA atom
:type mma_atom: cute.MmaAtom
:param a: First tensor
:type a: cute.Tensor
:param sfa: First scale factor tensor
:type sfa: cute.Tensor
:param b: Second tensor
:type b: cute.Tensor
:param sfb: Second scale factor tensor
:type sfb: cute.Tensor
:param c: Result tensor
:type c: cute.Tensor
"""
cutlass_lir.DotBlockScaledOp(
mma_atom._unpack(),
a.value,
sfa.value,
b.value,
sfb.value,
c.value,
loc=loc,
ip=ip,
)
@dsl_user_op
def dot(
mma_atom: cute.MmaAtom,
a: cute.Tensor,
b: cute.Tensor,
c: cute.Tensor,
loc=None,
ip=None,
):
"""
Computes the dot product of two tensors and accumulates the result into a third tensor.
:param mma_atom: MMA atom
:type mma_atom: cute.MmaAtom
:param a: First tensor
:type a: cute.Tensor
:param b: Second tensor
:type b: cute.Tensor
:param c: Result tensor
:type c: cute.Tensor
"""
cutlass_lir.DotOp(
mma_atom._unpack(),
a.value,
b.value,
c.value,
loc=loc,
ip=ip,
)
@@ -0,0 +1,256 @@
# SPDX-FileCopyrightText: Copyright (c) 2025 - 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: LicenseRef-NvidiaProprietary
#
# Use of this software is governed by the terms and conditions of the
# NVIDIA End User License Agreement (EULA), available at:
# https://docs.nvidia.com/cutlass/media/docs/pythonDSL/license.html
#
# Any use, reproduction, disclosure, or distribution of this software
# and related documentation outside the scope permitted by the EULA
# is strictly prohibited.
from typing import Type, Optional
from cutlass._mlir.dialects import cute_nvgpu as _cute_nvgpu_ir
from cutlass._mlir import ir
from cutlass.cutlass_dsl import dsl_user_op
from cutlass._mlir.dialects import (
lir as cutlass_lir,
cute as _cute_ir,
)
from cutlass._mlir.dialects.core import OperationTypeEnum
from cutlass import cute
def _get_tma_load_kind(tma_operation_type: OperationTypeEnum):
"""Convert OperationTypeEnum to TiledTmaLoadEnum."""
if tma_operation_type == OperationTypeEnum.SM100_TMA_LOAD_2SM_MULTICAST:
return _cute_ir.TiledTmaLoadEnum.sm_100_2sm_multicast
if tma_operation_type == OperationTypeEnum.SM90_TMA_LOAD_MULTICAST:
return _cute_ir.TiledTmaLoadEnum.sm_90_multicast
if tma_operation_type == OperationTypeEnum.SM100_TMA_LOAD_2SM:
return _cute_ir.TiledTmaLoadEnum.sm_100_2sm
if tma_operation_type == OperationTypeEnum.SM90_TMA_LOAD:
return _cute_ir.TiledTmaLoadEnum.sm_90
raise ValueError(f"Unsupported TMA operation type: {tma_operation_type}")
@dsl_user_op
def allocate(
type: Type[cute.Numeric],
address_space: cute.AddressSpace,
layout: cute.Layout | cute.ComposedLayout,
alignment: cute.Int32,
is2cta: bool = False,
loc=None,
ip=None,
) -> cute.Tensor:
"""
Allocate a buffer of the given type and layout.
:param type: The type of the buffer
:type type: cute.Tensor
:param layout: The layout of the buffer
:type layout: cute.Layout
:param address_space: The address space of the buffer
:type address_space: str
:param alignment: The alignment of the buffer
:type alignment: cute.Int32
:param is2cta: Whether TMEM allocation should span a CTA pair (2CTA TMEM)
:type is2cta: bool
"""
swizzle = None
if isinstance(layout, cute.ComposedLayout):
swizzle = layout.inner
layout = layout.outer
# Handle SparseElemType (pass through) vs regular types (get mlir_type)
if isinstance(type, _cute_ir.SparseElemType):
pass
else:
type = type.mlir_type
ptr_ty = _cute_ir.PtrType.get(
type,
address_space,
alignment,
swizzle.type.attribute if swizzle else None,
)
buffer_type = _cute_ir.MemRefType.get(ptr_ty, layout.type)
# `is2cta` is a UnitAttr flag in the IR:
# present => true, absent => false.
is2cta_attr = ir.UnitAttr.get() if is2cta else None
buffer_op = cutlass_lir.AllocateBufferOp(
buffer_type, is2cta=is2cta_attr, loc=loc, ip=ip
)
return buffer_op.result
@dsl_user_op
def tma_load(
src: cute.Tensor,
dst: cute.Tensor,
mbar,
*,
cta_v_map,
tma_operation_type: Optional[OperationTypeEnum] = None,
internal_type=None,
update_expect_tx: bool = True,
loc=None,
ip=None,
):
"""
Copies a tensor pointed by a !cute.memref into a Buffer using TMA.
update_expect_tx (bool): controls whether this operation increments the mbarrier's transaction bytes with the TMA copy size.
When used with Cute DSL pipelines, it must be set to False as the pipeline already initializes the mbarrier's transaction bytes.
tma_operation_type (optional): specifies the TMA operation type (SM90_TMA_LOAD, SM100_TMA_LOAD_2SM, etc.)
internal_type (optional): selects the TMA transfer's internal element encoding used by hardware.
Does not change src/dst memref types. For structured sparsity, use base storage types:
Float16 for 2:4 FP16 sparse element type, Uint8 for 8:1 uint8 sparse element type.
:param src: Source tensor in global memory
:type src: cute.Tensor
:param dst: Destination tensor in shared memory
:type dst: cute.Tensor
:param mbar: Memory barrier for synchronization
:type mbar: cute.core.Mbarrier
:param cta_v_map: CTA V-map for the tensor
:type cta_v_map: cute.core.CtaVMap
:param tma_operation_type: TMA operation type (e.g., SM90_TMA_LOAD, SM100_TMA_LOAD_2SM, etc.)
:type tma_operation_type: OperationTypeEnum
:param internal_type: Internal type of the TMA transfer
:type internal_type: cute.core.InternalType
:param update_expect_tx: Whether to update expected transaction bytes
:type update_expect_tx: bool
"""
if tma_operation_type is not None:
kind = _get_tma_load_kind(tma_operation_type)
else:
kind = _cute_ir.TiledTmaLoadEnum.sm_90
kwargs = {
"cta_v_map": cta_v_map.type.attribute,
"kind": kind,
"loc": loc,
"ip": ip,
}
# Map internal_type to tma_format per updated API
if internal_type is not None:
internal_mlir_ty = (
internal_type.mlir_type
if hasattr(internal_type, "mlir_type")
else internal_type
)
kwargs["tma_format"] = _cute_nvgpu_ir.TmaDataFormat(
_cute_nvgpu_ir.get_default_tma_format(internal_mlir_ty, False)
)
if update_expect_tx:
kwargs["update_expect_tx"] = True
cutlass_lir.TmaLoadOp(src.value, dst.value, mbar, **kwargs)
@dsl_user_op
def tma_load_multicast(
src: cute.Tensor,
dst: cute.Tensor,
mbar,
*,
vmnk_layout: cute.Layout,
cta_v_map,
tma_operation_type: OperationTypeEnum,
multicast_mode: int,
update_expect_tx: bool = True,
loc=None,
ip=None,
):
"""
Copies a tensor pointed by a !cute.memref into a Buffer using TMA with multicast.
:param src: Source tensor in global memory
:param dst: Destination tensor in shared memory
:param mbar: Memory barrier for synchronization
:param vmnk_layout: Layout describing the cluster configuration
:param cta_v_map: CTA V-map for the tensor
:param tma_operation_type: TMA operation type (e.g., SM90_TMA_LOAD_MULTICAST, SM100_TMA_LOAD_2SM_MULTICAST)
:param multicast_mode: Multicast projection mode (1=column, 2=row)
:param update_expect_tx: Whether to update expected transaction bytes
"""
kind = _get_tma_load_kind(tma_operation_type)
kwargs = {
"cta_v_map": cta_v_map.type.attribute,
"kind": kind,
"vmnk_layout": vmnk_layout,
"multicast_mode": multicast_mode,
"loc": loc,
"ip": ip,
}
if update_expect_tx:
kwargs["update_expect_tx"] = True
cutlass_lir.TmaLoadMulticastOp(
src.value,
dst.value,
mbar,
**kwargs,
)
@dsl_user_op
def tma_store(
src: cute.Tensor,
dst: cute.Tensor,
*,
cta_v_map,
internal_type=None,
loc=None,
ip=None,
):
"""
Copies a tensor from a Buffer to a tensor pointed to by a !cute.memref.
internal_type (optional): selects the TMA transfer's internal element encoding used by hardware.
Does not change src/dst memref types. For structured sparsity, use base storage types:
Float16 for 2:4 FP16 sparse element type, Uint8 for 8:1 uint8 sparse element type.
:param src: Source tensor in shared memory
:type src: cute.Tensor
:param dst: Destination tensor in global memory
:type dst: cute.Tensor
:param cta_v_map: CTA V-map for the tensor
:type cta_v_map: cute.core.CtaVMap
:param internal_type: Internal type of the TMA transfer
:type internal_type: cute.core.InternalType
"""
kwargs = {
"cta_v_map": cta_v_map.type.attribute,
"loc": loc,
"ip": ip,
}
# Map internal_type to tma_format per updated API
if internal_type is not None:
internal_mlir_ty = (
internal_type.mlir_type
if hasattr(internal_type, "mlir_type")
else internal_type
)
kwargs["tma_format"] = _cute_nvgpu_ir.TmaDataFormat(
_cute_nvgpu_ir.get_default_tma_format(internal_mlir_ty, False)
)
cutlass_lir.TmaStoreOp(src.value, dst.value, **kwargs)
@dsl_user_op
def copy(src: cute.Tensor, dst: cute.Tensor, *, copy_atom, loc=None, ip=None):
"""
Copy a tensor from src to dst using a given copy atom.
"""
copy_atom = ir.Attribute.parse(f"{copy_atom.type}")
cutlass_lir.CopyOp(src.value, dst.value, copy_atom=copy_atom, loc=loc, ip=ip)
@@ -0,0 +1,684 @@
# SPDX-FileCopyrightText: Copyright (c) 2025 - 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: LicenseRef-NvidiaProprietary
#
# NVIDIA CORPORATION, its affiliates and licensors retain all intellectual
# property and proprietary rights in and to this material, related
# documentation and any modifications thereto. Any use, reproduction,
# disclosure or distribution of this material and related documentation
# without an express license agreement from NVIDIA CORPORATION or
# its affiliates is strictly prohibited.
"""
Convenience pipeline classes that hide elect_one synchronization complexity
"""
from dataclasses import dataclass
from typing import Optional
import cutlass
import cutlass.cute as cute
from cutlass._mlir.dialects import lir as cutlass_lir_ir
from cutlass.base_dsl.typing import Int32
from cutlass._mlir.dialects.core import OperationTypeEnum
from cutlass.cute.experimental.core import (
create_pipeline,
create_pipeline_with_mask,
producer_acquire,
get_pipeline_produce_stage,
get_pipeline_consume_stage,
producer_commit,
consumer_release,
pipeline_advance_iterator,
consumer_wait,
consumer_tail,
)
from cutlass.cutlass_dsl import CuteExperimentalDSL
class GenericPipelineBase:
"""Base class for pipeline convenience wrappers"""
def __init__(
self,
raw_pipeline,
num_stages,
producer_state,
consumer_state,
):
self.raw_pipeline = raw_pipeline
self.num_stages = num_stages
# For convenience class, we always manage state internally
self.producer_state = producer_state
self.consumer_state = consumer_state
def __extract_mlir_values__(self):
"""Extract MLIR values for DynamicExpression protocol."""
# raw_pipeline is always ir.OpResult from create_pipeline (no __extract_mlir_values__)
pipeline_values = [self.raw_pipeline]
# Create DSL types and extract their underlying MLIR values
num_stages_dsl = Int32(self.num_stages)
# Pipeline states are already MLIR values (PipelineState objects)
producer_state_values = [self.producer_state]
consumer_state_values = [self.consumer_state]
return (
pipeline_values
+ [
num_stages_dsl.__extract_mlir_values__()[0],
]
+ producer_state_values
+ consumer_state_values
)
@classmethod
def __new_from_mlir_values__(cls, values):
"""Reconstruct object from MLIR values."""
# Parse the known structure: [pipeline] + [num_stages, producer_flag, consumer_flag] + [producer_state] + [consumer_state]
# All lir_* objects are single MLIR values
raw_pipeline = values[0] # Always single ir.OpResult
num_stages_val = values[1]
producer_state = values[2] # Always single PipelineState
consumer_state = values[3] # Always single PipelineState
# Create temporary DSL objects and extract Python values
temp_num_stages = Int32(0)
num_stages_dsl = temp_num_stages.__new_from_mlir_values__([num_stages_val])
return cls(
raw_pipeline,
(
num_stages_dsl.value
if hasattr(num_stages_dsl, "value")
else int(num_stages_dsl)
),
producer_state,
consumer_state,
)
def producer_acquire(self):
"""Acquire producer state."""
producer_acquire(self.raw_pipeline, self.producer_state)
return self
def get_producer_stage(self):
"""Get producer stage."""
return get_pipeline_produce_stage(self.raw_pipeline, self.producer_state)
def get_consumer_stage(self):
"""Get consumer stage."""
return get_pipeline_consume_stage(self.raw_pipeline, self.consumer_state)
# Instance methods that can now be used directly in kernel context
def producer_acquire_and_get_stage(self):
"""Combined producer acquire + get_stage with automatic elect_one using internal state."""
self.producer_acquire()
return get_pipeline_produce_stage(self.raw_pipeline, self.producer_state)
def producer_commit(self):
"""Commit producer state."""
producer_commit(self.raw_pipeline, self.producer_state)
return self
def consumer_release(self):
"""Release consumer state."""
consumer_release(self.raw_pipeline, self.consumer_state)
return self
def producer_commit_and_advance(self):
"""Combined producer commit + advance with automatic elect_one using internal state."""
self.producer_commit()
# Update internal state in-place for better performance
self.producer_state = pipeline_advance_iterator(
self.raw_pipeline, self.producer_state
)
return self
def consumer_wait_and_get_stage(self):
"""Combined consumer wait + get_stage with automatic elect_one using internal state."""
self.consumer_wait()
return get_pipeline_consume_stage(self.raw_pipeline, self.consumer_state)
def consumer_wait(self):
"""Wait for consumer to be ready."""
consumer_wait(self.raw_pipeline, self.consumer_state)
return self
def consumer_release_and_advance(self):
"""Combined consumer release + advance with automatic elect_one using internal state."""
self.consumer_release()
# Update internal state in-place for better performance
self.consumer_state = pipeline_advance_iterator(
self.raw_pipeline, self.consumer_state
)
return self
def consumer_tail(self):
"""Combined consumer tail with automatic elect_one using internal state."""
consumer_tail(self.raw_pipeline, self.consumer_state)
return self
class GenericPipeline(GenericPipelineBase):
"""
Generic pipeline for any combination of producer and consumer.
"""
@staticmethod
def create(
*,
producer: OperationTypeEnum,
consumer: OperationTypeEnum,
producer_arv_count: cute.Int32,
consumer_arv_count: cute.Int32,
num_stages: cute.Int32,
):
"""
Create a generic pipeline with parameterized producer and consumer.
Args:
producer: Producer operation type
consumer: Consumer operation type
producer_arv_count: Producer arrival count
consumer_arv_count: Consumer arrival count
num_stages: Number of pipeline stages
"""
raw_pipeline, producer_state, consumer_state = create_pipeline(
num_stages,
producer,
consumer,
producer_arv_count=producer_arv_count,
consumer_arv_count=consumer_arv_count,
)
return GenericPipeline(
raw_pipeline,
num_stages,
producer_state,
consumer_state,
)
def _validate_umma_operation_type(operation_type: OperationTypeEnum):
if operation_type not in [
OperationTypeEnum.SM100_MMA_1SM_SS,
OperationTypeEnum.SM100_MMA_1SM_TS,
OperationTypeEnum.SM100_MMA_2SM_SS,
OperationTypeEnum.SM100_MMA_2SM_TS,
OperationTypeEnum.SM100_MMA_SCALED_1SM_SS,
OperationTypeEnum.SM100_MMA_SCALED_1SM_TS,
OperationTypeEnum.SM100_MMA_SCALED_2SM_SS,
OperationTypeEnum.SM100_MMA_SCALED_2SM_TS,
]:
raise ValueError(f"Invalid UMMA operation type: {operation_type}")
def _is_2sm_umma_operation_type(operation_type: OperationTypeEnum) -> bool:
"""Check if the operation type is a 2SM UMMA operation."""
return operation_type in [
OperationTypeEnum.SM100_MMA_2SM_SS,
OperationTypeEnum.SM100_MMA_2SM_TS,
OperationTypeEnum.SM100_MMA_SCALED_2SM_SS,
OperationTypeEnum.SM100_MMA_SCALED_2SM_TS,
]
class TMAToUMMAPipeline(GenericPipelineBase):
"""
Pipeline for TMA to UMMA.
"""
@staticmethod
def create(
*,
num_stages: cute.Int32,
mma_operation_type: OperationTypeEnum,
tma_operation_type: Optional[OperationTypeEnum] = None,
cluster_layout_vmnk: Optional[cute.Layout] = None,
):
"""
Create a TMA to UMMA pipeline.
For 2SM MMA with TMA_LOAD_2SM, provide cluster_layout_vmnk for proper mask computation.
"""
_validate_umma_operation_type(
mma_operation_type,
)
# Default to SM90_TMA_LOAD if not specified
if tma_operation_type is None:
tma_operation_type = OperationTypeEnum.SM90_TMA_LOAD
if tma_operation_type == OperationTypeEnum.SM100_TMA_LOAD_2SM:
if cluster_layout_vmnk is None:
raise ValueError(
"cluster_layout_vmnk is required if using 2CTA MMA with TMA"
)
# If using 2CTA MMA, need consumer_mask == local_cta | peer_cta
cta_rank_in_cluster = cute.arch.make_warp_uniform(
cute.arch.block_idx_in_cluster()
)
cta_in_cluster_coord_vmnk = cluster_layout_vmnk.get_flat_coord(
cta_rank_in_cluster
)
arrival_mask = cute.make_layout_image_mask(
cluster_layout_vmnk, cta_in_cluster_coord_vmnk, mode=0
)
raw_pipeline, producer_state, consumer_state = create_pipeline_with_mask(
num_stages,
tma_operation_type,
mma_operation_type,
producer_arv_count=1,
consumer_arv_count=1,
arrival_mask=arrival_mask,
)
else:
raw_pipeline, producer_state, consumer_state = create_pipeline(
num_stages,
tma_operation_type,
mma_operation_type,
producer_arv_count=1,
consumer_arv_count=1,
)
return TMAToUMMAPipeline(
raw_pipeline,
num_stages,
producer_state,
consumer_state,
)
@staticmethod
def create_with_mask(
*,
num_stages: cute.Int32,
tma_operation_type: OperationTypeEnum,
mma_operation_type: OperationTypeEnum,
cluster_layout_vmnk: cute.Layout,
):
"""
Create a TMA to UMMA pipeline with multicast mask for 2CTA operations.
"""
_validate_umma_operation_type(
mma_operation_type,
)
# Calculate TMA multicasting masks
tma_mcast_proj_A = 2 # multicast across CTAs in same row
tma_mcast_proj_B = 1 # multicast across CTAs in same column
cta_rank_in_cluster = cute.arch.make_warp_uniform(
cute.arch.block_idx_in_cluster()
)
cta_in_cluster_coord_vmnk = cluster_layout_vmnk.get_flat_coord(
cta_rank_in_cluster
)
# For 2CTA MMA (v-size==2), the peer CTA is the other v-slice (xor 1).
# For 1CTA MMA (v-size==1), the peer is the local CTA (no flip).
v_size = cute.size(cluster_layout_vmnk.shape[0])
peer_v = (
(cta_in_cluster_coord_vmnk[0] ^ 1)
if cutlass.const_expr(v_size > 1)
else cta_in_cluster_coord_vmnk[0]
)
cta_in_cluster_coord_vmnk_peer = (
peer_v,
*cta_in_cluster_coord_vmnk[1:],
)
arrival_mask_a = cute.nvgpu.cpasync.create_tma_multicast_mask(
cluster_layout_vmnk, cta_in_cluster_coord_vmnk, tma_mcast_proj_A
)
arrival_mask_b = cute.nvgpu.cpasync.create_tma_multicast_mask(
cluster_layout_vmnk, cta_in_cluster_coord_vmnk, tma_mcast_proj_B
)
arrival_mask_a_peer = cute.nvgpu.cpasync.create_tma_multicast_mask(
cluster_layout_vmnk,
cta_in_cluster_coord_vmnk_peer,
mcast_mode=tma_mcast_proj_A,
)
arrival_mask_b_peer = cute.nvgpu.cpasync.create_tma_multicast_mask(
cluster_layout_vmnk,
cta_in_cluster_coord_vmnk_peer,
mcast_mode=tma_mcast_proj_B,
)
# if 1SM MMA, arrival_mask_a_peer==arrival_mask_a && arrival_mask_b==arrival_mask_b_peer
arrival_mask_c = (
arrival_mask_a | arrival_mask_a_peer | arrival_mask_b | arrival_mask_b_peer
)
num_mcast_ctas_a = cute.size(cluster_layout_vmnk.shape[2])
num_mcast_ctas_b = cute.size(cluster_layout_vmnk.shape[1])
num_mcast_participants = num_mcast_ctas_a + num_mcast_ctas_b - 1
raw_pipeline, producer_state, consumer_state = create_pipeline_with_mask(
num_stages,
tma_operation_type,
mma_operation_type,
producer_arv_count=1,
consumer_arv_count=num_mcast_participants,
arrival_mask=arrival_mask_c,
)
return TMAToUMMAPipeline(
raw_pipeline, num_stages, producer_state, consumer_state
)
def producer_commit(self):
"""Commit producer state."""
with cute.arch.elect_one():
super().producer_commit()
return self
def consumer_release(self):
"""Release consumer state."""
with cute.arch.elect_one():
super().consumer_release()
return self
class TMAToAsyncPipeline(GenericPipelineBase):
"""
Pipeline for TMA to * (except UMMA).
"""
@staticmethod
def create(
*,
num_stages: cute.Int32,
consumer: OperationTypeEnum,
consumer_arv_count: cute.Int32,
):
"""
Create a TMA to * (except UMMA) pipeline.
"""
raw_pipeline, producer_state, consumer_state = create_pipeline(
num_stages,
OperationTypeEnum.SM90_TMA_LOAD,
consumer,
producer_arv_count=1,
consumer_arv_count=consumer_arv_count,
)
return TMAToAsyncPipeline(
raw_pipeline,
num_stages,
producer_state,
consumer_state,
)
def producer_commit(self):
"""Commit producer state."""
with cute.arch.elect_one():
super().producer_commit()
return self
class AsyncToUMMAPipeline(GenericPipelineBase):
"""
Pipeline for * (except TMA) to UMMA.
"""
@staticmethod
def create(
*,
num_stages: cute.Int32,
producer: OperationTypeEnum,
producer_arv_count: cute.Int32,
mma_operation_type: OperationTypeEnum,
):
"""
Create a * (except TMA) to UMMA pipeline.
"""
_validate_umma_operation_type(
mma_operation_type,
)
if producer == OperationTypeEnum.SM90_TMA_LOAD:
raise ValueError("TMA to UMMA is not supported.")
raw_pipeline, producer_state, consumer_state = create_pipeline(
num_stages,
producer,
mma_operation_type,
producer_arv_count=producer_arv_count,
consumer_arv_count=1,
)
return AsyncToUMMAPipeline(
raw_pipeline,
num_stages,
producer_state,
consumer_state,
)
def consumer_release(self):
"""Release consumer state."""
with cute.arch.elect_one():
super().consumer_release()
return self
class UMMAtoAsyncPipeline(GenericPipelineBase):
"""
Pipeline for UMMA to * (except TMA).
"""
@staticmethod
def create(
*,
num_stages: cute.Int32,
consumer: OperationTypeEnum,
consumer_arv_count: cute.Int32,
mma_operation_type: OperationTypeEnum,
cluster_layout_vmnk: Optional[cute.Layout] = None,
):
"""
Create a UMMA to * (except TMA) pipeline.
For 2SM MMA, provide cluster_layout_vmnk for proper mask computation.
"""
_validate_umma_operation_type(
mma_operation_type,
)
if consumer == OperationTypeEnum.SM90_TMA_LOAD:
raise ValueError("UMMA to TMA is not supported.")
if _is_2sm_umma_operation_type(mma_operation_type):
if cluster_layout_vmnk is None:
raise ValueError("cluster_layout_vmnk cannot be None if using 2SM MMA")
return UMMAtoAsyncPipeline.create_with_mask(
num_stages=num_stages,
consumer_type=consumer,
consumer_arv_count=consumer_arv_count,
mma_operation_type=mma_operation_type,
cluster_layout_vmnk=cluster_layout_vmnk,
)
else: # 1SM MMA
raw_pipeline, producer_state, consumer_state = create_pipeline(
num_stages,
mma_operation_type,
consumer,
producer_arv_count=1,
consumer_arv_count=consumer_arv_count,
)
return UMMAtoAsyncPipeline(
raw_pipeline,
num_stages,
producer_state,
consumer_state,
)
@staticmethod
def create_with_mask(
*,
num_stages: cute.Int32,
consumer_type: OperationTypeEnum,
consumer_arv_count: cute.Int32,
mma_operation_type: OperationTypeEnum,
cluster_layout_vmnk: cute.Layout,
):
"""
Create a UMMA to * pipeline with arrival mask for 2CTA operations.
"""
tmem_sync_mask = cutlass.pipeline.PipelineUmmaAsync._compute_tmem_sync_mask(
cta_layout_vmnk=cluster_layout_vmnk
)
raw_pipeline, producer_state, consumer_state = create_pipeline_with_mask(
num_stages,
mma_operation_type,
consumer_type,
producer_arv_count=1,
consumer_arv_count=consumer_arv_count,
arrival_mask=tmem_sync_mask,
)
return UMMAtoAsyncPipeline(
raw_pipeline,
num_stages,
producer_state,
consumer_state,
)
def producer_commit(self):
"""Commit producer state."""
with cute.arch.elect_one():
super().producer_commit()
return self
@dataclass
class TMAStorePipeline:
"""
TMA Store Pipeline modeling SMEM producer to TMA consumer pipeline.
A number of epilogue warps participate in the pipeline as producers, and one of them is designated as the consumer to perform TMA store.
Named barrier is used to synchronize all warps so that producers write SMEM after the pipeline stage is available, and the consumer waits for all producers before issuing TMA store.
The canonical pipeline flow is:
1. acquire_sync(): wait for pipeline stage availability + barrier
2. Each producer performs SMEM writes
3. commit_sync(): fence SMEM writes + barrier
4. Consumer performs TMA store
5. release_advance(): commit TMA store + advance stage
Args:
stages: Number of pipeline stages (type parameter)
arv_count: Number of threads participating in barriers
barrier_id: Barrier ID for synchronization
tma_warp_id: Which warp issues TMA stores (None = no TMA operations)
index: Initial stage index
"""
stages: cutlass.Constexpr[int]
arv_count: int
barrier_id: int
tma_warp_id: int
index: int = 0
def get_num_stages(self):
return self.stages
def acquire_sync(self):
"""
Acquire pipeline stage and synchronize all warps.
TMA warp waits for previous TMA operation to the same stage to complete (allowing writes to other stages to be in flight).
All warps then synchronize before producers write to SMEM.
"""
@CuteExperimentalDSL.jit
def acquire_sync_impl():
# Only TMA warp needs to wait for bulk async operations
warp_idx = cute.arch.warp_idx()
warp_idx = cute.arch.make_warp_uniform(warp_idx)
# Use Python if with @Cutlass_LIR.jit preprocessor
if warp_idx == self.tma_warp_id:
# Allow N-1 TMA operations in flight for pipelining
# Now we can use the compile-time constant from type parameter
num_stages = self.get_num_stages()
wait_count = num_stages - 1 if num_stages > 1 else 0
cute.arch.cp_async_bulk_wait_group(wait_count, read=True)
# All warps must synchronize before producers write to SMEM
self._barrier()
return self
return acquire_sync_impl()
def commit_sync(self):
"""
Fence SMEM writes and synchronize all warps.
All warps fence their SMEM writes to make them visible to consumer
All warps then synchronize before TMA store operation.
"""
# All warps fence their SMEM writes for TMA visibility
cute.arch.fence_proxy("async.shared", space="cta")
# All warps synchronize before TMA store
self._barrier()
return self
def release_advance(self):
"""
Release current stage and advance to next stage.
TMA warp commits the TMA store operations to a bulk group.
All warps advance to the next pipeline stage.
"""
@CuteExperimentalDSL.jit
def release_advance_impl():
# Only TMA warp commits the TMA operations
warp_idx = cute.arch.warp_idx()
warp_idx = cute.arch.make_warp_uniform(warp_idx)
# Use Python if with @Cutlass_LIR.jit preprocessor
if warp_idx == self.tma_warp_id:
cute.arch.cp_async_bulk_commit_group()
# All warps advance to next stage
self.index = (self.index + 1) % self.get_num_stages()
return self
return release_advance_impl()
def get_index(self):
"""Get current pipeline stage index."""
return self.index
def tail(self):
"""
Wait for all remaining TMA operations to complete.
Should be called at the end of the pipeline to ensure all TMA stores finish.
"""
@CuteExperimentalDSL.jit
def tail_impl():
warp_idx = cute.arch.warp_idx()
warp_idx = cute.arch.make_warp_uniform(warp_idx)
# Use Python if with @Cutlass_LIR.jit preprocessor
if warp_idx == self.tma_warp_id:
# Wait for all TMA operations to complete
cute.arch.cp_async_bulk_wait_group(0, read=True)
self._barrier()
return self
return tail_impl()
def _barrier(self):
"""Internal barrier synchronization."""
cute.arch.barrier(
barrier_id=self.barrier_id,
number_of_threads=self.arv_count,
)
@@ -0,0 +1,79 @@
# SPDX-FileCopyrightText: Copyright (c) 2025 - 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: LicenseRef-NvidiaProprietary
#
# Use of this software is governed by the terms and conditions of the
# NVIDIA End User License Agreement (EULA), available at:
# https://docs.nvidia.com/cutlass/media/docs/pythonDSL/license.html
#
# Any use, reproduction, disclosure, or distribution of this software
# and related documentation outside the scope permitted by the EULA
# is strictly prohibited.
from cutlass import cute
def get_cta_v_map_ab(
gmem_tensor,
mma_tiler_mnk,
tiled_mma,
input_operand,
*,
loc=None,
ip=None,
):
"""
Build the **CTA-to-value map** (aka **CTA V-map**) layout for a TMA load of A/B
(and scale-factor variants SFA/SFB).
In practice, `cta_v_map` is a `cute.Layout` that tells TMA how this CTA’s
portion of a global tensor tile maps onto the values being transferred into
shared memory.
:param gmem_tensor: Global-memory tensor being loaded by TMA.
:type gmem_tensor: cute.Tensor
:param mma_tiler_mnk: The (M,N,K,...) tiler describing the CTA tile shape.
:type mma_tiler_mnk: tuple
:param tiled_mma: The tiled MMA object used to derive the per-operand thread/value mapping.
:type tiled_mma: cute.core.TiledMma
:param input_operand: One of {"A","B","SFA","SFB"} selecting which operand mapping to use.
:type input_operand: str
:returns: A layout suitable to pass as `cta_v_map=...` to `tma_load` / `tma_load_multicast`.
:rtype: cute.Layout
"""
ident = cute.core.make_identity_layout(gmem_tensor.shape, loc=loc, ip=ip)
mode = 0 if (input_operand in ("A", "SFA")) else 1
mma_tiler_mk = (mma_tiler_mnk[mode], *mma_tiler_mnk[2:])
g_tile = cute.core.composition(ident, mma_tiler_mk, loc=loc, ip=ip)
if input_operand in ("A", "SFA"):
cta_v_map = tiled_mma._thrfrg_A(g_tile)
if input_operand in ("B", "SFB"):
cta_v_map = tiled_mma._thrfrg_B(g_tile)
cta_v_map = cute.core.get(cta_v_map, mode=[1])
cta_v_map = cute.core.dice(cta_v_map, (1, (1,) * cute.core.rank(g_tile)))
return cta_v_map
def get_cta_v_map_c(
gmem_tensor,
epi_tile,
*,
loc=None,
ip=None,
):
"""
Build the **CTA-to-value map** (aka **CTA V-map**) layout for a TMA store/load
of the output tensor C/D.
This returns an identity layout over the global tensor composed with the
epilogue tile, yielding a `cute.Layout` that describes which global indices
this CTA is responsible for.
:param gmem_tensor: Global-memory tensor being stored/loaded by TMA.
:type gmem_tensor: cute.Tensor
:param epi_tile: Epilogue tile layout describing the CTA’s output tile shape.
:type epi_tile: cute.Layout
:returns: A layout suitable to pass as `cta_v_map=...` to `tma_store` / `tma_load`.
:rtype: cute.Layout
"""
ident = cute.core.make_identity_layout(gmem_tensor.shape, loc=loc, ip=ip)
return cute.core.composition(ident, epi_tile, loc=loc, ip=ip)
@@ -169,4 +169,3 @@ Examples:
if __name__ == "__main__":
main()
-2
View File
@@ -16,8 +16,6 @@ from .tensor import TensorSSA
from cutlass._mlir.dialects import math, arith
from typing import Callable, Union
def _math_op(func: Callable, fastmath: bool, *args, **kwargs):
"""Dispatch the function to either a TensorSSA or a Numeric(Float).
+30 -3
View File
@@ -24,6 +24,7 @@ from ..typing import Float16, Float32, Float64, Numeric
__all__ = [
"OpError",
"normalize_field_to_ir_name",
"MmaUniversalOp",
"MmaUniversalTrait",
"CopyUniversalOp",
@@ -33,6 +34,33 @@ __all__ = [
"CacheEvictionPriority",
]
def normalize_field_to_ir_name(field, admissible_fields) -> str:
"""
Normalize a field specifier to its IR logical field name.
Accepted inputs:
- Enum value present in admissible_fields (must expose _to_ir_field_name()).
- Exact string IR name (e.g., "accum_c", "neg_a", "sf_a").
Any other form is rejected.
"""
# Enum path
if any(field is f for f in admissible_fields):
return field._to_ir_field_name()
# String path (must match exactly one of the IR names exposed by admissible_fields)
if isinstance(field, str):
allowed = {f._to_ir_field_name() for f in admissible_fields}
if field in allowed:
return field
# Otherwise, reject
allowed_pretty = [f._to_ir_field_name() for f in admissible_fields]
raise ValueError(
f"invalid field, must be one of {allowed_pretty} or their enum counterparts, but got {field}"
)
class OpError(DSLBaseError):
"""
An exception class for Op construction errors.
@@ -178,8 +206,8 @@ class CopyUniversalOp(atom.CopyOp):
op = cute.nvgpu.CopyUniversalOp()
atom = cute.make_copy_atom(
op,
tensor_dtype,
op,
tensor_dtype,
num_bits_per_copy=64,
l1c_evict_priority=cute.nvgpu.CacheEvictionPriority.EVICT_NORMAL
)
@@ -195,7 +223,6 @@ class CopyUniversalOp(atom.CopyOp):
- ``invariant`` is a kw argument specifying whether the load is invariant (read-only data \
that never changes). This enables compiler optimizations like instruction reordering. \
Defaults to ``False`` if not provided.
"""
def __str__(self) -> str:
@@ -24,6 +24,7 @@ __all__ = [
"CopyBulkTensorTileG2SMulticastOp",
"CopyBulkTensorTileS2GOp",
"CopyReduceBulkTensorTileS2GOp",
"CopyDsmemStoreOp",
#
# helpers.py
#
@@ -36,5 +37,4 @@ __all__ = [
"fence_tma_desc_acquire",
"cp_fence_tma_desc_release",
"fence_tma_desc_release",
"group_bulk_copy_modes",
]
@@ -21,7 +21,7 @@ from cutlass._mlir.dialects.cute import ReductionOp as ReductionOp
from cutlass._mlir import ir
from ...atom import CopyOp, Trait, make_atom
from ...typing import Int16, Int64, Pointer, Integer, Numeric
from ...typing import Int16, Int32, Int64, Pointer, Integer, Numeric
from ..common import OpError
from ..tcgen05.mma import CtaGroup
@@ -112,6 +112,7 @@ TMA_MBAR_PTR_FIELD_NAME = "tma_bar"
TMA_MCAST_MASK_FIELD_NAME = "mcast_mask"
TMA_DESC_PTR_FIELD_NAME = "tma_descriptor_ptr"
TMA_BYTE_MASK_FIELD_NAME = "byte_mask"
TMA_CTA_RANK_FIELD_NAME = "cta_rank"
TMA_CACHE_POLICY_FIELD_NAME = "cache_policy"
@@ -249,6 +250,7 @@ class CopyBulkTensorTileG2SNonExecTrait(Trait):
class CopyBulkTensorTileG2STrait(Trait):
pass
#
# TMA GMEM -> SMEM multicast copies
#
@@ -374,6 +376,7 @@ class CopyBulkTensorTileG2SMulticastNonExecTrait(Trait):
)
return exec_value
class CopyBulkTensorTileG2SMulticastTrait(Trait):
pass
@@ -457,10 +460,6 @@ class CopyBulkTensorTileS2GTrait(Trait):
pass
class CopyBulkTensorTileS2GTrait(Trait):
pass
@dataclass
class CopyReduceBulkTensorTileS2GOp(TmaCopyOp):
"""
@@ -800,7 +799,7 @@ class CopyBulkS2GByteMaskOp(CopyOp):
def __post_init__(self) -> None:
# Arch verification
arch: Arch = CuTeDSL._get_dsl().get_arch_enum()
arch: Arch = BaseDSL._get_dsl().get_arch_enum()
if not arch >= Arch.sm_100:
raise OpError(
self,
@@ -874,7 +873,7 @@ class CopyBulkS2SOp(CopyOp):
def __post_init__(self) -> None:
# Arch verification
arch: Arch = CuTeDSL._get_dsl().get_arch_enum()
arch: Arch = BaseDSL._get_dsl().get_arch_enum()
if not arch >= Arch.sm_90:
raise OpError(
self,
@@ -958,7 +957,7 @@ class CopyDsmemStoreOp(CopyOp):
def __post_init__(self) -> None:
# Arch verification
arch: Arch = CuTeDSL._get_dsl().get_arch_enum()
arch: Arch = BaseDSL._get_dsl().get_arch_enum()
if not arch >= Arch.sm_90:
raise OpError(
self,
@@ -984,6 +983,11 @@ class CopyDsmemStoreOp(CopyOp):
"expects a 'num_bits_per_copy' kw argument of type int that is non-negative "
f"when creating a copy Atom for {self.__class__.__name__}"
)
if num_bits_per_copy not in [0, 32, 64, 128]:
raise ValueError(
"expects a 'num_bits_per_copy' kw argument that is one of {0, 32, 64, 128} "
f"when creating a copy Atom for {self.__class__.__name__}"
)
ty = _cute_nvgpu_ir.CopyAtomDsmemStoreType.get(
copy_internal_type.mlir_type, num_bits_per_copy
)
@@ -10,7 +10,6 @@
# is strictly prohibited.
from typing import Optional, Tuple, Type, Union
from typing_extensions import deprecated
from cutlass.cutlass_dsl import dsl_user_op
@@ -47,11 +46,12 @@ TMAOp = Union[
CopyReduceBulkTensorTileS2GOp,
]
@dsl_user_op
def make_tiled_tma_atom(
op: TMAOp,
gmem_tensor: Tensor,
smem_layout: Union[Layout, ComposedLayout],
smem_layout_: Union[Layout, ComposedLayout],
cta_tiler: Tiler,
num_multicast: int = 1,
*,
@@ -84,7 +84,7 @@ def make_tiled_tma_atom(
:type op: TMAOp
:param gmem_tensor: The GMEM tensor involved in the Copy
:type gmem_tensor: Tensor
:param smem_layout: The SMEM layout to construct the Copy Atom
:param smem_layout: The SMEM layout to construct the Copy Atom, either w/ or w/o the stage mode
:type smem_layout: Union[Layout, ComposedLayout]
:param cta_tiler: The CTA Tiler to use
:type cta_tiler: Tiler
@@ -95,6 +95,26 @@ def make_tiled_tma_atom(
:return: A TMA Copy Atom associated with the TMA tensor
:rtype: Tuple[atom.CopyAtom, Tensor]
"""
smem_rank = core.rank(smem_layout_)
tiler_rank = core.rank(cta_tiler)
assert smem_rank == tiler_rank or smem_rank == tiler_rank + 1, (
f"smem_layout must be non-staged (rank(smem_layout) == rank(cta_tiler)) "
f"or staged (rank(smem_layout) == rank(cta_tiler) + 1)"
)
# Set the smem_layout on the operation for later retrieval
op.smem_layout = (
smem_layout_.value
if isinstance(smem_layout_, core._ComposedLayout)
else smem_layout_
)
# Slice the smem_layout if it is staged
if smem_rank == tiler_rank + 1:
smem_layout = core.select(smem_layout_, mode=list(range(tiler_rank)))
else:
smem_layout = smem_layout_
cta_v_map = core.composition(
core.make_identity_layout(gmem_tensor.shape, loc=loc, ip=ip),
cta_tiler,
@@ -105,22 +125,21 @@ def make_tiled_tma_atom(
if isinstance(smem_layout, core._ComposedLayout):
smem_layout = smem_layout.value
# Set the smem_layout on the operation for later retrieval
op.smem_layout = (
smem_layout.value
if isinstance(smem_layout, core._ComposedLayout)
else smem_layout
)
tma_format = None
if internal_type is not None:
if not isinstance(internal_type, NumericMeta):
raise TypeError(f"internal_type must be a Numeric, but got {internal_type}")
use_unpack = (internal_type.width == 8 and
isinstance(gmem_tensor.element_type, NumericMeta) and
gmem_tensor.element_type.width < 8)
internal_mlir_type = gmem_tensor.element_type.mlir_type if use_unpack else internal_type.mlir_type
use_unpack = (
internal_type.width == 8
and isinstance(gmem_tensor.element_type, NumericMeta)
and gmem_tensor.element_type.width < 8
)
internal_mlir_type = (
gmem_tensor.element_type.mlir_type
if use_unpack
else internal_type.mlir_type
)
tma_format = _cute_nvgpu_ir.TmaDataFormat(
_cute_nvgpu_ir.get_default_tma_format(internal_mlir_type, use_unpack)
)
@@ -380,14 +399,3 @@ def fence_tma_desc_release(*, loc=None, ip=None) -> None:
loc=loc,
ip=ip,
)
@dsl_user_op
@deprecated("`group_bulk_copy_modes` is deprecated, use `group_modes` instead")
def group_bulk_copy_modes(src: Tensor, dst: Tensor, loc=None, ip=None) -> Tuple:
"""
Copy async bulk need group mode 0, acquiring whole tensor for bulk copy
"""
mSrc = core.group_modes(src, 0, core.rank(src))
mDst = core.group_modes(dst, 0, core.rank(dst))
return (mSrc, mDst)
+20 -23
View File
@@ -17,7 +17,6 @@ import cutlass._mlir.dialects.cute_nvgpu as _cute_nvgpu_ir
from .. import core, atom
from ..typing import Shape, Layout, ComposedLayout, Tensor, Numeric, NumericMeta
from ...impl_utils import check_type_in
from .cpasync.copy import (
CopyBulkTensorTileG2SOp,
CopyBulkTensorTileG2SNonExecTrait,
@@ -96,13 +95,6 @@ def make_tiled_tma_atom_A(
"""
check_type_in(
op,
[CopyBulkTensorTileG2SOp, CopyBulkTensorTileG2SMulticastOp],
"op",
"make_tiled_tma_atom_A",
)
# Set the smem_layout on the operation for later retrieval
op.smem_layout = (
smem_layout.value
@@ -136,10 +128,16 @@ def make_tiled_tma_atom_A(
if not isinstance(internal_type, NumericMeta):
raise TypeError(f"internal_type must be a Numeric, but got {internal_type}")
use_unpack = (internal_type.width == 8 and
isinstance(gmem_tensor.element_type, NumericMeta) and
gmem_tensor.element_type.width < 8)
internal_mlir_type = gmem_tensor.element_type.mlir_type if use_unpack else internal_type.mlir_type
use_unpack = (
internal_type.width == 8
and isinstance(gmem_tensor.element_type, NumericMeta)
and gmem_tensor.element_type.width < 8
)
internal_mlir_type = (
gmem_tensor.element_type.mlir_type
if use_unpack
else internal_type.mlir_type
)
tma_format = _cute_nvgpu_ir.TmaDataFormat(
_cute_nvgpu_ir.get_default_tma_format(internal_mlir_type, use_unpack)
)
@@ -224,13 +222,6 @@ def make_tiled_tma_atom_B(
"""
check_type_in(
op,
[CopyBulkTensorTileG2SOp, CopyBulkTensorTileG2SMulticastOp],
"op",
"make_tiled_tma_atom_B",
)
# Set the smem_layout on the operation for later retrieval
op.smem_layout = (
smem_layout.value
@@ -264,10 +255,16 @@ def make_tiled_tma_atom_B(
if not isinstance(internal_type, NumericMeta):
raise TypeError(f"internal_type must be a Numeric, but got {internal_type}")
use_unpack = (internal_type.width == 8 and
isinstance(gmem_tensor.element_type, NumericMeta) and
gmem_tensor.element_type.width < 8)
internal_mlir_type = gmem_tensor.element_type.mlir_type if use_unpack else internal_type.mlir_type
use_unpack = (
internal_type.width == 8
and isinstance(gmem_tensor.element_type, NumericMeta)
and gmem_tensor.element_type.width < 8
)
internal_mlir_type = (
gmem_tensor.element_type.mlir_type
if use_unpack
else internal_type.mlir_type
)
tma_format = _cute_nvgpu_ir.TmaDataFormat(
_cute_nvgpu_ir.get_default_tma_format(internal_mlir_type, use_unpack)
)
@@ -60,4 +60,5 @@ __all__ = [
"make_tmem_copy",
"make_s2t_copy",
"get_s2t_smem_desc_tensor",
"make_umma_smem_desc",
]
@@ -9,14 +9,16 @@
# and related documentation outside the scope permitted by the EULA
# is strictly prohibited.
from typing import overload, Type, Tuple, Union
from typing import overload, Type, Tuple, Union, Optional
from cutlass.cutlass_dsl import dsl_user_op
from cutlass._mlir import ir
import cutlass._mlir.dialects.cute_nvgpu as _cute_nvgpu_ir
from cutlass._mlir.dialects import nvvm
from cutlass._mlir.dialects import nvvm, builtin
from ...typing import (
Pointer,
Shape,
IntTuple,
Layout,
@@ -27,6 +29,7 @@ from ...typing import (
NumericMeta,
Int16,
Int32,
Int64,
)
from ... import core
from ...tensor import recast_tensor
@@ -102,17 +105,27 @@ def make_smem_layout_atom(
SmemLayoutAtomKind.MN_SW128_32B,
):
# M/N-major layout
outer = core.make_layout(
(num_contiguous_elems, 8), stride=(1, num_contiguous_elems), loc=loc, ip=ip
return core.make_composed_layout(
sw,
0,
core.make_layout(
(num_contiguous_elems, 8), stride=(1, num_contiguous_elems)
),
loc=loc,
ip=ip,
)
else:
# K-major layout
outer = core.make_layout(
(8, num_contiguous_elems), stride=(num_contiguous_elems, 1), loc=loc, ip=ip
return core.make_composed_layout(
sw,
0,
core.make_layout(
(8, num_contiguous_elems), stride=(num_contiguous_elems, 1)
),
loc=loc,
ip=ip,
)
return core.make_composed_layout(sw, 0, outer, loc=loc, ip=ip)
@overload
def tile_to_mma_shape(
@@ -190,14 +203,27 @@ def commit(
mbar_ptr = mbar_ptr.llvm_ptr
if mask is not None:
mask = Int16(mask).ir_value(loc=loc, ip=ip)
nvvm.tcgen05_commit_arrive(
mbar_ptr, multicast_mask=mask, group=group, loc=loc, ip=ip
)
nvvm.tcgen05_commit(mbar_ptr, multicast_mask=mask, group=group, loc=loc, ip=ip)
else:
nvvm.tcgen05_commit_arrive(mbar_ptr, group=group, loc=loc, ip=ip)
nvvm.tcgen05_commit(mbar_ptr, group=group, loc=loc, ip=ip)
return
@dsl_user_op
def int_to_smem_descriptor(i, *, loc=None, ip=None) -> ir.Value:
desc_type = _cute_nvgpu_ir.SmemDescType.get()
return builtin.unrealized_conversion_cast(
[desc_type], [Int64(i).ir_value(loc=loc, ip=ip)], loc=loc, ip=ip
)
@dsl_user_op
def smem_descriptor_to_int(desc: ir.Value, *, loc=None, ip=None) -> Int64:
return Int64(
builtin.unrealized_conversion_cast([Int64.mlir_type], [desc], loc=loc, ip=ip)
)
####################################################################################################
#
# Helper functions for Copies
@@ -324,3 +350,55 @@ def get_s2t_smem_desc_tensor(
atom._trait.value, smem_tensor.value, loc=loc, ip=ip
)
return smem_desc_tensor
def make_umma_smem_desc(
src: Pointer,
layout: Layout,
major: str,
next_src: Optional[Pointer] = None,
*,
loc=None,
ip=None,
):
"""
Construct shared memory descriptor for UMMA.
The `make_umma_smem_desc` operation accepts an input cute.ptr (optionally a nextSrc
pointer for the second buffer in a circular buffer scheme), alongside a cute.layout
and a major attr, then constructs the shared memory descriptor and returns it.
The layout must be describing the buffer pointed to by the input pointer and the
iterator must carry valid swizzle information.
There are 5 supported swizzle variants:
- S<0, 4, 3> | SWIZZLE_NONE
- S<1, 4, 3> | SWIZZLE_32B
- S<2, 4, 3> | SWIZZLE_64B
- S<3, 4, 3> | SWIZZLE_128B
- S<2, 5, 2> | SWIZZLE_128B_BASE32B
The cute.ptr must carry shared address space and must be aligned to 16B.
:param src: The source pointer to shared memory
:type src: Pointer
:param layout: The layout describing the buffer
:type layout: Layout
:param major: The major mode attribute
:type major: str
:param next_src: Optional next source pointer for circular buffer scheme
:type next_src: Optional[Pointer]
:return: The shared memory descriptor
:rtype: SmemDescType
"""
src = src.value
if next_src is not None:
next_src = next_src.value
return _cute_nvgpu_ir.make_umma_smem_desc(
src=src,
layout=layout.type.attribute,
major=major,
next_src=next_src,
loc=loc,
ip=ip,
)
+247 -41
View File
@@ -20,7 +20,7 @@ import cutlass._mlir.dialects.cute as _cute_ir
import cutlass._mlir.dialects.cute_nvgpu as _cute_nvgpu_ir
from cutlass._mlir import ir
from ..common import OpError
from ..common import OpError, normalize_field_to_ir_name
from ... import core, atom
from ...core import _pack_shape, rank, depth
from ...typing import (
@@ -141,6 +141,7 @@ class Field(enum.Enum):
return self.value
# Base class for all tcgen05 MMA Ops with syntax `tcgen05.mma.cta_group.kind` used to factor out some internal code
@dataclass(frozen=True)
class MmaOp(Tcgen05MmaOp):
@@ -268,26 +269,30 @@ class MmaTraits(Trait):
admissible_fields = [Field.ACCUMULATE, Field.NEGATE_A, Field.NEGATE_B]
def set(self, field, value, *, loc=None, ip=None) -> None:
if field not in self.admissible_fields:
raise ValueError(
f"expects field to be one of {self.admissible_fields}, but got {field}"
field_ir = normalize_field_to_ir_name(field, self.admissible_fields)
bool_val = Boolean(value).ir_value(loc=loc, ip=ip)
try:
self.value = _cute_nvgpu_ir.atom_set_value(
self.value, field_ir, bool_val, loc=loc, ip=ip
)
except (TypeError, AttributeError):
# Legacy fallback
attr = ir.Attribute.parse(f"#cute_nvgpu.atom_mma_field_sm100<{field_ir}>")
self.value = _cute_nvgpu_ir.atom_set_value(
self.value, attr, bool_val, loc=loc, ip=ip
)
field_name = f"#cute_nvgpu.atom_mma_field_sm100<{field._to_ir_field_name()}>"
attr = ir.Attribute.parse(field_name)
self.value = _cute_nvgpu_ir.atom_set_value(
self.value, attr, Boolean(value).ir_value(loc=loc, ip=ip), loc=loc, ip=ip
)
def get(self, field, *, loc=None, ip=None) -> Any:
if field not in self.admissible_fields:
raise ValueError(
f"expects field to be one of {self.admissible_fields}, but got {field}"
field_ir = normalize_field_to_ir_name(field, self.admissible_fields)
try:
return _cute_nvgpu_ir.atom_get_value(
Boolean.mlir_type, self.value, field_ir, loc=loc, ip=ip
)
except (TypeError, AttributeError):
attr = ir.Attribute.parse(f"#cute_nvgpu.atom_mma_field_sm100<{field_ir}>")
return _cute_nvgpu_ir.atom_get_value(
Boolean.mlir_type, self.value, attr, loc=loc, ip=ip
)
field_name = f"#cute_nvgpu.atom_mma_field_sm100<{field._to_ir_field_name()}>"
attr = ir.Attribute.parse(field_name)
return _cute_nvgpu_ir.atom_get_value(
Boolean.mlir_type, self.value, attr, loc=loc, ip=ip
)
# Base class for all tcgen05 BlockScaled MMA Ops with syntax `tcgen05.mma.cta_group.kind.block_scale` used to factor out some internal code
@@ -420,33 +425,58 @@ class BlockScaledMmaTraits(Trait):
]
def set(self, field, value, *, loc=None, ip=None) -> None:
if field not in self.admissible_fields:
raise ValueError(
f"expects field to be one of {self.admissible_fields}, but got {field}"
)
if field in [Field.ACCUMULATE, Field.NEGATE_A, Field.NEGATE_B]:
value = Boolean(value).ir_value(loc=loc, ip=ip)
elif field in [Field.SFA, Field.SFB]:
field_ir = normalize_field_to_ir_name(field, self.admissible_fields)
# Derive boolean/pointer IR names from enum values, no hard-coded strings.
bool_field_ir = {
f._to_ir_field_name()
for f in self.admissible_fields
if f in (Field.ACCUMULATE, Field.NEGATE_A, Field.NEGATE_B)
}
ptr_field_ir = {
f._to_ir_field_name()
for f in self.admissible_fields
if f in (Field.SFA, Field.SFB)
}
# Coerce value based on field kind
if field_ir in bool_field_ir:
val = Boolean(value).ir_value(loc=loc, ip=ip)
elif field_ir in ptr_field_ir:
if not isinstance(value, Pointer):
raise ValueError(
f"expects value to be a pointer for {field}, but got {type(value).__name__}"
f"expects value to be a pointer for {field_ir}, but got {type(value).__name__}"
)
value = value.value
field_name = f"#cute_nvgpu.atom_mma_field_sm100_block_scaled<{field._to_ir_field_name()}>"
attr = ir.Attribute.parse(field_name)
self.value = _cute_nvgpu_ir.atom_set_value(
self.value, attr, value, loc=loc, ip=ip
)
val = value.value
else:
raise ValueError(f"unsupported field: {field_ir}")
try:
self.value = _cute_nvgpu_ir.atom_set_value(
self.value, field_ir, val, loc=loc, ip=ip
)
except (TypeError, AttributeError):
attr = ir.Attribute.parse(
f"#cute_nvgpu.atom_mma_field_sm100_block_scaled<{field_ir}>"
)
self.value = _cute_nvgpu_ir.atom_set_value(
self.value, attr, val, loc=loc, ip=ip
)
def get(self, field, *, loc=None, ip=None) -> Any:
if field not in [Field.ACCUMULATE, Field.NEGATE_A, Field.NEGATE_B]:
raise ValueError(f"the get method for {field} is not supported")
field_name = f"#cute_nvgpu.atom_mma_field_sm100_block_scaled<{field._to_ir_field_name()}>"
attr = ir.Attribute.parse(field_name)
return _cute_nvgpu_ir.atom_get_value(
Boolean.mlir_type, self.value, attr, loc=loc, ip=ip
)
# Only boolean-returning fields supported for get. Derive from admissible_fields.
gettable_fields = [
f for f in self.admissible_fields if f not in (Field.SFA, Field.SFB)
]
field_ir = normalize_field_to_ir_name(field, gettable_fields)
try:
return _cute_nvgpu_ir.atom_get_value(
Boolean.mlir_type, self.value, field_ir, loc=loc, ip=ip
)
except (TypeError, AttributeError):
attr = ir.Attribute.parse(
f"#cute_nvgpu.atom_mma_field_sm100_block_scaled<{field_ir}>"
)
return _cute_nvgpu_ir.atom_get_value(
Boolean.mlir_type, self.value, attr, loc=loc, ip=ip
)
#
@@ -802,6 +832,7 @@ class MmaFP8Trait(MmaTraits):
pass
#
# MXF8F6F4 MMA
#
@@ -946,7 +977,7 @@ class MmaMXF4Op(BlockScaledMmaOp):
f"but got {self.shape_mnk[2]}",
)
def _make_trait(self, *, loc=None, ip=None, **kwargs) -> "MmaMXF8Trait":
def _make_trait(self, *, loc=None, ip=None, **kwargs) -> "MmaMXF4Trait":
shape_mnk = _pack_shape(self.shape_mnk, loc=loc, ip=ip)
ty = _cute_nvgpu_ir.MmaAtomSM100UMMABlockScaledType.get(
shape_mnk.type.attribute,
@@ -1039,7 +1070,7 @@ class MmaMXF4NVF4Op(BlockScaledMmaOp):
f"but got {self.shape_mnk[2]}",
)
def _make_trait(self, *, loc=None, ip=None, **kwargs) -> "MmaMXF8Trait":
def _make_trait(self, *, loc=None, ip=None, **kwargs) -> "MmaMXF4NVF4Trait":
shape_mnk = _pack_shape(self.shape_mnk, loc=loc, ip=ip)
ty = _cute_nvgpu_ir.MmaAtomSM100UMMABlockScaledType.get(
shape_mnk.type.attribute,
@@ -1077,6 +1108,181 @@ class MmaMXF4NVF4Trait(BlockScaledMmaTraits):
pass
#
# SM103 MXF4 MMA
#
@dataclass(frozen=True)
class SM103MmaMXF4Op(BlockScaledMmaOp):
"""
SM103 MXF4 tcgen05 BlockScaled MMA Operation.
See the `PTX documentation <https://docs.nvidia.com/cuda/parallel-thread-execution/#tcgen05-mma-instructions-mma>`__.
This Operation corresponds to the ``.kind::mxf4`` qualifier.
This Operation is for SM103.
"""
descriptive_name = "tcgen05 SM103 MXF4 BlockScaled MMA Operation"
def __init__(
self,
instruction_shape: Shape,
cta_group: CtaGroup,
a_src: OperandSource,
) -> None:
super().__init__(
Float4E2M1FN,
Float4E2M1FN,
Float32,
Float8E8M0FNU,
32,
instruction_shape,
cta_group,
a_src,
OperandMajorMode.K,
OperandMajorMode.K,
)
self._verify()
def _verify(self) -> None:
# Instruction shape verification
instruction_k = 96
if rank(self.shape_mnk) == 2:
object.__setattr__(self, "shape_mnk", (*self.shape_mnk, instruction_k))
if self.shape_mnk[2] != instruction_k:
raise OpError(
self,
f"expects the instruction extent in the K-mode to be {instruction_k}, "
f"but got {self.shape_mnk[2]}",
)
def _make_trait(self, *, loc=None, ip=None, **kwargs) -> "MmaMXF4Trait":
shape_mnk = _pack_shape(self.shape_mnk, loc=loc, ip=ip)
ty = _cute_nvgpu_ir.MmaAtomSM100UMMABlockScaledType.get(
shape_mnk.type.attribute,
self.cta_group.value,
self.a_major_mode._to_ir(),
self.b_major_mode._to_ir(),
self.a_dtype.mlir_type,
self.b_dtype.mlir_type,
self.acc_dtype.mlir_type,
self.sf_dtype.mlir_type,
self.a_src._to_ir(),
self.sf_vec_size,
1030,
)
return MmaMXF4Trait(
make_atom(
ty,
(
Boolean(False).ir_value(loc=loc, ip=ip),
Boolean(False).ir_value(loc=loc, ip=ip),
Boolean(False).ir_value(loc=loc, ip=ip),
core.make_ptr(
self.sf_dtype, 0, _cute_ir.AddressSpace.tmem, loc=loc, ip=ip
).value,
core.make_ptr(
self.sf_dtype, 0, _cute_ir.AddressSpace.tmem, loc=loc, ip=ip
).value,
),
loc=loc,
ip=ip,
)
)
#
# SM103 MXF4NVF4 MMA
#
@dataclass(frozen=True)
class SM103MmaMXF4NVF4Op(BlockScaledMmaOp):
"""
SM103 MXF4NVF4 tcgen05 BlockScaled MMA Operation.
See the `PTX documentation <https://docs.nvidia.com/cuda/parallel-thread-execution/#tcgen05-mma-instructions-mma>`__.
This Operation corresponds to the ``.kind::mxf4nvf4`` qualifier.
This Operation is for SM103.
"""
descriptive_name = "tcgen05 SM103 MXF4NVF4 BlockScaled MMA Operation"
def __init__(
self,
sf_dtype: Type[Numeric],
instruction_shape: Shape,
cta_group: CtaGroup,
a_src: OperandSource,
) -> None:
super().__init__(
Float4E2M1FN,
Float4E2M1FN,
Float32,
sf_dtype,
16,
instruction_shape,
cta_group,
a_src,
OperandMajorMode.K,
OperandMajorMode.K,
)
self._verify()
def _verify(self) -> None:
# Scale Factor data type verification
if self.sf_dtype not in [Float8E8M0FNU, Float8E4M3FN]:
raise OpError(
self,
"expects the 'sf_dtype' Op parameter to be one of Float8E8M0FNU",
)
# Instruction shape verification
instruction_k = 96
if rank(self.shape_mnk) == 2:
object.__setattr__(self, "shape_mnk", (*self.shape_mnk, instruction_k))
if self.shape_mnk[2] != instruction_k:
raise OpError(
self,
f"expects the instruction extent in the K-mode to be {instruction_k}, "
f"but got {self.shape_mnk[2]}",
)
def _make_trait(self, *, loc=None, ip=None, **kwargs) -> "MmaMXF4NVF4Trait":
shape_mnk = _pack_shape(self.shape_mnk, loc=loc, ip=ip)
ty = _cute_nvgpu_ir.MmaAtomSM100UMMABlockScaledType.get(
shape_mnk.type.attribute,
self.cta_group.value,
self.a_major_mode._to_ir(),
self.b_major_mode._to_ir(),
self.a_dtype.mlir_type,
self.b_dtype.mlir_type,
self.acc_dtype.mlir_type,
self.sf_dtype.mlir_type,
self.a_src._to_ir(),
self.sf_vec_size,
1030,
)
return MmaMXF4NVF4Trait(
make_atom(
ty,
(
Boolean(False).ir_value(loc=loc, ip=ip),
Boolean(False).ir_value(loc=loc, ip=ip),
Boolean(False).ir_value(loc=loc, ip=ip),
core.make_ptr(
self.sf_dtype, 0, _cute_ir.AddressSpace.tmem, loc=loc, ip=ip
).value,
core.make_ptr(
self.sf_dtype, 0, _cute_ir.AddressSpace.tmem, loc=loc, ip=ip
).value,
),
loc=loc,
ip=ip,
)
)
####################################################################################################
#
# SMEM layout atoms
@@ -82,6 +82,7 @@ class LdMatrix8x8x16bOp(BaseOp):
class LdMatrix8x8x16bTrait(Trait):
pass
@dataclass(frozen=True)
class LdMatrix8x16x8bOp(BaseOp):
"""
@@ -125,11 +126,12 @@ class LdMatrix8x16x8bOp(BaseOp):
class LdMatrix8x16x8bTrait(Trait):
pass
@dataclass(frozen=True)
class LdMatrix16x8x8bOp(BaseOp):
"""
16x8 8b ``ldmatrix`` Operation with transpose
There is no direct PTX correspondance to this Op.
This actually lowers to ldmatrix with the ``.m16n16`` qualifier and
additional address and value permutations to match stmatrix.m16n8.trans.
@@ -166,6 +168,7 @@ class LdMatrix16x8x8bOp(BaseOp):
)
return LdMatrix16x8x8bTrait(make_atom(ty, loc=loc, ip=ip))
class LdMatrix16x8x8bTrait(Trait):
pass
@@ -176,7 +179,7 @@ class LdMatrix16x16x8bOp(BaseOp):
16x16 ``ldmatrix`` Operation with transpose and optional unpacking to 8b container.
Packed source container is 16x4b elements with 64b padding
or 16x6b elements with 32b padding (total 128b per 16 elements)
See the `PTX documentation <https://docs.nvidia.com/cuda/parallel-thread-execution/#warp-level-matrix-load-instruction-ldmatrix>`__.
This operation corresponds to the ``.m16n16`` and the ``.b4x16_p64``,``.b6x16_p32``,``.b8`` qualifiers.
"""
@@ -15,7 +15,7 @@ from typing import Type, Any
import enum
from cutlass import cute
from cutlass.base_dsl.arch import Arch
from cutlass.cutlass_dsl import CuTeDSL
from cutlass.cutlass_dsl import BaseDSL
from ..common import OpError
@@ -134,7 +134,7 @@ class MmaSM120BlockScaledOp(MmaOp):
def __post_init__(self) -> None:
# Verify arch
arch = CuTeDSL._get_dsl().get_arch_enum()
arch = BaseDSL._get_dsl().get_arch_enum()
if not arch == Arch.sm_120a:
raise OpError(
self,
@@ -174,6 +174,7 @@ class MmaSM120BlockScaledOp(MmaOp):
self,
"expects the 'sf_vec_size' Op parameter to be 16 or 32",
)
def __str__(self) -> str:
return (
"warp-level MXF4/MXF4NVF4 MMA Operation"
@@ -190,6 +191,7 @@ class MmaSM120BlockScaledOp(MmaOp):
def _verify_fragment_B(self, input: Tensor, *, loc=None, ip=None):
pass
class Field(enum.Enum):
"""
An enumeration for the fields of the MMA Atom that can be modified at runtime.
@@ -15,12 +15,13 @@ from typing import Type, Any
from cutlass.base_dsl.arch import Arch
from cutlass.cutlass_dsl import BaseDSL, T
from typing_extensions import deprecated
import cutlass._mlir.dialects.cute as _cute_ir
import cutlass._mlir.dialects.cute_nvgpu as _cute_nvgpu_ir
from cutlass._mlir import ir
from ..common import OpError
from ..common import OpError, normalize_field_to_ir_name
from ...core import _pack_shape, rank, depth
from ...typing import (
Shape,
@@ -208,27 +209,44 @@ class MmaOp(WarpGroupMmaOp):
class MmaTraits(Trait):
admissible_fields = [Field.ACCUMULATE]
def _normalize_field_name(self, field: Any) -> str:
"""
Normalize a field specifier (enum or string) into the IR logical field name.
Accepted inputs:
- Field.ACCUMULATE
- "accum_c"
"""
return normalize_field_to_ir_name(field, self.admissible_fields)
def set(self, field, value, *, loc=None, ip=None) -> None:
if field not in self.admissible_fields:
raise ValueError(
f"invalid field, must be {Field.ACCUMULATE}, but got {field}"
field_ir_name = self._normalize_field_name(field)
# Prefer the newer builder that accepts a logical field name, but keep
# a fallback for legacy attribute-based construction to avoid breaking changes.
bool_val = Boolean(value).ir_value(loc=loc, ip=ip)
try:
self.value = _cute_nvgpu_ir.atom_set_value(
self.value, field_ir_name, bool_val, loc=loc, ip=ip
)
except (TypeError, AttributeError):
# Legacy path: construct the per-arch field attribute explicitly
attr_asm = f"#cute_nvgpu.atom_mma_field_sm90<{field_ir_name}>"
attr = ir.Attribute.parse(attr_asm)
self.value = _cute_nvgpu_ir.atom_set_value(
self.value, attr, bool_val, loc=loc, ip=ip
)
field_name = f"#cute_nvgpu.atom_mma_field_sm90<{field._to_ir_field_name()}>"
attr = ir.Attribute.parse(field_name)
self.value = _cute_nvgpu_ir.atom_set_value(
self.value, attr, Boolean(value).ir_value(loc=loc, ip=ip), loc=loc, ip=ip
)
def get(self, field, *, loc=None, ip=None) -> Any:
if field not in self.admissible_fields:
raise ValueError(
f"invalid field, must be {Field.ACCUMULATE}, but got {field}"
field_ir_name = self._normalize_field_name(field)
try:
return _cute_nvgpu_ir.atom_get_value(
Boolean.mlir_type, self.value, field_ir_name, loc=loc, ip=ip
)
except (TypeError, AttributeError):
attr_asm = f"#cute_nvgpu.atom_mma_field_sm90<{field_ir_name}>"
attr = ir.Attribute.parse(attr_asm)
return _cute_nvgpu_ir.atom_get_value(
Boolean.mlir_type, self.value, attr, loc=loc, ip=ip
)
field_name = f"#cute_nvgpu.atom_mma_field_sm90<{field._to_ir_field_name()}>"
attr = ir.Attribute.parse(field_name)
return _cute_nvgpu_ir.atom_get_value(
Boolean.mlir_type, self.value, attr, loc=loc, ip=ip
)
@dataclass(frozen=True)
+50 -53
View File
@@ -20,6 +20,7 @@ from cutlass.cutlass_dsl import (
T,
cutlass_arith,
_binary_op_type_promote,
MLIR_DYNAMIC,
BaseDSL,
)
from cutlass._mlir import ir
@@ -75,35 +76,8 @@ from .core import (
recast_layout,
)
from .typing import (
IntTuple,
Coord,
Shape,
Stride,
Pointer,
Layout,
ComposedLayout,
Tensor,
AddressSpace,
is_integer,
is_int_tuple,
as_numeric,
)
from .typing import (
Numeric,
Integer,
Boolean,
Int4,
Uint8,
Int8,
Int32,
Float4E2M1FN,
Float16,
Float32,
BFloat16,
)
from .tuple import transform_leaf, product, product_like, flatten_to_tuple
from .arch import cvt_i8_bf16_intrinsic, cvt_i4_bf16_intrinsic, cvt_f4e2m1_f16_intrinsic
from .arch import cvt_i8_bf16_intrinsic, cvt_i4_bf16_intrinsic
__all__ = [
@@ -439,10 +413,9 @@ class _Tensor(Tensor):
return _cute_ir.get_layout(self.value, loc=loc, ip=ip)
@property
@dsl_user_op
@lru_cache_ir()
def shape(self, *, loc=None, ip=None) -> Shape:
return self.layout.shape_method(loc=loc, ip=ip)
def shape(self) -> Shape:
return self.layout.shape
@property
@lru_cache_ir()
@@ -480,12 +453,23 @@ class _Tensor(Tensor):
raise ValueError(f"{self} doesn't have memspace")
@dsl_user_op
def load(self, *, loc=None, ip=None) -> "TensorSSA":
def load(
self,
*,
mask: Optional["TensorSSA"] = None,
pass_thru: Optional["TensorSSA"] = None,
loc=None,
ip=None,
) -> "TensorSSA":
"""Load tensor elements as a vector.
Loads all elements of the tensor into a vector representation, assuming the tensor
has a static shape and is in a memory space that supports load operations.
:param mask: Mask vector, defaults to None
:type mask: Optional[TensorSSA]
:param pass_thru: Pass through vector, defaults to None
:type pass_thru: Optional[TensorSSA]
:param loc: Source location for MLIR operation tracking, defaults to None
:type loc: Optional[Location]
:param ip: Insertion point for MLIR operation, defaults to None
@@ -501,9 +485,15 @@ class _Tensor(Tensor):
if not is_static(self.shape):
raise ValueError("dynamic layout doesn't support load")
self._check_can_load_store()
self._check_can_load_store(vectorized=True)
res_vect = _cute_ir.memref_load_vec(self.value, loc=loc, ip=ip)
mask_val = None if mask is None else mask.ir_value(loc=loc, ip=ip)
pass_thru_val = (
None if pass_thru is None else self._cvt_to_dest(pass_thru, loc=loc, ip=ip)
)
res_vect = _cute_ir.memref_load_vec(
self.value, mask=mask_val, pass_thru=pass_thru_val, loc=loc, ip=ip
)
if self.element_type is Boolean:
assert res_vect.type.element_type == T.i8(), (
f"Boolean tensor must be stored as i8 in memory, but got {res_vect.type.element_type}"
@@ -515,7 +505,14 @@ class _Tensor(Tensor):
return TensorSSA(res_vect, self.shape, self.element_type)
@dsl_user_op
def store(self, data: "TensorSSA", *, loc=None, ip=None):
def store(
self,
data: "TensorSSA",
*,
mask: Optional["TensorSSA"] = None,
loc=None,
ip=None,
):
"""Store vector data into tensor.
Stores vector data into the tensor, assuming matching shapes and a memory space
@@ -523,6 +520,8 @@ class _Tensor(Tensor):
:param data: Vector data to store into tensor
:type data: TensorSSA
:param mask: Mask vector, defaults to None
:type mask: Optional[TensorSSA]
:param loc: Source location for MLIR operation tracking, defaults to None
:type loc: Optional[Location]
:param ip: Insertion point for MLIR operation, defaults to None
@@ -538,7 +537,7 @@ class _Tensor(Tensor):
if not is_static(self.shape):
raise ValueError("Dynamic layout doesn't support vectorized store")
self._check_can_load_store()
self._check_can_load_store(vectorized=True)
n_elems = size(self.shape, loc=loc, ip=ip)
if n_elems != size(data.shape, loc=loc, ip=ip):
@@ -556,7 +555,11 @@ class _Tensor(Tensor):
# Implicit upcast to wider type
new_data = self._cvt_to_dest(data, loc=loc, ip=ip)
return _cute_ir.memref_store_vec(new_data, self.value, loc=loc, ip=ip)
mask_val = None if mask is None else mask.ir_value(loc=loc, ip=ip)
return _cute_ir.memref_store_vec(
new_data, self.value, mask=mask_val, loc=loc, ip=ip
)
@dsl_user_op
def fill(self, value: Numeric, *, loc=None, ip=None) -> None:
@@ -585,7 +588,7 @@ class _Tensor(Tensor):
# Fill tensor with constant value
tensor.fill(0.5) # All elements become 0.5
"""
self._check_can_load_store()
self._check_can_load_store(vectorized=True)
sz = size(self, loc=loc, ip=ip)
if type(sz) is not int:
@@ -599,7 +602,7 @@ class _Tensor(Tensor):
)
self.store(vect_val, loc=loc, ip=ip)
def _check_can_load_store(self):
def _check_can_load_store(self, vectorized: bool = False):
if not isinstance(self.type, _cute_ir.MemRefType) or self.memspace not in (
AddressSpace.rmem,
AddressSpace.smem,
@@ -608,9 +611,9 @@ class _Tensor(Tensor):
):
raise ValueError(f"{self} doesn't support load and store")
if self.type.is_swizzled:
if vectorized and isinstance(self.layout, ComposedLayout):
raise NotImplementedError(
f"load & store swizzled memory is not supported yet: {self}"
"vectorized load/store on tensor with composed layout is not supported yet"
)
def _check_can_dereference(self):
@@ -1038,8 +1041,10 @@ def print_tensor(
signed = tensor.element_type.signed
else:
signed = False
else:
elif isinstance(tensor.type, _cute_ir.CoordTensorType):
signed = True
else:
raise ValueError(f"unsupported tensor type for print_tensor, got {tensor.type}")
_cute_ir.print_view(tensor.value, verbose=verbose, is_signed=signed, loc=loc, ip=ip)
@@ -1750,7 +1755,8 @@ class TensorSSA(cutlass_arith.ArithValue):
idx = crd2idx(crd, self._layout, loc=loc, ip=ip)
assert not isinstance(idx, tuple), "index must be scalar"
idx_val = as_numeric(idx).ir_value(loc=loc, ip=ip)
res_val = vector.extractelement(self, position=idx_val, loc=loc, ip=ip)
idx_val = arith.index_cast(T.index(), idx_val, loc=loc, ip=ip)
res_val = vector.extract(self, [idx_val], [MLIR_DYNAMIC], loc=loc, ip=ip)
return self.dtype(res_val)
if not is_static(crd):
@@ -1817,16 +1823,7 @@ class TensorSSA(cutlass_arith.ArithValue):
# maybe downcast can lose signedness
src = self.maybe_downcast().with_signedness(self.signed)
if src_dtype.is_float and dtype.is_float:
if src_dtype == Float4E2M1FN and dtype in (Float16, Float32):
res_vect = cvt_f4e2m1_f16_intrinsic(
src, size(self.shape), loc=loc, ip=ip
)
if dtype == Float32:
res_vect = cutlass_arith.cvtf(
res_vect, dtype.mlir_type, loc=loc, ip=ip
)
else:
res_vect = cutlass_arith.cvtf(src, dtype.mlir_type, loc=loc, ip=ip)
res_vect = cutlass_arith.cvtf(src, dtype.mlir_type, loc=loc, ip=ip)
elif src_dtype.is_float and issubclass(dtype, Integer):
res_vect = cutlass_arith.fptoi(
src, dtype.signed, dtype.mlir_type, loc=loc, ip=ip
+6 -6
View File
@@ -20,15 +20,12 @@ from typing import Type, Union, Callable, Optional, Dict, List, Any
import cuda.bindings.driver as cuda_driver
import cuda.bindings.runtime as cuda_runtime
import cutlass
import cutlass.base_dsl.jit_executor
import cutlass.cutlass_dsl.cuda_jit_executor
from cutlass.cutlass_dsl import Constexpr, CuTeDSL, T, dsl_user_op, const_expr
from .typing import Numeric, Int8, Boolean, Tensor, Layout, Shape
from . import nvgpu
from .core import recast_layout, make_layout, composition, get, rank, size, zipped_divide
from .core import recast_layout, make_layout, composition, get, rank, size
from .tuple import elem_less
from .tensor import (
make_rmem_tensor,
@@ -39,6 +36,7 @@ from .tensor import (
)
from .atom import make_copy_atom
from .algorithm import copy
from .core import zipped_divide
from .runtime import from_dlpack
from cutlass._mlir.dialects import builtin, cf, nvvm, vector
@@ -76,7 +74,7 @@ class _CompileTimeAssertion(Assertion):
def __init__(
self,
tensor: _Tensor,
tensor: Tensor,
num_assertions: int = 1,
msgs=None,
device=None,
@@ -849,7 +847,9 @@ def get_workspace_count(
:return: Number of workspaces needed
:rtype: int
"""
num_l2_cache_bytes = cutlass.utils.HardwareInfo().get_l2_cache_size_in_bytes()
from cutlass.utils import HardwareInfo
num_l2_cache_bytes = HardwareInfo().get_l2_cache_size_in_bytes()
num_workspaces = (num_l2_cache_bytes * 3) // one_workspace_bytes + 1
num_iters = warmup_iterations + iterations
return num_iters if num_iters < num_workspaces else num_workspaces
+5 -3
View File
@@ -12,7 +12,6 @@
from abc import ABC, abstractmethod
import ctypes
from typing import ForwardRef, Tuple, Union, Any, Type, List, Optional, Literal
from functools import lru_cache
from cutlass.base_dsl.typing import *
@@ -28,9 +27,13 @@ class SymInt:
def __init__(self, width: Literal[32, 64] = 32, *, divisibility=1):
if width not in [32, 64]:
raise ValueError(f"Unsupported width: {width}")
self._width = width
self._divisibility = divisibility
def __hash__(self):
return hash((self._width, self._divisibility))
@property
def width(self):
return self._width
@@ -80,6 +83,7 @@ class SymInt:
else:
assert False, f"Unsupported width: {self.width}"
return self
def sym_int(width: Literal[32, 64] = 32, *, divisibility=1) -> SymInt:
return SymInt(width, divisibility=divisibility)
@@ -403,6 +407,4 @@ __all__ = [
"XTuple",
"is_integer",
"is_int_tuple",
"Pointer",
"Tensor",
]
@@ -53,7 +53,6 @@ from ..base_dsl.compiler import (
KeepCUBIN,
KeepPTX,
GPUArch,
LinkLibraries,
EnableTVMFFI,
)
from ..base_dsl.runtime.jit_arg_adapters import *
@@ -58,8 +58,6 @@ class CudaDialectJitModule:
for library in self.cuda_library:
cuda_runtime.cudaLibraryUnload(library)
self.cuda_library.clear()
except Exception as e:
pass
finally:
self._unloaded = True
+87 -37
View File
@@ -22,6 +22,7 @@ from typing import (
List,
Tuple,
Sequence,
Iterable,
ForwardRef,
Any,
get_origin,
@@ -34,7 +35,6 @@ from dataclasses import is_dataclass, fields
from math import ceil
from itertools import chain
from pathlib import Path
from collections.abc import Sequence
import builtins
import ctypes
import hashlib
@@ -65,10 +65,14 @@ from cutlass._mlir.dialects import (
from cutlass._mlir.dialects._ods_common import (
get_op_result_or_op_results as _get_op_result_or_op_results,
)
from cutlass._mlir.dialects import lir as cutlass_lir
from cutlass._mlir.extras import types as T
# Helpers
from ..base_dsl._mlir_helpers import arith as cutlass_arith
from ..base_dsl._mlir_helpers import lru_cache_ir
from ..base_dsl._mlir_helpers.op import dsl_user_op
from ..base_dsl._mlir_helpers.arith import const
@@ -94,6 +98,7 @@ from .cutlass_ast_decorators import (
_loop_execute_range_dynamic,
_if_execute_dynamic,
_while_execute_dynamic,
_ifexp_execute_dynamic,
)
from ..base_dsl.runtime.jit_arg_adapters import JitArgAdapterRegistry
@@ -275,6 +280,11 @@ class CutlassBaseDSL(BaseDSL):
log().info(f"self: {self}")
log().info(f"Entering GPU module for {self.name}")
log().info(f"GPU module: {self.gpu_module}")
if not self.gpu_module:
raise DSLRuntimeError(
f"GPU module is not set, probably compilation of a kernel from different DSL decorator",
suggestion=f"Use the same DSL decorator to build the GPU module, DSL: {type(self).__name__}",
)
return ir.InsertionPoint(self.gpu_module.bodyRegion.blocks[0])
@staticmethod
@@ -290,9 +300,9 @@ class CutlassBaseDSL(BaseDSL):
)
def _generate_kernel_attrs(self, config: BaseDSL.LaunchConfig) -> dict:
assert isinstance(config, BaseDSL.LaunchConfig), (
f"Expect LaunchConfig for @kernel, but got {type(config)}"
)
assert isinstance(
config, BaseDSL.LaunchConfig
), f"Expect LaunchConfig for @kernel, but got {type(config)}"
ret = {}
if config.has_max_number_threads():
@@ -381,16 +391,7 @@ class CutlassBaseDSL(BaseDSL):
) from e
files.append((giant_dso_name, so_path, so_size))
def handle_import_error(exc):
"""Handle errors during package walking, ignoring ImportError and NotImplementedError."""
if isinstance(exc, (ImportError, NotImplementedError)):
log().info(f"Skipping module due to {type(exc).__name__}: {exc}")
else:
log().warning(f"Unexpected error during package walk: {exc}")
for lib in pkgutil.walk_packages(
[dsl_path], prefix="cutlass.", onerror=handle_import_error
):
for lib in pkgutil.walk_packages([dsl_path], prefix="cutlass."):
spec = lib.module_finder.find_spec(lib.name)
if not spec or not spec.origin:
continue
@@ -624,12 +625,7 @@ class CutlassBaseDSL(BaseDSL):
loc=None,
ip=None,
):
# set to 3 for PDL, cluster size, and cooperative
max_num_attributes = 3
if preferred_cluster_size_x is not None:
max_num_attributes += 1
max_num_attributes = 17
launch_config_type = cuda_dialect.LaunchConfigType.get(max_num_attributes)
if len(stream) == 0:
@@ -653,8 +649,6 @@ class CutlassBaseDSL(BaseDSL):
cfg = cuda_dialect.launch_cfg_create(
# Launch config type
launch_config_type,
# Max num of attributes the launch config can hold
# set to 3 for PDL, cluster size, and cooperative
ir.IntegerAttr.get(ir.IntegerType.get_signless(32), max_num_attributes),
block_size_x,
block_size_y,
@@ -793,15 +787,15 @@ class CutlassBaseDSL(BaseDSL):
requiredArgs = kwargs.get("requiredArgs", None)
loc = kwargs.get("loc", None)
assert kernelSym is not None, "kernelSym being None is not expected!"
assert requiredArgs is not None, (
"requiredArgs being None is not expected!"
)
assert kernelOperands is not None, (
"kernelOperands being None is not expected!"
)
assert isinstance(requiredArgs.config, BaseDSL.LaunchConfig), (
f"Expect LaunchConfig for @kernel, but got {type(requiredArgs.config)}"
)
assert (
requiredArgs is not None
), "requiredArgs being None is not expected!"
assert (
kernelOperands is not None
), "kernelOperands being None is not expected!"
assert isinstance(
requiredArgs.config, BaseDSL.LaunchConfig
), f"Expect LaunchConfig for @kernel, but got {type(requiredArgs.config)}"
cfg = requiredArgs.config
@@ -822,6 +816,7 @@ class CutlassBaseDSL(BaseDSL):
if not isinstance(cfg.async_deps, (list, tuple)):
async_deps = [cfg.async_deps]
# Prepare launch kwargs
launch_kwargs = {}
if cfg.has_fallback_cluster:
@@ -1091,6 +1086,58 @@ class CuTeDSL(CutlassBaseDSL):
return cuda_dialect.ReturnOp([], loc=loc, ip=ip)
# =============================================================================
# CuteExperimental DSL Class
# =============================================================================
class CuteExperimentalDSL(CutlassBaseDSL):
def __init__(self):
name = "CUTE_EXPERIMENTAL_DSL"
compiler_provider = compiler.Compiler(passmanager, execution_engine)
pass_sm_arch_name = "cubin-chip"
super().__init__(name, compiler_provider, pass_sm_arch_name, preprocess=True)
def _get_pipeline(self, pipeline):
if pipeline == None:
return "builtin.module(gpu.module(lir-to-cute{enable-cuda-dialect enable-lir-func-finalization=false}), lir-func-finalization{enable-cuda-dialect=true}, cute-to-nvvm{check-inline-asm=false cubin-format=bin enable-cuda-dialect})"
return pipeline
@staticmethod
def generate_func_op(arg_types, arg_attrs, kernel_name, loc=None):
func_op = cutlass_lir.FuncOp(
ir.StringAttr.get(kernel_name),
ir.TypeAttr.get(ir.FunctionType.get(arg_types, [])),
loc=loc,
)
func_op.attributes["cu_attrs"] = ir.DictAttr.get(
{
str(
cuda_dialect.CUFunctionAttribute.non_portable_cluster_size_allowed
): ir.IntegerAttr.get(ir.IntegerType.get_signless(32), 1),
str(
cuda_dialect.CUFunctionAttribute.max_dynamic_shared_size_bytes
): cuda_dialect.DevMaxSharedMemoryOptinAttr.get(),
}
)
# Monkey patch FuncOp to add an add_entry_block method, if not already defined.
if not hasattr(func_op, "add_entry_block"):
def add_entry_block(arg_locs):
if len(func_op.body.blocks) != 0:
raise RuntimeError("The function already has an entry block.")
func_op.body.blocks.append(*arg_types)
return func_op.body.blocks[0]
func_op.add_entry_block = add_entry_block
return func_op
@staticmethod
def generate_func_ret_op(loc=None, ip=None):
return cutlass_lir.ReturnOp([])
# =============================================================================
# KernelLauncher
# =============================================================================
@@ -1331,9 +1378,9 @@ def to_index(value):
if is_dynamic_expression(value):
if isinstance(value, Numeric):
value = value.ir_value()
assert ir.IntegerType.isinstance(value.type), (
f"expects integer type, but got {value.type}"
)
assert ir.IntegerType.isinstance(
value.type
), f"expects integer type, but got {value.type}"
res = arith.index_cast(T.index(), value)
else:
res = const(int(value), ty=T.index())
@@ -1379,7 +1426,7 @@ def _validate_iter_args_structure(iter_args, ir_values):
def _minmax(op, *args, loc=None, ip=None):
"""Computes the minimum or maximum value from the provided arguments."""
from ..base_dsl.typing import _binary_op_type_promote
from ..base_dsl.typing import _binary_op, _binary_op_type_promote
# AST Traversal doesn't support early exit in if executor
x = None
@@ -1830,7 +1877,7 @@ def for_generate(
def _createI32Attr(value):
if not isinstance(value, int):
raise DSLRuntimeError("value must be int.")
raise DSLRuntimeError(f"value must be int.")
return ir.IntegerAttr.get(ir.IntegerType.get_signless(32), value)
ir_iter_args = extract_mlir_values(iter_args) if iter_args is not None else None
@@ -1951,7 +1998,9 @@ def if_generate(
# Collect MLIR results.
mlir_results = _get_op_result_or_op_results(if_op)
if not isinstance(mlir_results, list):
if not isinstance(mlir_results, list) and not isinstance(
mlir_results, ir.OpResultList
):
mlir_results = [mlir_results]
# Wrap the results with their DSL types.
@@ -2245,6 +2294,7 @@ executor.set_functions(
any_executor=any_,
all_executor=all_,
builtin_redirector=_builtin_redirector,
ifexp_dynamic=_ifexp_execute_dynamic,
)
@@ -374,7 +374,7 @@ def _loop_execute_range_dynamic(
for i, d in enumerate(dyn_yield_ops)
)
raise DSLRuntimeError(
f"Failed to create scf.ForOp \n\t\tstart={start_}: type : {type(start_)}"
f"Failed to create dynamic for loop \n\t\tstart={start_}: type : {type(start_)}"
f"\n\t\tstop={stop_}: type : {type(stop_)}\n\t\tstep={step_}: type : {type(step_)}"
f", \n\tdyn_yield_ops:\n{yield_ops}"
) from e
@@ -461,7 +461,7 @@ def _if_execute_dynamic(
)
except Exception as e:
raise DSLRuntimeError(
f"Failed to create scf.IfOp \n\t\tpred={pred_}: type : {type(pred_)}"
f"Failed to create dynamic if \n\t\tpred={pred_}: type : {type(pred_)}"
) from e
return if_op
@@ -550,7 +550,7 @@ def _while_execute_dynamic(
for i, d in enumerate(dyn_yield_ops)
)
raise DSLRuntimeError(
f"Failed to create scf.WhileOp with yield_ops:\n{yield_ops}"
f"Failed to create dynamic while loop with yield_ops:\n{yield_ops}"
) from e
def before_block_builder(
@@ -643,3 +643,121 @@ def _while_execute_dynamic(
before_block_builder: before_block_terminator
}, # Only customize the before block
)
def _ifexp_execute_dynamic(
pred: "ir.Value",
generator_targets: tuple,
then_block: Callable,
else_block: Callable,
):
"""
Dynamically execute a Python inline if-expression (ternary) as a runtime-dispatched control flow op.
This function builds an SCF (Structured Control Flow) `if` operation in the IR, using the given
predicate and block functions for the 'then' and 'else' branches, and infers the result types
from the return signature of those blocks. It ensures that both branches return values of the same
tree structure and types, so that the IR op can properly yield their results.
Parameters
----------
pred : ir.Value
The predicate value (a boolean IR value) that determines which branch is executed.
generator_targets : tuple
The generator targets that are passed to the then and else blocks.
then_block : Callable
A Python function that executes the 'then' branch and returns the result(s). This will be
executed if `pred` evaluates to True.
else_block : Callable
A Python function that executes the 'else' branch and returns the result(s). This will be
executed if `pred` evaluates to False.
Returns
-------
list
The evaluated result(s) of the selected branch, in a standardized (possibly list-wrapped) format.
Raises
------
DSLRuntimeError
If the 'then' and 'else' blocks return values of different tree structures or types,
or if IR construction fails.
Notes
-----
This function is a low-level implementation intended for use by the AST transformation machinery,
and not for direct user invocation. It acts as the backend for transformed Python inline if-expressions.
"""
# Infer result types by running both branches with dummy arguments in a temporary region
execution_region = scf.ExecuteRegionOp(result=[])
execution_region.region.blocks.append()
result_types = []
mix_iter_args = []
with ir.InsertionPoint(execution_region.region.blocks[0]):
# Call the then block and unpack its results to IR values and tree structure
then_results = ScfGenerator._normalize_region_result_to_list(
then_block(*generator_targets)
)
ir_values, then_tree = cutlass_dsl.unpack_to_irvalue(then_results, "ifexp", 0)
# Call the else block and unpack its results to IR values and tree structure
else_results = ScfGenerator._normalize_region_result_to_list(
else_block(*generator_targets)
)
_, else_tree = cutlass_dsl.unpack_to_irvalue(else_results, "ifexp", 0)
# Check that both branches are structurally and type compatible
if check_tree_equal(then_tree, else_tree) != -1:
raise DSLRuntimeError(
"Then and else blocks of ifexp return different types"
)
# Collect result types for the SCF IfOp
result_types.extend([arg.type for arg in ir_values])
mix_iter_args.extend(then_results)
# Set up a generator for SCF op creation
scf_gen = ScfGenerator()
# Function to create the IfOp with correct predicate and result types
def create_if_op(_):
pred_ = Boolean(pred)
try:
if_op = scf.IfOp(
pred_.ir_value(),
hasElse=True,
results_=result_types,
)
except Exception as e:
raise DSLRuntimeError(
f"Failed to create dynamic if-expression \n\t\tpred={pred_}: type : {type(pred_)}"
) from e
return if_op
# SCF region builder for then block
def then_builder(*args):
# Just call the then_block as no arguments are passed to it
return then_block(*generator_targets)
# SCF region builder for else block
def else_builder(*args):
return else_block(*generator_targets)
# Prepare the list of region builders for the SCF IfOp: first for "then", then for "else"
region_builders = [then_builder, else_builder]
ret = scf_gen.scf_execute_dynamic(
op_type_name="if",
mix_iter_args=mix_iter_args,
full_write_args_count=0,
mix_iter_arg_names=["unknown" for _ in mix_iter_args],
create_op_func=create_if_op,
region_builders=region_builders,
)
# Clean up: Remove the temporary execution region from the IR graph
execution_region.operation.erase()
return ret
@@ -21,7 +21,6 @@ from cutlass._mlir.dialects import llvm
from cutlass._mlir._mlir_libs._cutlass_ir import _aot_support
from cutlass.cutlass_dsl.cuda_jit_executor import CudaDialectJitCompiledFunction
from cutlass.base_dsl.common import DSLRuntimeError
from cutlass.base_dsl.jit_executor import ExecutionArgs
from typing import Optional, Callable
import tvm_ffi
@@ -211,6 +210,7 @@ class TVMFFICuteCallProvider(DynamicParamPackCallProvider):
global_dtors = llvm.mlir_global_dtors(
dtors=[],
priorities=[],
data=[],
)
else:
# use the existing global destructors
@@ -223,6 +223,9 @@ class TVMFFICuteCallProvider(DynamicParamPackCallProvider):
global_dtors.attributes["priorities"] += [
ir.IntegerAttr.get(self.i32_type, 65535)
] # the default priority
global_dtors.attributes["data"] += [
ir.FlatSymbolRefAttr.get(unload_func_wrapper_symbol)
] # the data will not be used, but we need to pass something to satisfy the llvm.mlir.global_dtors op
return current_block
@@ -255,7 +258,8 @@ class TVMFFICuteCallProvider(DynamicParamPackCallProvider):
current_device: Optional[ir.Value],
target_device: Optional[ir.Value],
) -> ir.Block:
"""Set the CUDA device index if it differs from the target device."""
"""Set the CUDA device index if it differs from the target device.
"""
# If either device is None, no switching needed
if current_device is None:
assert target_device is None
@@ -273,7 +277,7 @@ class TVMFFICuteCallProvider(DynamicParamPackCallProvider):
self.cond_br(
cond=devices_differ,
true_block=switch_device_block,
false_block=continuation_block,
false_block=continuation_block
)
# Switch device block: call cudaSetDevice
@@ -287,9 +291,7 @@ class TVMFFICuteCallProvider(DynamicParamPackCallProvider):
)
# Check for errors and branch to continuation
switch_device_block = self.check_cuda_error(
result, switch_device_block, context
)
switch_device_block = self.check_cuda_error(result, switch_device_block, context)
with ir.InsertionPoint(switch_device_block):
self.br(continuation_block)
@@ -320,9 +322,7 @@ class TVMFFICuteCallProvider(DynamicParamPackCallProvider):
op_bundle_sizes=[],
op_bundle_operands=[],
)
current_block = self.check_cuda_error(
get_device_result, current_block, context
)
current_block = self.check_cuda_error(get_device_result, current_block, context)
# Load the current device index from the alloca
with ir.InsertionPoint(current_block):
@@ -354,6 +354,7 @@ class TVMFFICuteCallProvider(DynamicParamPackCallProvider):
return current_block
def find_cuda_device_index_from_params(self, context: CallContext):
"""Find the CUDA device index from tensor parameters."""
for param in context.params:
@@ -365,9 +366,12 @@ class TVMFFICuteCallProvider(DynamicParamPackCallProvider):
return None
def create_shared_cuda_error_block(
self, current_block: ir.Block, context: CallContext
self,
current_block: ir.Block,
context: CallContext
) -> ir.Block:
"""Create a shared error handling block for all CUDA errors."""
"""Create a shared error handling block for all CUDA errors.
"""
# Create the shared error block after the current block (setup phase)
# This block will be branched to from multiple error checking sites
# It accepts the error code as a block argument
@@ -397,9 +401,7 @@ class TVMFFICuteCallProvider(DynamicParamPackCallProvider):
current_block = self.append_unload_to_global_dtors(current_block, context)
# Create shared CUDA error handling block after the setup blocks
# This reduces code duplication - all CUDA errors branch to this single block
self.cuda_error_handle_block = self.create_shared_cuda_error_block(
current_block, context
)
self.cuda_error_handle_block = self.create_shared_cuda_error_block(current_block, context)
# setup device index, will be set around the call to the target function
self.cuda_device_index = self.find_cuda_device_index_from_params(context)
current_block = super().__call__(current_block, context)
@@ -458,6 +460,10 @@ class TVMFFIJitCompiledFunctionBase(CudaDialectJitCompiledFunction):
def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
# use direct call to the tvm_ffi.Function.__call__
# to avoid most of python overhead
__call__ = tvm_ffi.Function.__call__
def to(self, device=None):
"""TVM FFI function itself is already support all devices."""
return self
+1
View File
@@ -285,6 +285,7 @@ class JaxArrayValue(JaxArray):
llvm.PointerType.get(),
shape_array,
[],
no_wrap_flags=0,
raw_constant_indices=ir.DenseI32ArrayAttr.get([i]),
elem_type=i64,
loc=loc,
+30 -58
View File
@@ -10,22 +10,14 @@
# is strictly prohibited.
import enum
import inspect
from abc import ABC, abstractmethod
from dataclasses import dataclass
from typing import Optional, Union
import warnings
import cutlass.cute as cute
from cutlass.cutlass_dsl import (
Boolean,
Int32,
Int64,
if_generate,
dsl_user_op,
dsl_user_op,
)
from cutlass._mlir.dialects import llvm
import cutlass._mlir.dialects.cute as _cute_ir
from cutlass.cutlass_dsl import Boolean, Int32, if_generate, dsl_user_op
##############################################################################
@@ -111,7 +103,6 @@ class PipelineOp(enum.Enum):
# Async load without TMA
AsyncLoad = enum.auto()
def _get_pipeline_op(type_str):
return PipelineOp(type_str)
@@ -336,7 +327,9 @@ class MbarrierArray(SyncObject):
def arrive_and_expect_tx_with_dst(
self, index: int, tx_count: int, dst: Optional[int] = None, *, loc=None, ip=None
) -> None:
cute.arch.mbarrier_arrive_and_expect_tx(self.get_barrier(index), tx_count, dst, loc=loc, ip=ip)
cute.arch.mbarrier_arrive_and_expect_tx(
self.get_barrier(index, loc=loc, ip=ip), tx_count, dst, loc=loc, ip=ip
)
@dsl_user_op
def try_wait(self, index: int, phase: int, *, loc=None, ip=None) -> Boolean:
@@ -386,6 +379,14 @@ class MbarrierArray(SyncObject):
)
# Set explicit signature for Sphinx documentation to avoid issues with @dsl_user_op decorator
MbarrierArray.__init__.__signature__ = inspect.Signature(
[
inspect.Parameter("self", inspect.Parameter.POSITIONAL_OR_KEYWORD),
]
)
##############################################################################
# NamedBarrier class
##############################################################################
@@ -429,14 +430,11 @@ class NamedBarrier(SyncObject):
"""
The unaligned flavor of arrive can be used with an arbitrary number of threads in the CTA.
"""
llvm.inline_asm(
None,
[Int32(self.barrier_id).ir_value(), Int32(self.num_threads).ir_value()],
"barrier.arrive $0, $1;",
"r,r",
has_side_effects=True,
is_align_stack=False,
asm_dialect=llvm.AsmDialect.AD_ATT,
cute.arch.barrier_arrive(
barrier_id=self.barrier_id,
number_of_threads=self.num_threads,
loc=loc,
ip=ip,
)
@dsl_user_op
@@ -453,15 +451,13 @@ class NamedBarrier(SyncObject):
)
self.arrive_and_wait(loc=loc, ip=ip)
def wait_unaligned(self) -> None:
llvm.inline_asm(
None,
[Int32(self.barrier_id).ir_value(), Int32(self.num_threads).ir_value()],
"barrier.sync $0, $1;",
"r,r",
has_side_effects=True,
is_align_stack=False,
asm_dialect=llvm.AsmDialect.AD_ATT,
@dsl_user_op
def wait_unaligned(self, *, loc=None, ip=None) -> None:
cute.arch.barrier(
barrier_id=self.barrier_id,
number_of_threads=self.num_threads,
loc=loc,
ip=ip,
)
@dsl_user_op
@@ -751,18 +747,6 @@ def agent_sync(group: Agent, is_relaxed: bool = False, *, loc=None, ip=None):
)
def _mbarrier_i64_to_ptr(val: Int64) -> cute.Pointer:
"""
Converts a smem pointer of type Int64 to cute.Pointer with 8B alignment
"""
return cute.make_ptr(
Int64,
val.ir_value(),
mem_space=_cute_ir.AddressSpace.smem,
assumed_align=8,
)
# NamedBarrier free functions
@dsl_user_op
def arrive(barrier_id: int, num_threads: int, *, loc=None, ip=None):
@@ -780,19 +764,13 @@ def arrive_unaligned(barrier_id: int, num_threads: int, *, loc=None, ip=None):
"""
The unaligned flavor of arrive can be used with an arbitrary number of threads in the CTA.
"""
llvm.inline_asm(
None,
[Int32(barrier_id).ir_value(), Int32(num_threads).ir_value()],
"barrier.arrive $0, $1;",
"r,r",
has_side_effects=True,
is_align_stack=False,
asm_dialect=llvm.AsmDialect.AD_ATT,
cute.arch.barrier_arrive(
barrier_id=barrier_id, number_of_threads=num_threads, loc=loc, ip=ip
)
@dsl_user_op
def wait(barrier_id: int, num_threads: int):
def wait(*, loc=None, ip=None):
"""
NamedBarriers do not have a standalone wait like mbarriers, only an arrive_and_wait.
If synchronizing two warps in a producer/consumer pairing, the arrive count would be
@@ -811,14 +789,8 @@ def wait_unaligned(barrier_id: int, num_threads: int, *, loc=None, ip=None):
warnings.warn(
"NamedBarrier wait also arrives on the barrier. Routing call to NamedBarrier.arrive_and_wait()."
)
llvm.inline_asm(
None,
[Int32(barrier_id).ir_value(), Int32(num_threads).ir_value()],
"barrier.sync $0, $1;",
"r,r",
has_side_effects=True,
is_align_stack=False,
asm_dialect=llvm.AsmDialect.AD_ATT,
cute.arch.barrier(
barrier_id=barrier_id, number_of_threads=num_threads, loc=loc, ip=ip
)
+12 -3
View File
@@ -238,7 +238,10 @@ class PipelineTmaUmma(PipelineAsync):
if not defer_sync:
cute.arch.mbarrier_init_fence()
if cta_layout_vmnk is None or cute.size(cta_layout_vmnk, loc=loc, ip=ip) == 1:
if (
cta_layout_vmnk is None
or cute.size(cta_layout_vmnk, loc=loc, ip=ip) == 1
):
agent_sync(Agent.ThreadBlock)
else:
agent_sync(Agent.ThreadBlockCluster, is_relaxed=True)
@@ -449,7 +452,10 @@ class PipelineAsyncUmma(PipelineAsync):
if not defer_sync:
cute.arch.mbarrier_init_fence()
if cta_layout_vmnk is None or cute.size(cta_layout_vmnk, loc=loc, ip=ip) == 1:
if (
cta_layout_vmnk is None
or cute.size(cta_layout_vmnk, loc=loc, ip=ip) == 1
):
agent_sync(Agent.ThreadBlock)
else:
agent_sync(Agent.ThreadBlockCluster, is_relaxed=True)
@@ -587,7 +593,10 @@ class PipelineUmmaAsync(PipelineAsync):
if not defer_sync:
cute.arch.mbarrier_init_fence()
if cta_layout_vmnk is None or cute.size(cta_layout_vmnk, loc=loc, ip=ip) == 1:
if (
cta_layout_vmnk is None
or cute.size(cta_layout_vmnk, loc=loc, ip=ip) == 1
):
agent_sync(Agent.ThreadBlock)
else:
agent_sync(Agent.ThreadBlockCluster, is_relaxed=True)
+2 -1
View File
@@ -308,7 +308,7 @@ class PipelineAsync:
@dataclass(frozen=True)
class PipelineCpAsync(PipelineAsync):
"""
PipelineCpAsync is used for CpAsync producers and AsyncThread consumers (e.g. Hopper non-TMA mainloops).
PipelineCpAsync is used for CpAsync producers and AsyncThread consumers
"""
@staticmethod
@@ -656,6 +656,7 @@ class PipelineTmaMultiConsumersAsync(PipelineAsync):
)
if not defer_sync:
cute.arch.mbarrier_init_fence()
if cta_layout_vmnk is None or cute.size(cta_layout_vmnk) == 1:
agent_sync(Agent.ThreadBlock)
else:
+13 -5
View File
@@ -70,7 +70,6 @@ from .tmem_allocator import TmemAllocator, get_num_tmem_alloc_cols
from .layout import LayoutEnum
from . import gemm
from . import distributed
from .mixed_input_helpers import (
@@ -99,11 +98,17 @@ from .mixed_input_helpers import (
store_transformed_a,
)
from . import gemm
from . import hopper_helpers as sm90
from . import blackwell_helpers as sm100
from .print_latex import print_latex, print_latex_tv
from .tensor_helpers import (
is_fp8_dtype,
create_cute_tensor_for_fp8,
)
__all__ = [
"get_smem_capacity_in_bytes",
"SmemAllocator",
@@ -135,6 +140,7 @@ __all__ = [
"get_divisibility",
"epilogue_tma_store",
"epilogue",
"create_tensor_a",
"compute_epilogue_tile_shape",
"get_smem_store_op",
"get_tmem_load_op",
@@ -145,10 +151,12 @@ __all__ = [
"make_blockscaled_trivial_tiled_mma",
"sm90",
"sm100",
"print_latex",
"print_latex_tv",
"gemm",
"distributed",
"ClcDynamicPersistentTileSchedulerParams",
"ClcDynamicPersistentTileScheduler",
"print_latex",
"print_latex_tv",
"is_fp8_dtype",
"create_cute_tensor_for_fp8",
"distributed",
]
@@ -658,7 +658,11 @@ def make_smem_layout_a(
:rtype: Union[cute.Layout, cute.ComposedLayout]
"""
is_k_major = (tiled_mma.op.a_major_mode == OperandMajorMode.K) if is_k_major is None else is_k_major
is_k_major = (
(tiled_mma.op.a_major_mode == OperandMajorMode.K)
if is_k_major is None
else is_k_major
)
a_major_mode = OperandMajorMode.K if is_k_major else OperandMajorMode.MN
a_smem_shape = tiled_mma.partition_shape_A(
cute.dice(mma_tiler_mnk, (1, None, 1), loc=loc, ip=ip), loc=loc, ip=ip
@@ -712,7 +716,11 @@ def make_smem_layout_b(
:rtype: Union[cute.Layout, cute.ComposedLayout]
"""
is_k_major = (tiled_mma.op.b_major_mode == OperandMajorMode.K) if is_k_major is None else is_k_major
is_k_major = (
(tiled_mma.op.b_major_mode == OperandMajorMode.K)
if is_k_major is None
else is_k_major
)
b_major_mode = OperandMajorMode.K if is_k_major else OperandMajorMode.MN
b_smem_shape = tiled_mma.partition_shape_B(
cute.dice(mma_tiler_mnk, (None, 1, 1), loc=loc, ip=ip), loc=loc, ip=ip
@@ -84,6 +84,38 @@ def tile_atom_to_shape_SF(
return sf_layout
@dsl_user_op
def make_smem_layout_sf(
tile_shape: cute.Tile,
sf_vec_size: int,
num_stages: int,
*,
loc=None,
ip=None,
) -> cute.Layout:
"""
A helper function to get dynamic SFA/SFB layout by filling dynamic A/B shape to the scale factor atom layout.
:param Shape: The shape of the A/B tensor
:param sf_vec_size: Scale factor vector size
:param num_stages: Number of stages
:return: The layout of the SFA/SFB tensor
:rtype: cute.Layout
"""
smem_layout = cute.tile_to_shape(
BlockScaledBasicChunk(sf_vec_size).layout, tile_shape, (2, 1)
)
smem_layout_staged = cute.append(
smem_layout,
cute.make_layout(
num_stages, stride=cute.cosize(cute.filter_zeros(smem_layout))
),
)
return smem_layout_staged
@dsl_user_op
def make_smem_layout_sfa(
tiled_mma: cute.TiledMma,
@@ -214,6 +246,176 @@ def make_smem_layout_sfb(
return sfb_smem_layout_staged
@dsl_user_op
def sm120_make_smem_layout_sfa(
tiled_mma: cute.TiledMma,
tile_shape_mnk: cute.Tile,
sf_vec_size: int,
num_stages: int,
*,
loc=None,
ip=None,
) -> cute.Layout:
"""
Make smem layout for SFA based on:
1. BlockScaledBasicChunk
2. MMA tiler shape
3. Scale factor vector size
4. Number of stages
:param tiled_mma: The tiled MMA
:type tiled_mma: cute.TiledMma
:param mma_tiler_mnk: The mma tiler shape
:type mma_tiler_mnk: cute.Tile
:param sf_vec_size: The scale factor vector size
:type sf_vec_size: int
:param num_stages: The number of stages
:type num_stages: int
:return: Smem layout for SFA
:rtype: cute.Layout
"""
assert sf_vec_size == 16 or sf_vec_size == 32, "sf_vec_size must be 16 or 32"
blk_mn = 128
blk_sf = 4
blk_elems = blk_mn * blk_sf
mma_nsf = tiled_mma.shape_mnk[2] // sf_vec_size
mn_basic_block_shape = (32, 4)
mn_basic_block_stride = (16, 4)
k_basic_block_shape = (sf_vec_size, mma_nsf)
k_basic_block_stride = (0, 1)
assert tile_shape_mnk[0] % blk_mn == 0, (
"tile_shape_mnk[0] must be divisible by blk_mn"
)
sSFA_shapeM = (mn_basic_block_shape, tile_shape_mnk[0] // blk_mn)
sSF_strideM = (mn_basic_block_stride, blk_elems)
assert tile_shape_mnk[2] % (blk_sf * mma_nsf) == 0, (
"tile_shape_mnk[2] must be divisible by blk_sf * mma_nsf"
)
sSFA_shapeK = (
k_basic_block_shape,
blk_sf // mma_nsf,
tile_shape_mnk[2] // sf_vec_size // blk_sf,
)
sSF_strideK = (
k_basic_block_stride,
mma_nsf,
tile_shape_mnk[0] // blk_mn * blk_elems,
)
sSFA_shape = (sSFA_shapeM, sSFA_shapeK)
sSFA_stride = (sSF_strideM, sSF_strideK)
smem_layout = cute.make_layout(sSFA_shape, stride=sSFA_stride)
# (((Atom_Inst_M, Rest_M),(Atom_Inst_K, Rest_K)), MMA_M, MMA_K, STAGE)
sfa_smem_layout_staged = cute.append(
smem_layout,
cute.make_layout(
num_stages, stride=cute.cosize(cute.filter_zeros(smem_layout))
),
)
return sfa_smem_layout_staged
@dsl_user_op
def sm120_make_smem_layout_sfb(
tiled_mma: cute.TiledMma,
tile_shape_mnk: cute.Tile,
sf_vec_size: int,
num_stages: int,
*,
loc=None,
ip=None,
) -> cute.Layout:
"""
Make smem layout for SFB based on:
1. BlockScaledBasicChunk
2. MMA tiler shape
3. Scale factor vector size
4. Number of stages
:param tiled_mma: The tiled MMA
:type tiled_mma: cute.TiledMma
:param mma_tiler_mnk: The mma tiler shape
:type mma_tiler_mnk: cute.Tile
:param sf_vec_size: The scale factor vector size
:type sf_vec_size: int
:param num_stages: The number of stages
:type num_stages: int
:return: Smem layout for SFA
:rtype: cute.Layout
"""
# A single indivisible block will hold 4 scale factors of 128 rows/columns (A/B matrix).
# 4 is chosen to make consecutive 32bits of data to have scale factors for only a single row(col).
blk_mn = 128
blk_sf = 4
blk_elems = blk_mn * blk_sf
assert sf_vec_size == 16 or sf_vec_size == 32, "sf_vec_size must be 16 or 32"
assert tile_shape_mnk[1] % blk_mn == 0, (
"tile_shape_mnk[1] must be divisible by blk_mn"
)
assert tile_shape_mnk[2] % sf_vec_size == 0, (
"tile_shape_mnk[2] must be divisible by sf_vec_size"
)
mma_nsf = tiled_mma.shape_mnk[2] // sf_vec_size
mn_basic_block_shape = (32, 4)
mn_basic_block_stride = (16, 4)
k_basic_block_shape = (sf_vec_size, mma_nsf)
k_basic_block_stride = (0, 1)
assert tile_shape_mnk[1] % blk_mn == 0, (
"tile_shape_mnk[1] must be divisible by blk_mn"
)
sSFA_shapeN = (mn_basic_block_shape, tile_shape_mnk[1] // blk_mn)
sSF_strideN = (mn_basic_block_stride, blk_elems)
assert tile_shape_mnk[2] % (blk_sf * mma_nsf) == 0, (
"tile_shape_mnk[2] must be divisible by blk_sf * mma_nsf"
)
sSFA_shapeK = (
k_basic_block_shape,
blk_sf // mma_nsf,
tile_shape_mnk[2] // sf_vec_size // blk_sf,
)
sSF_strideK = (
k_basic_block_stride,
mma_nsf,
tile_shape_mnk[1] // blk_mn * blk_elems,
)
sSFA_shape = (sSFA_shapeN, sSFA_shapeK)
sSFA_stride = (sSF_strideN, sSF_strideK)
smem_layout = cute.make_layout(sSFA_shape, stride=sSFA_stride)
# (((Atom_Inst_M, Rest_M),(Atom_Inst_K, Rest_K)), MMA_M, MMA_K, STAGE)
sfb_smem_layout_staged = cute.append(
smem_layout,
cute.make_layout(
num_stages, stride=cute.cosize(cute.filter_zeros(smem_layout))
),
)
return sfb_smem_layout_staged
@dsl_user_op
def make_tmem_layout_sfa(
+2 -1
View File
@@ -71,6 +71,7 @@ def ld_bypass(input_tensor: cute.Tensor):
@dsl_user_op
def multimem_red_release_gpu_add1(
lock_ptr: Pointer,
*,
loc=None,
ip=None,
) -> None:
@@ -89,6 +90,7 @@ def multimem_red_release_gpu_add1(
@dsl_user_op
def multimem_red_release_sys_add1(
lock_ptr: Pointer,
*,
loc=None,
ip=None,
) -> None:
@@ -285,7 +287,6 @@ def spin_lock_atom_cas_relaxed_wait(
ip=ip,
)
########################################################
# Multimem Load & Store
########################################################
@@ -27,6 +27,7 @@ from cutlass.utils.static_persistent_tile_scheduler import (
)
import cutlass.cute as cute
class ClcDynamicPersistentTileSchedulerParams:
"""A class to represent parameters for a dynamic persistent tile scheduler.
@@ -98,6 +99,7 @@ class ClcDynamicPersistentTileSchedulerParams:
)
return problem_ceiling_cta_mnl
class ClcDynamicPersistentTileScheduler:
"""A scheduler for dynamic persistent tile execution in CUTLASS/CuTe kernels.
@@ -243,7 +245,10 @@ class ClcDynamicPersistentTileScheduler:
result_addr: 16-byte response data (simulating shared memory access)
"""
m_idx, n_idx, l_idx, vld = cute.arch.clc_response(result_addr, loc=loc, ip=ip)
cute.arch.fence_proxy("async.shared", space="cta")
cute.arch.fence_proxy(
"async.shared",
space="cta",
)
cta_idx_in_cluster, cta_idy_in_cluster, _ = self.cta_id_in_cluster
cur_tile_coord = (m_idx + cta_idx_in_cluster, n_idx + cta_idy_in_cluster, l_idx)
return WorkTileInfo(cur_tile_coord, vld)
@@ -7,7 +7,7 @@
#
# Any use, reproduction, disclosure, or distribution of this software
# and related documentation outside the scope permitted by the EULA
# is strictly prohibited
# is strictly prohibited.
from . import sm100
+136 -229
View File
@@ -14,9 +14,6 @@ import cutlass.cute as cute
from cutlass.cutlass_dsl import Int32, Boolean, Constexpr, const_expr
import cutlass.pipeline as pipeline
from cutlass.utils.static_persistent_tile_scheduler import StaticPersistentTileScheduler
from cutlass.utils.dynamic_persistent_tile_scheduler import (
ClcDynamicPersistentTileScheduler,
)
from cutlass.utils.blackwell_helpers import get_tmem_load_op, get_smem_store_op
from cutlass.cute.nvgpu import cpasync, tcgen05
from cutlass.cute.nvgpu.common import CacheEvictionPriority
@@ -161,8 +158,6 @@ def epilogue_tma_store(
gemm_kernel,
epi_tidx: Int32,
warp_idx: Int32,
acc_pipeline: pipeline.PipelineAsync,
tiled_mma: cute.TiledMma,
tma_atom_c: cute.CopyAtom,
# Input of epilogue
tCtAcc_base: cute.Tensor,
@@ -171,11 +166,13 @@ def epilogue_tma_store(
# Output of epilogue
tCgC_base: cute.Tensor,
epi_tile: cute.Tile,
tile_sched: Union[StaticPersistentTileScheduler, ClcDynamicPersistentTileScheduler],
num_tiles_executed: Int32,
epilogue_op: Constexpr,
clc_pipeline: Union[pipeline.PipelineClcFetchAsync, None] = None,
clc_consumer_state: Union[pipeline.PipelineState, None] = None,
) -> None:
mma_tile_coord_mnl: Tuple[Int32, Int32, Int32],
acc_consumer_state: pipeline.PipelineState,
acc_pipeline: pipeline.PipelineAsync,
c_pipeline: pipeline.PipelineTmaStore,
) -> pipeline.PipelineState:
# Layout transformation for tCgC_base
# ((MMA_ATOM_M, MMA_ATOM_N), MMA_M, MMA_N, TILE_M, TILE_N, TILE_K)
# -> ((MMA_ATOM_M, MMA_M), (MMA_ATOM_N, MMA_N), TILE_M, TILE_N, TILE_K)
@@ -207,142 +204,96 @@ def epilogue_tma_store(
cute.group_modes(tCgC_epi, 0, 2),
)
acc_consumer_state = pipeline.make_pipeline_state(
pipeline.PipelineUserType.Consumer, gemm_kernel.num_acc_stage
)
# Threads/warps participating in tma store pipeline
c_producer_group = pipeline.CooperativeGroup(
pipeline.Agent.Thread,
32 * len(gemm_kernel.epilogue_warp_id),
)
c_pipeline = pipeline.PipelineTmaStore.create(
num_stages=gemm_kernel.num_c_stage, producer_group=c_producer_group
)
epilog_sync_barrier = pipeline.NamedBarrier(
barrier_id=gemm_kernel.epilog_sync_bar_id,
num_threads=32 * len(gemm_kernel.epilogue_warp_id),
)
work_tile = tile_sched.initial_work_tile_info()
while work_tile.is_valid_tile:
# Get tile coord from tile scheduler
cur_tile_coord = work_tile.tile_idx
mma_tile_coord_mnl = (
cur_tile_coord[0] // cute.size(tiled_mma.thr_id.shape),
cur_tile_coord[1],
cur_tile_coord[2],
)
#
# Slice to per mma tile index
#
# ((ATOM_V, REST_V), EPI_M, EPI_N)
bSG_gC = bSG_gC_partitioned[(None, None, None, *mma_tile_coord_mnl)]
# Set tensor memory buffer for current tile
# (T2R, T2R_M, T2R_N, EPI_M, EPI_N)
tTR_tAcc = tTR_tAcc_base[(None, None, None, None, None, acc_consumer_state.index)]
#
# Wait for accumulator buffer full
#
acc_pipeline.consumer_wait(acc_consumer_state)
tTR_tAcc = cute.group_modes(tTR_tAcc, 3, cute.rank(tTR_tAcc))
bSG_gC = cute.group_modes(bSG_gC, 1, cute.rank(bSG_gC))
#
# Store accumulator to global memory in subtiles
#
subtile_cnt = cute.size(tTR_tAcc.shape, mode=[3])
num_prev_subtiles = num_tiles_executed * subtile_cnt
for subtile_idx in range(subtile_cnt):
#
# Load accumulator from tensor memory buffer to register
#
tTR_tAcc_mn = tTR_tAcc[(None, None, None, subtile_idx)]
cute.copy(tiled_copy_t2r, tTR_tAcc_mn, tTR_rAcc)
#
# Slice to per mma tile index
# Convert to C type
#
# ((ATOM_V, REST_V), EPI_M, EPI_N)
bSG_gC = bSG_gC_partitioned[(None, None, None, *mma_tile_coord_mnl)]
# Set tensor memory buffer for current tile
# (T2R, T2R_M, T2R_N, EPI_M, EPI_N)
tTR_tAcc = tTR_tAcc_base[
(None, None, None, None, None, acc_consumer_state.index)
]
acc_vec = tiled_copy_r2s.retile(tTR_rAcc).load()
acc_vec = epilogue_op(acc_vec.to(gemm_kernel.c_dtype))
tRS_rC.store(acc_vec)
#
# Wait for accumulator buffer full
# Store C to shared memory
#
acc_pipeline.consumer_wait(acc_consumer_state)
tTR_tAcc = cute.group_modes(tTR_tAcc, 3, cute.rank(tTR_tAcc))
bSG_gC = cute.group_modes(bSG_gC, 1, cute.rank(bSG_gC))
#
# Store accumulator to global memory in subtiles
#
subtile_cnt = cute.size(tTR_tAcc.shape, mode=[3])
num_prev_subtiles = tile_sched.num_tiles_executed * subtile_cnt
for subtile_idx in range(subtile_cnt):
#
# Load accumulator from tensor memory buffer to register
#
tTR_tAcc_mn = tTR_tAcc[(None, None, None, subtile_idx)]
cute.copy(tiled_copy_t2r, tTR_tAcc_mn, tTR_rAcc)
#
# Convert to C type
#
acc_vec = tiled_copy_r2s.retile(tTR_rAcc).load()
acc_vec = epilogue_op(acc_vec.to(gemm_kernel.c_dtype))
tRS_rC.store(acc_vec)
#
# Store C to shared memory
#
c_buffer = (num_prev_subtiles + subtile_idx) % gemm_kernel.num_c_stage
cute.copy(tiled_copy_r2s, tRS_rC, tRS_sC[(None, None, None, c_buffer)])
# Fence and barrier to make sure shared memory store is visible to TMA store
cute.arch.fence_proxy("async.shared", space="cta")
epilog_sync_barrier.arrive_and_wait()
#
# TMA store C to global memory
#
if warp_idx == gemm_kernel.epilogue_warp_id[0]:
cute.copy(
tma_atom_c,
bSG_sC[(None, c_buffer)],
bSG_gC[(None, subtile_idx)],
)
# Fence and barrier to make sure shared memory store is visible to TMA store
c_pipeline.producer_commit()
c_pipeline.producer_acquire()
epilog_sync_barrier.arrive_and_wait()
c_buffer = (num_prev_subtiles + subtile_idx) % gemm_kernel.num_c_stage
cute.copy(tiled_copy_r2s, tRS_rC, tRS_sC[(None, None, None, c_buffer)])
# Fence and barrier to make sure shared memory store is visible to TMA store
cute.arch.fence_proxy("async.shared", space="cta")
epilog_sync_barrier.arrive_and_wait()
#
# Async arrive accumulator buffer empty
# TMA store C to global memory
#
with cute.arch.elect_one():
acc_pipeline.consumer_release(acc_consumer_state)
acc_consumer_state.advance()
if warp_idx == gemm_kernel.epilogue_warp_id[0]:
cute.copy(
tma_atom_c,
bSG_sC[(None, c_buffer)],
bSG_gC[(None, subtile_idx)],
)
# Fence and barrier to make sure shared memory store is visible to TMA store
c_pipeline.producer_commit()
c_pipeline.producer_acquire()
epilog_sync_barrier.arrive_and_wait()
#
# Advance to next tile
#
# Check if tile_sched is StaticPersistentTileScheduler or any subclass inheriting from it
if const_expr(isinstance(tile_sched, StaticPersistentTileScheduler)):
tile_sched.advance_to_next_work()
work_tile = tile_sched.get_current_work()
elif const_expr(isinstance(tile_sched, ClcDynamicPersistentTileScheduler)):
clc_pipeline.consumer_wait(clc_consumer_state)
work_tile = tile_sched.get_current_work()
clc_pipeline.consumer_release(clc_consumer_state)
clc_consumer_state.advance()
else:
# Not match
pass
epilog_sync_barrier.arrive_and_wait()
# Wait for C store complete
c_pipeline.producer_tail()
#
# Async arrive accumulator buffer empty
#
with cute.arch.elect_one():
acc_pipeline.consumer_release(acc_consumer_state)
acc_consumer_state.advance()
return acc_consumer_state
@cute.jit
def epilogue(
gemm_kernel,
epi_tidx: Int32,
acc_pipeline: pipeline.PipelineAsync,
tiled_mma: cute.TiledMma,
tCtAcc_base: cute.Tensor,
tCgC_base: cute.Tensor,
epi_tile: cute.Tile,
tile_sched: Union[StaticPersistentTileScheduler, ClcDynamicPersistentTileScheduler],
epilogue_op: Constexpr,
tmem_dealloc_barrier: pipeline.NamedBarrier,
mma_tile_coord_mnl: Tuple[Int32, Int32, Int32],
acc_consumer_state: pipeline.PipelineState,
acc_pipeline: pipeline.PipelineAsync,
tCcC_base: cute.Tensor = None,
mC_mnl: cute.Tensor = None,
clc_pipeline: Union[pipeline.PipelineClcFetchAsync, None] = None,
clc_consumer_state: Union[pipeline.PipelineState, None] = None,
) -> None:
) -> pipeline.PipelineState:
"""
Epilogue function that stores accumulator results directly to global memory.
Used when TMA store is not enabled.
@@ -351,32 +302,20 @@ def epilogue(
:type gemm_kernel: Any
:param epi_tidx: Thread index in epilogue warp groups
:type epi_tidx: Int32
:param acc_pipeline: Accumulator pipeline for async operations
:type acc_pipeline: pipeline.PipelineAsync
:param tiled_mma: The tiled MMA configuration
:type tiled_mma: cute.TiledMma
:param tCtAcc_base: Base accumulator tensor in tensor memory
:type tCtAcc_base: cute.Tensor
:param tCgC_base: The global memory tensor C to be copied and partitioned
:type tCgC_base: cute.Tensor
:param epi_tile: Epilogue tile configuration
:type epi_tile: cute.Tile
:param tile_sched: Tile scheduler for persistent scheduling
:type tile_sched: StaticPersistentTileScheduler
:param epilogue_op: Optional elementwise operation to apply
:type epilogue_op: Constexpr
:param tmem_dealloc_barrier: Barrier for tensor memory deallocation
:type tmem_dealloc_barrier: pipeline.NamedBarrier
:param alignment_bytes: Alignment bytes for global memory store
:type alignment_bytes: int
:param tCcC_base: Identity/coordinate tensor C
:type tCcC_base: cute.Tensor
:param mC_mnl: Global memory tensor C (full tensor for predicate computation)
:type mC_mnl: cute.Tensor
:param clc_pipeline: Pipeline for dynamic persistent tile scheduling
:type clc_pipeline: Union[pipeline.PipelineClcFetchAsync, None]
:param clc_consumer_state: Consumer state for dynamic persistent tile scheduling
:type clc_consumer_state: Union[pipeline.PipelineState, None]
"""
# Layout transformation for tCgC_base
@@ -434,29 +373,21 @@ def epilogue(
cC_epi = cute.flat_divide(tCcC, epi_tile)
tTR_cC_partitioned = thr_copy_t2r.partition_D(cC_epi)
acc_consumer_state = pipeline.make_pipeline_state(
pipeline.PipelineUserType.Consumer, gemm_kernel.num_acc_stage
)
work_tile = tile_sched.initial_work_tile_info()
while work_tile.is_valid_tile:
#
# Pre-advance to next tile
#
if const_expr(isinstance(tile_sched, StaticPersistentTileScheduler)):
tile_sched.advance_to_next_work()
next_work_tile = tile_sched.get_current_work()
# Get tile coord from current work tile
cur_tile_coord = work_tile.tile_idx
mma_tile_coord_mnl = (
cur_tile_coord[0] // cute.size(tiled_mma.thr_id.shape),
cur_tile_coord[1],
cur_tile_coord[2],
# (T2R, T2R_M, T2R_N, EPI_M, EPI_N)
tTR_gC = tTR_gC_partitioned[
(
None,
None,
None,
None,
None,
*mma_tile_coord_mnl,
)
]
if const_expr(use_predication):
# (T2R, T2R_M, T2R_N, EPI_M, EPI_N)
tTR_gC = tTR_gC_partitioned[
tTR_cC = tTR_cC_partitioned[
(
None,
None,
@@ -466,88 +397,65 @@ def epilogue(
*mma_tile_coord_mnl,
)
]
tTR_cC = cute.group_modes(tTR_cC, 3, cute.rank(tTR_cC))
# Set tensor memory buffer for current tile
# (T2R, T2R_M, T2R_N, EPI_M, EPI_M)
tTR_tAcc = tTR_tAcc_base[(None, None, None, None, None, acc_consumer_state.index)]
#
# Wait for accumulator buffer full
#
acc_pipeline.consumer_wait(acc_consumer_state)
tTR_tAcc = cute.group_modes(tTR_tAcc, 3, cute.rank(tTR_tAcc))
tTR_gC = cute.group_modes(tTR_gC, 3, cute.rank(tTR_gC))
#
# Store accumulator to global memory in subtiles
#
subtile_cnt = cute.size(tTR_tAcc.shape, mode=[3])
for subtile_idx in range(subtile_cnt):
#
# Get the destination and coordinate slices for this subtile
#
tTR_gC_subtile = tTR_gC[(None, None, None, subtile_idx)]
#
# Load accumulator from tensor memory buffer to register
#
tTR_tAcc_mn = tTR_tAcc[(None, None, None, subtile_idx)]
cute.copy(tiled_copy_t2r, tTR_tAcc_mn, tTR_rAcc)
# Async arrive accumulator buffer empty
# Release early for perf
if subtile_idx == subtile_cnt - 1:
with cute.arch.elect_one():
acc_pipeline.consumer_release(acc_consumer_state)
acc_consumer_state.advance()
#
# Convert to C type
#
acc_vec = tTR_rAcc.load()
acc_vec = epilogue_op(acc_vec.to(gemm_kernel.c_dtype))
tTR_rC.store(acc_vec)
if const_expr(use_predication):
# (T2R, T2R_M, T2R_N, EPI_M, EPI_N)
tTR_cC = tTR_cC_partitioned[
(
None,
None,
None,
None,
None,
*mma_tile_coord_mnl,
)
]
tTR_cC = cute.group_modes(tTR_cC, 3, cute.rank(tTR_cC))
# compute predicate
tTR_cC_subtile = tTR_cC[(None, None, None, subtile_idx)]
pred_C_shape = (1, *tTR_cC_subtile.shape[1:])
pred_C = cute.make_rmem_tensor(pred_C_shape, Boolean)
for m_idx in range(tTR_cC_subtile.shape[1]):
for n_idx in range(tTR_cC_subtile.shape[2]):
vector_first_coord = tTR_cC_subtile[(0, m_idx, n_idx)]
pred_C[(0, m_idx, n_idx)] = cute.elem_less(
vector_first_coord, mC_mnl.shape
)
# Store C to global memory with predication
cute.copy(simt_atom, tTR_rC, tTR_gC_subtile, pred=pred_C)
else:
# Store C directly to global memory
cute.copy(simt_atom, tTR_rC, tTR_gC_subtile)
# Set tensor memory buffer for current tile
# (T2R, T2R_M, T2R_N, EPI_M, EPI_M)
tTR_tAcc = tTR_tAcc_base[
(None, None, None, None, None, acc_consumer_state.index)
]
#
# Wait for accumulator buffer full
#
acc_pipeline.consumer_wait(acc_consumer_state)
tTR_tAcc = cute.group_modes(tTR_tAcc, 3, cute.rank(tTR_tAcc))
tTR_gC = cute.group_modes(tTR_gC, 3, cute.rank(tTR_gC))
#
# Store accumulator to global memory in subtiles
#
subtile_cnt = cute.size(tTR_tAcc.shape, mode=[3])
for subtile_idx in range(subtile_cnt):
#
# Get the destination and coordinate slices for this subtile
#
tTR_gC_subtile = tTR_gC[(None, None, None, subtile_idx)]
#
# Load accumulator from tensor memory buffer to register
#
tTR_tAcc_mn = tTR_tAcc[(None, None, None, subtile_idx)]
cute.copy(tiled_copy_t2r, tTR_tAcc_mn, tTR_rAcc)
# Async arrive accumulator buffer empty
# Release early for perf
if subtile_idx == subtile_cnt - 1:
with cute.arch.elect_one():
acc_pipeline.consumer_release(acc_consumer_state)
acc_consumer_state.advance()
#
# Convert to C type
#
acc_vec = tTR_rAcc.load()
acc_vec = epilogue_op(acc_vec.to(gemm_kernel.c_dtype))
tTR_rC.store(acc_vec)
if const_expr(use_predication):
# compute predicate
tTR_cC_subtile = tTR_cC[(None, None, None, subtile_idx)]
pred_C_shape = (1, *tTR_cC_subtile.shape[1:])
pred_C = cute.make_rmem_tensor(pred_C_shape, Boolean)
for m_idx in range(tTR_cC_subtile.shape[1]):
for n_idx in range(tTR_cC_subtile.shape[2]):
vector_first_coord = tTR_cC_subtile[(0, m_idx, n_idx)]
pred_C[(0, m_idx, n_idx)] = cute.elem_less(
vector_first_coord, mC_mnl.shape
)
# Store C to global memory with predication
cute.copy(simt_atom, tTR_rC, tTR_gC_subtile, pred=pred_C)
else:
# Store C directly to global memory
cute.copy(simt_atom, tTR_rC, tTR_gC_subtile)
if const_expr(isinstance(tile_sched, StaticPersistentTileScheduler)):
work_tile = next_work_tile
elif const_expr(isinstance(tile_sched, ClcDynamicPersistentTileScheduler)):
clc_pipeline.consumer_wait(clc_consumer_state)
work_tile = tile_sched.get_current_work()
clc_pipeline.consumer_release(clc_consumer_state)
clc_consumer_state.advance()
# Synchronize before TMEM dealloc (done by the caller)
tmem_dealloc_barrier.arrive_and_wait()
return acc_consumer_state
@cute.jit
@@ -908,4 +816,3 @@ def epilogue_release_flag(
# Synchronize before TMEM dealloc (done by the caller)
tmem_dealloc_barrier.arrive_and_wait()
@@ -178,13 +178,17 @@ class HardwareInfo:
# Create a temporary directory for dumping artifacts
with tempfile.TemporaryDirectory() as temp_dir:
# keep-cubin will keep the cubin in the artifacts
compiled_func = cute.compile(self._host_function, options=f"--dump-dir={temp_dir} --keep-cubin")
compiled_func = cute.compile(
self._host_function, options=f"--dump-dir={temp_dir} --keep-cubin"
)
# Get the CUBIN from artifacts
cubin_data = compiled_func.artifacts.CUBIN
cuda_library = self._checkCudaErrors(
driver.cuLibraryLoadData(cubin_data, None, None, 0, None, None, 0)
)
# Enumerate kernels from the library
kernels = self._checkCudaErrors(driver.cuLibraryEnumerateKernels(1, cuda_library))
kernels = self._checkCudaErrors(
driver.cuLibraryEnumerateKernels(1, cuda_library)
)
# Get the function from the kernel
return self._checkCudaErrors(driver.cuKernelGetFunction(kernels[0]))
@@ -513,9 +513,9 @@ def get_smem_layout_scale(
cute.size(mma_tiler[2]) % cute.size(smem_layout_scale_per_stage.outer[1]) == 0
), "smem_layout_scale_per_stage must evenly divide tile k shape."
# Shared memory buffer for scale must be at least 128B to satisfy TMA requirement
assert (
cute.size_in_bytes(a_scale_dtype, smem_layout_scale_per_stage) >= 128
), "smem size for scale must be at least 128B"
assert cute.size_in_bytes(a_scale_dtype, smem_layout_scale_per_stage) >= 128, (
"smem size for scale must be at least 128B"
)
# Scale layout in smem with multiple stages
smem_layout_scale = cute.append(
smem_layout_scale_per_stage,
@@ -972,10 +972,9 @@ def cvt_tensor_a(
for int4-to-bf16 conversion.
"""
from cutlass import CUDA_VERSION
# shuffle is supported since CUDA 13.1
shuffle_supported = True
if CUDA_VERSION.major < 13 or (CUDA_VERSION == 13 and CUDA_VERSION.minor < 1):
if CUDA_VERSION.major < 13 or (CUDA_VERSION.major == 13 and CUDA_VERSION.minor < 1):
shuffle_supported = False
shuffle = shuffle and shuffle_supported
rst = src.load()
@@ -80,7 +80,7 @@ class SmemAllocator:
GPU compute capability.
:param compute_capability: The compute capability string (e.g. "70", "75", "80")
:type compute_capability: str
:type compute_capability: Optional[str]
:return: The shared memory capacity in bytes
:rtype: int
:raises ValueError: If the compute capability is not supported
@@ -9,6 +9,7 @@
# and related documentation outside the scope permitted by the EULA
# is strictly prohibited.
import inspect
from typing import Tuple
from cutlass.cutlass_dsl import (
@@ -325,6 +326,14 @@ class PersistentTileSchedulerParams:
return (*self.cluster_shape_mn, num_persistent_clusters)
# Set explicit signature for Sphinx documentation to avoid issues with @dsl_user_op decorator
PersistentTileSchedulerParams.__init__.__signature__ = inspect.Signature(
[
inspect.Parameter("self", inspect.Parameter.POSITIONAL_OR_KEYWORD),
]
)
class StaticPersistentTileScheduler:
"""A scheduler for static persistent tile execution in CUTLASS/CuTe kernels.
@@ -0,0 +1,63 @@
# SPDX-FileCopyrightText: Copyright (c) 2025 - 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: LicenseRef-NvidiaProprietary
#
# Use of this software is governed by the terms and conditions of the
# NVIDIA End User License Agreement (EULA), available at:
# https://docs.nvidia.com/cutlass/media/docs/pythonDSL/license.html
#
# Any use, reproduction, disclosure, or distribution of this software
# and related documentation outside the scope permitted by the EULA
# is strictly prohibited.
"""Utility functions for tensor creation and type handling."""
from typing import Type, Optional
# Import only the specific types needed to avoid circular import with cutlass module
from cutlass.cute.typing import Float8E5M2, Float8E4M3FN, TFloat32, Numeric
from cutlass.cute.runtime import from_dlpack
def is_fp8_dtype(dtype: Type[Numeric]) -> bool:
"""Check if dtype is a float8 type that doesn't support dlpack.
params dtype: The cutlass numeric type to check
type dtype: Type[cutlass.Numeric]
return: True if the dtype is Float8E5M2 or Float8E4M3FN, False otherwise
"""
return dtype in {Float8E5M2, Float8E4M3FN}
def create_cute_tensor_for_fp8(
storage_tensor,
dtype: Type[Numeric],
leading_dim: int,
source_f32_tensor=None,
):
"""Create cute tensor, handling float8 types that don't support dlpack.
For float8 types, the storage_tensor should be uint8 (for DLPack compatibility).
The source_f32_tensor provides the actual float32 values to convert to fp8.
params storage_tensor: Tensor for DLPack (uint8 for fp8, otherwise the actual dtype)
params dtype: Target cutlass dtype
params leading_dim: Leading dimension for dynamic layout
paramas source_f32_tensor: Float32 source data for fp8 conversion (required for fp8)
return: A cute tensor with the appropriate dtype and layout
"""
import cutlass.torch as cutlass_torch
cute_tensor = from_dlpack(
storage_tensor, assumed_align=16, force_tf32=dtype == TFloat32
)
# For float8 types, set element_type explicitly since storage is uint8
if is_fp8_dtype(dtype):
cute_tensor.element_type = dtype
cute_tensor = cute_tensor.mark_layout_dynamic(leading_dim=leading_dim)
# For float8 types, convert data from float32 using GPU kernel
if is_fp8_dtype(dtype):
if source_f32_tensor is None:
raise ValueError("source_f32_tensor is required for fp8 types")
cute_tensor = cutlass_torch.convert_cute_tensor(
source_f32_tensor, cute_tensor, dtype, is_dynamic_layout=True
)
return cute_tensor
@@ -9,8 +9,8 @@
# and related documentation outside the scope permitted by the EULA
# is strictly prohibited.
from math import log2, ceil
from typing import Optional, Type, Union, List
from math import ceil, log2
import inspect
from cutlass import const_expr
@@ -29,7 +29,7 @@ from cutlass.cute.arch import get_max_tmem_alloc_cols, get_min_tmem_alloc_cols
class TmemAllocator:
"""A class for managing tensor memory allocation.
"""A class for managing tensor memory allocation on GPUs.
This class manages allocation/deallocation of tensor memory, including the mbarrier
synchronization for two cta use case.
@@ -81,6 +81,7 @@ class TmemAllocator:
two_cta_tmem_dealloc_mbar_ptr: Optional[cute.Pointer] = None,
*,
arch: str = "sm_100",
dealloc_mbarrier_initialized: bool = False,
loc=None,
ip=None,
):
@@ -126,7 +127,7 @@ class TmemAllocator:
self._max_tmem_columns = get_max_tmem_alloc_cols(arch)
# Init tmem dealloc mbarrier if two cta
if const_expr(self._is_two_cta):
if not dealloc_mbarrier_initialized and const_expr(self._is_two_cta):
self._init_dealloc_mbarrier(loc=loc, ip=ip)
def __extract_mlir_values__(self) -> list[ir.Value]:
@@ -158,7 +159,8 @@ class TmemAllocator:
self._is_two_cta,
self._num_allocated_columns,
new_two_cta_tmem_dealloc_mbar_ptr,
arch=self._arch,
arch=self._arch, # Preserve the architecture parameter
dealloc_mbarrier_initialized=True,
)
@cute.jit
+8 -29
View File
@@ -40,27 +40,6 @@ class CutlassDSLSetupError(Exception):
pass
def get_package_spec(requirements_path: Optional[Path] = None) -> str:
"""
Return the pip requirement spec for nvidia-cutlass-dsl from requirements.txt.
If anything goes wrong (file not found, parse failure, line missing),
return PACKAGE_NAME as a safe default.
"""
try:
req_path = requirements_path or Path(__file__).with_name("requirements.txt")
with open(req_path, "r", encoding="utf-8") as f:
for raw_line in f:
line = raw_line.strip()
if not line or line.startswith("#"):
continue
if line.lower().startswith(PACKAGE_NAME):
return line.split("#", 1)[0].strip()
except Exception:
pass
return PACKAGE_NAME
def download_wheel(temp_dir: Path) -> Path:
"""
Download the nvidia-cutlass-dsl wheel to a temporary directory.
@@ -74,10 +53,7 @@ def download_wheel(temp_dir: Path) -> Path:
Raises:
CutlassDSLSetupError: If download fails or wheel not found
"""
# Resolve package spec from requirements, or fall back to PACKAGE_NAME
package_spec = get_package_spec()
logger.info(f"Downloading {package_spec} wheel to {temp_dir}")
logger.info(f"Downloading {PACKAGE_NAME} wheel to {temp_dir}")
try:
subprocess.check_call(
@@ -87,7 +63,7 @@ def download_wheel(temp_dir: Path) -> Path:
"pip",
"download",
"--no-deps",
package_spec,
PACKAGE_NAME,
"--dest",
str(temp_dir),
],
@@ -103,7 +79,7 @@ def download_wheel(temp_dir: Path) -> Path:
raise CutlassDSLSetupError(error_msg)
# Find the downloaded wheel file
wheel_pattern = f"*.whl"
wheel_pattern = f"{PACKAGE_NAME.replace('-', '_')}-*.whl"
wheel_files = list(temp_dir.glob(wheel_pattern))
if not wheel_files:
raise CutlassDSLSetupError(
@@ -132,7 +108,7 @@ def extract_version_from_wheel(wheel_path: Path) -> str:
# Construct version regex from package name
# Wheel filename format: {package_name_with_underscores}-{version}-{python}-{abi}-{platform}.whl
package_pattern = PACKAGE_NAME.replace("-", "_")
version_regex = rf"{re.escape(package_pattern)}-([^-]+)"
version_regex = rf"{re.escape(package_pattern)}-([^-]+)-"
version_match = re.match(version_regex, wheel_filename)
if version_match:
@@ -156,7 +132,10 @@ def extract_version_from_wheel(wheel_path: Path) -> str:
return dev_version
else:
return "9.9.9.dev0"
raise CutlassDSLSetupError(
f"Could not parse version from wheel filename: {wheel_filename}"
)
def extract_wheel_contents(wheel_path: Path, extract_dir: Path) -> None:
"""
+1 -1
View File
@@ -1,3 +1,3 @@
# Use `pip install -r requirements.txt` with the present file to install a
# wheel consistent with the present state of the github repository
nvidia-cutlass-dsl==4.4.0.dev0
nvidia-cutlass-dsl==4.4.0.dev1