v4.3 tag release update. (#2789)
This commit is contained in:
@@ -29,16 +29,14 @@
|
||||
import argparse
|
||||
from typing import Optional, Tuple, Type, Union
|
||||
|
||||
import torch
|
||||
import cuda.bindings.driver as cuda
|
||||
|
||||
import cutlass
|
||||
import cutlass.cute as cute
|
||||
import cutlass.cute.testing as testing
|
||||
import cutlass.torch as cutlass_torch
|
||||
import cutlass.utils as utils
|
||||
import cutlass.pipeline as pipeline
|
||||
import cutlass.utils.blackwell_helpers as sm100_utils
|
||||
from cutlass.pipeline import pipeline_init_arrive, pipeline_init_wait
|
||||
from cutlass.cute.nvgpu import cpasync, tcgen05
|
||||
|
||||
"""
|
||||
@@ -152,10 +150,10 @@ def _compute_stages(
|
||||
num_c_stage = 2 if use_tma_store else 0
|
||||
|
||||
# Calculate smem layout and size for one stage of A, B, and C with 1-stage
|
||||
a_smem_layout_stage_one = sm100_utils.make_smem_layout_a(
|
||||
a_smem_layout_stage_one = utils.sm100.make_smem_layout_a(
|
||||
tiled_mma, mma_tiler_mnk, a_dtype, 1
|
||||
)
|
||||
b_smem_layout_staged_one = sm100_utils.make_smem_layout_b(
|
||||
b_smem_layout_staged_one = utils.sm100.make_smem_layout_b(
|
||||
tiled_mma, mma_tiler_mnk, b_dtype, 1
|
||||
)
|
||||
|
||||
@@ -229,14 +227,14 @@ class PersistentDenseGemmKernel:
|
||||
- Cluster shape M must be multiple of 2 if use_2cta_instrs=True
|
||||
- Cluster shape M/N must be positive and power of 2, total cluster size <= 16
|
||||
|
||||
Example:
|
||||
>>> gemm = PersistentDenseGemmKernel(
|
||||
... acc_dtype=cutlass.Float32,
|
||||
... use_2cta_instrs=True,
|
||||
... mma_tiler_mn=(128, 128),
|
||||
... cluster_shape_mn=(2, 2)
|
||||
... )
|
||||
>>> gemm(a_tensor, b_tensor, c_tensor, max_active_clusters, stream)
|
||||
**Example:**
|
||||
gemm = PersistentDenseGemmKernel(
|
||||
acc_dtype=cutlass.Float32,
|
||||
use_2cta_instrs=True,
|
||||
mma_tiler_mn=(128, 128),
|
||||
cluster_shape_mn=(2, 2)
|
||||
)
|
||||
gemm(a, b, c, max_active_clusters, stream)
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
@@ -316,7 +314,7 @@ class PersistentDenseGemmKernel:
|
||||
- Computing tensor memory allocation columns
|
||||
"""
|
||||
# Configure tiled mma
|
||||
tiled_mma = sm100_utils.make_trivial_tiled_mma(
|
||||
tiled_mma = utils.sm100.make_trivial_tiled_mma(
|
||||
self.a_dtype,
|
||||
self.a_major_mode,
|
||||
self.b_major_mode,
|
||||
@@ -353,7 +351,7 @@ class PersistentDenseGemmKernel:
|
||||
|
||||
# Compute epilogue subtile
|
||||
if cutlass.const_expr(self.use_tma_store):
|
||||
self.epi_tile = sm100_utils.compute_epilogue_tile_shape(
|
||||
self.epi_tile = utils.sm100.compute_epilogue_tile_shape(
|
||||
self.cta_tile_shape_mnk,
|
||||
self.use_2cta_instrs,
|
||||
self.c_layout,
|
||||
@@ -364,7 +362,7 @@ class PersistentDenseGemmKernel:
|
||||
|
||||
c_smem_layout = None
|
||||
if cutlass.const_expr(self.use_tma_store):
|
||||
c_smem_layout = sm100_utils.make_smem_layout_epi(
|
||||
c_smem_layout = utils.sm100.make_smem_layout_epi(
|
||||
self.c_dtype, self.c_layout, self.epi_tile, 1
|
||||
)
|
||||
|
||||
@@ -382,16 +380,16 @@ class PersistentDenseGemmKernel:
|
||||
)
|
||||
|
||||
# Compute A/B/C shared memory layout
|
||||
self.a_smem_layout_staged = sm100_utils.make_smem_layout_a(
|
||||
self.a_smem_layout_staged = utils.sm100.make_smem_layout_a(
|
||||
tiled_mma, self.mma_tiler, self.a_dtype, self.num_ab_stage
|
||||
)
|
||||
self.b_smem_layout_staged = sm100_utils.make_smem_layout_b(
|
||||
self.b_smem_layout_staged = utils.sm100.make_smem_layout_b(
|
||||
tiled_mma, self.mma_tiler, self.b_dtype, self.num_ab_stage
|
||||
)
|
||||
|
||||
self.c_smem_layout_staged = None
|
||||
if self.use_tma_store:
|
||||
self.c_smem_layout_staged = sm100_utils.make_smem_layout_epi(
|
||||
self.c_smem_layout_staged = utils.sm100.make_smem_layout_epi(
|
||||
self.c_dtype, self.c_layout, self.epi_tile, self.num_c_stage
|
||||
)
|
||||
|
||||
@@ -447,7 +445,7 @@ class PersistentDenseGemmKernel:
|
||||
# Setup attributes that dependent on gemm inputs
|
||||
self._setup_attributes()
|
||||
|
||||
tiled_mma = sm100_utils.make_trivial_tiled_mma(
|
||||
tiled_mma = utils.sm100.make_trivial_tiled_mma(
|
||||
self.a_dtype,
|
||||
self.a_major_mode,
|
||||
self.b_major_mode,
|
||||
@@ -458,7 +456,7 @@ class PersistentDenseGemmKernel:
|
||||
atom_thr_size = cute.size(tiled_mma.thr_id.shape)
|
||||
|
||||
# Setup TMA load for A
|
||||
a_op = sm100_utils.cluster_shape_to_tma_atom_A(
|
||||
a_op = utils.sm100.cluster_shape_to_tma_atom_A(
|
||||
self.cluster_shape_mn, tiled_mma.thr_id
|
||||
)
|
||||
a_smem_layout = cute.slice_(self.a_smem_layout_staged, (None, None, None, 0))
|
||||
@@ -475,7 +473,7 @@ class PersistentDenseGemmKernel:
|
||||
)
|
||||
|
||||
# Setup TMA load for B
|
||||
b_op = sm100_utils.cluster_shape_to_tma_atom_B(
|
||||
b_op = utils.sm100.cluster_shape_to_tma_atom_B(
|
||||
self.cluster_shape_mn, tiled_mma.thr_id
|
||||
)
|
||||
b_smem_layout = cute.slice_(self.b_smem_layout_staged, (None, None, None, 0))
|
||||
@@ -614,6 +612,7 @@ class PersistentDenseGemmKernel:
|
||||
consumer_group=ab_pipeline_consumer_group,
|
||||
tx_count=self.num_tma_load_bytes,
|
||||
cta_layout_vmnk=cluster_layout_vmnk,
|
||||
defer_sync=True,
|
||||
).make_participants()
|
||||
|
||||
# Initialize acc_pipeline (barrier) and states
|
||||
@@ -630,6 +629,7 @@ class PersistentDenseGemmKernel:
|
||||
producer_group=acc_pipeline_producer_group,
|
||||
consumer_group=acc_pipeline_consumer_group,
|
||||
cta_layout_vmnk=cluster_layout_vmnk,
|
||||
defer_sync=True,
|
||||
)
|
||||
|
||||
tmem_alloc_barrier = pipeline.NamedBarrier(
|
||||
@@ -652,8 +652,7 @@ class PersistentDenseGemmKernel:
|
||||
)
|
||||
|
||||
# Cluster arrive after barrier init
|
||||
if cute.size(self.cluster_shape_mn) > 1:
|
||||
cute.arch.cluster_arrive_relaxed()
|
||||
pipeline_init_arrive(cluster_shape_mn=cluster_layout_vmnk, is_relaxed=True)
|
||||
|
||||
#
|
||||
# Setup smem tensor A/B/C
|
||||
@@ -761,10 +760,7 @@ class PersistentDenseGemmKernel:
|
||||
#
|
||||
# Cluster wait before tensor memory alloc
|
||||
#
|
||||
if cute.size(self.cluster_shape_mn) > 1:
|
||||
cute.arch.cluster_wait()
|
||||
else:
|
||||
cute.arch.sync_threads()
|
||||
pipeline_init_wait(cluster_shape_mn=cluster_layout_vmnk)
|
||||
|
||||
#
|
||||
# Specialized TMA load warp
|
||||
@@ -1297,7 +1293,7 @@ class PersistentDenseGemmKernel:
|
||||
:rtype: Tuple[cute.TiledCopy, cute.Tensor, cute.Tensor]
|
||||
"""
|
||||
# Make tiledCopy for tensor memory load
|
||||
copy_atom_t2r = sm100_utils.get_tmem_load_op(
|
||||
copy_atom_t2r = utils.sm100.get_tmem_load_op(
|
||||
self.cta_tile_shape_mnk,
|
||||
self.c_layout,
|
||||
self.c_dtype,
|
||||
@@ -1354,7 +1350,7 @@ class PersistentDenseGemmKernel:
|
||||
- tRS_sC: The partitioned tensor C (smem destination)
|
||||
:rtype: Tuple[cute.TiledCopy, cute.Tensor, cute.Tensor]
|
||||
"""
|
||||
copy_atom_r2s = sm100_utils.get_smem_store_op(
|
||||
copy_atom_r2s = utils.sm100.get_smem_store_op(
|
||||
self.c_layout, self.c_dtype, self.acc_dtype, tiled_copy_t2r
|
||||
)
|
||||
tiled_copy_r2s = cute.make_tiled_copy_D(copy_atom_r2s, tiled_copy_t2r)
|
||||
@@ -1617,97 +1613,187 @@ class PersistentDenseGemmKernel:
|
||||
is_valid = False
|
||||
return is_valid
|
||||
|
||||
def can_implement(self, a: cute.Tensor, b: cute.Tensor, c: cute.Tensor) -> bool:
|
||||
"""Check if the given tensors can be implemented by this kernel.
|
||||
def can_implement(
|
||||
self,
|
||||
mnkl: Tuple[int, int, int, int],
|
||||
ab_dtype: Type[cutlass.Numeric],
|
||||
c_dtype: Type[cutlass.Numeric],
|
||||
a_major: str,
|
||||
b_major: str,
|
||||
c_major: str,
|
||||
) -> bool:
|
||||
"""
|
||||
Determine if the given tensor configuration can be implemented by this kernel.
|
||||
|
||||
:param a: Input tensor A
|
||||
:type a: cute.Tensor
|
||||
:param b: Input tensor B
|
||||
:type b: cute.Tensor
|
||||
:param c: Output tensor C
|
||||
:type c: cute.Tensor
|
||||
|
||||
:return: True if the gemm supports the given config, False otherwise
|
||||
: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 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").
|
||||
:type a_major: str
|
||||
:param b_major: Major dimension of the B tensor layout ("n" or "k").
|
||||
:type b_major: str
|
||||
:param c_major: Major dimension of the C tensor layout ("m" or "n").
|
||||
:type c_major: str
|
||||
:return: True if the kernel supports the given configuration, False otherwise.
|
||||
:rtype: bool
|
||||
"""
|
||||
m, n, k, l = a.shape[0], b.shape[0], a.shape[1], a.shape[2]
|
||||
|
||||
# infer a_major, b_major, c_major
|
||||
is_m_major_a = utils.LayoutEnum.from_tensor(a).is_m_major_a()
|
||||
is_n_major_b = utils.LayoutEnum.from_tensor(b).is_n_major_b()
|
||||
is_m_major_c = utils.LayoutEnum.from_tensor(c).is_m_major_c()
|
||||
a_major = "m" if is_m_major_a else "k"
|
||||
b_major = "n" if is_n_major_b else "k"
|
||||
c_major = "m" if is_m_major_c else "n"
|
||||
|
||||
can_implement = True
|
||||
# Skip unsupported types
|
||||
if not self.is_valid_dtypes(a.element_type, c.element_type):
|
||||
can_implement = False
|
||||
if not self.is_valid_dtypes(ab_dtype, c_dtype):
|
||||
return False
|
||||
|
||||
# Skip invalid mma tile shape and cluster shape
|
||||
if not self.is_valid_mma_tiler_and_cluster_shape():
|
||||
can_implement = False
|
||||
return False
|
||||
|
||||
# Unpack mnkl for clarity in calling the epilog check
|
||||
m, n, k, l = mnkl
|
||||
# Skip illegal problem shape for load/store alignment
|
||||
if not self.is_valid_tensor_alignment(
|
||||
m, n, k, l, a.element_type, c.element_type, a_major, b_major, c_major
|
||||
m, n, k, l, ab_dtype, c_dtype, a_major, b_major, c_major
|
||||
):
|
||||
can_implement = False
|
||||
return False
|
||||
# Skip invalid epilogue store option
|
||||
if not self.is_valid_epilog_store_option(m, n):
|
||||
can_implement = False
|
||||
return False
|
||||
|
||||
return can_implement
|
||||
return True
|
||||
|
||||
|
||||
def create_tensors(l, m, n, k, a_major, b_major, c_major, ab_dtype, c_dtype):
|
||||
torch.manual_seed(1111)
|
||||
@cute.jit
|
||||
def bmm(
|
||||
gemm_op: cutlass.Constexpr,
|
||||
a: cute.Tensor, # (l, m, k)
|
||||
b: cute.Tensor, # (l, k, n)
|
||||
c: cute.Tensor, # (l, m, n)
|
||||
max_active_clusters: cutlass.Constexpr,
|
||||
stream: cuda.CUstream,
|
||||
epilogue_op: cutlass.Constexpr = lambda x: x,
|
||||
):
|
||||
"""
|
||||
Wrapper API for persistent GEMM kernel to follow the convention of PyTorch's batch matrix-multiply (bmm).
|
||||
|
||||
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)
|
||||
c_torch_cpu = cutlass_torch.matrix(l, m, n, c_major == "m", c_dtype)
|
||||
Internally, the tensors are permuted to match CuTe's convention:
|
||||
- a: (m, k, l)
|
||||
- b: (n, k, l)
|
||||
- c: (m, n, l)
|
||||
|
||||
a_tensor, _ = cutlass_torch.cute_tensor_like(
|
||||
a_torch_cpu, ab_dtype, is_dynamic_layout=True, assumed_align=16
|
||||
:param gemm_op: Kernel operation, expects (a, b, c, max_active_clusters, stream, epilogue_op)
|
||||
:type gemm_op: cutlass.Constexpr
|
||||
:param a: Input tensor of shape (l, m, k)
|
||||
:type a: cute.Tensor
|
||||
:param b: Input tensor of shape (l, k, n)
|
||||
:type b: cute.Tensor
|
||||
:param c: Output tensor of shape (l, m, n)
|
||||
:type c: cute.Tensor
|
||||
:param max_active_clusters: Maximum number of hardware clusters to launch
|
||||
:type max_active_clusters: cutlass.Constexpr
|
||||
:param epilogue_op: Optional elementwise lambda function to apply per output element, defaults to identity
|
||||
:type epilogue_op: cutlass.Constexpr, optional
|
||||
"""
|
||||
# (l,m,k) -> (m,k,l)
|
||||
a = cute.make_tensor(a.iterator, cute.select(a.layout, mode=[1, 2, 0]))
|
||||
# (l,k,n) -> (n,k,l)
|
||||
b = cute.make_tensor(b.iterator, cute.select(b.layout, mode=[2, 1, 0]))
|
||||
# (l,m,n) -> (m,n,l)
|
||||
c = cute.make_tensor(c.iterator, cute.select(c.layout, mode=[1, 2, 0]))
|
||||
|
||||
gemm_op(a, b, c, max_active_clusters, stream, epilogue_op)
|
||||
|
||||
|
||||
def compile_bmm(
|
||||
gemm_op: PersistentDenseGemmKernel,
|
||||
a_dtype: Type[cutlass.Numeric],
|
||||
b_dtype: Type[cutlass.Numeric],
|
||||
c_dtype: Type[cutlass.Numeric],
|
||||
a_major: str,
|
||||
b_major: str,
|
||||
c_major: str,
|
||||
max_active_clusters: cutlass.Constexpr,
|
||||
stream: cuda.CUstream,
|
||||
epilogue_op: cutlass.Constexpr = lambda x: x,
|
||||
options: str = "",
|
||||
):
|
||||
from cutlass.cute.runtime import make_fake_compact_tensor
|
||||
|
||||
a_shape = (cute.sym_int(), cute.sym_int(divisibility=16), cute.sym_int())
|
||||
b_shape = (cute.sym_int(), cute.sym_int(divisibility=16), cute.sym_int())
|
||||
c_shape = (cute.sym_int(), cute.sym_int(divisibility=16), cute.sym_int())
|
||||
|
||||
if a_major == "k":
|
||||
a_order = (2, 1, 0) # k is leading dimension
|
||||
elif a_major == "m":
|
||||
a_order = (2, 0, 1) # m is leading dimension
|
||||
|
||||
if b_major == "n":
|
||||
b_order = (2, 1, 0) # n is leading dimension
|
||||
elif b_major == "k":
|
||||
b_order = (2, 0, 1) # k is leading dimension
|
||||
|
||||
if c_major == "n":
|
||||
c_order = (2, 1, 0) # n is leading dimension
|
||||
elif c_major == "m":
|
||||
c_order = (2, 0, 1) # m is leading dimension
|
||||
|
||||
a = make_fake_compact_tensor(
|
||||
a_dtype, a_shape, stride_order=a_order, assumed_align=16
|
||||
)
|
||||
b_tensor, _ = cutlass_torch.cute_tensor_like(
|
||||
b_torch_cpu, ab_dtype, is_dynamic_layout=True, assumed_align=16
|
||||
b = make_fake_compact_tensor(
|
||||
b_dtype, b_shape, stride_order=b_order, assumed_align=16
|
||||
)
|
||||
c_tensor, c_torch_gpu = cutlass_torch.cute_tensor_like(
|
||||
c_torch_cpu, c_dtype, is_dynamic_layout=True, assumed_align=16
|
||||
c = make_fake_compact_tensor(
|
||||
c_dtype, c_shape, stride_order=c_order, assumed_align=16
|
||||
)
|
||||
|
||||
return cute.compile(
|
||||
bmm, gemm_op, a, b, c, max_active_clusters, stream, epilogue_op, options=options
|
||||
)
|
||||
|
||||
|
||||
def prepare_tensors(
|
||||
mnkl: Tuple[int, int, int, int],
|
||||
ab_dtype: Type[cutlass.Numeric],
|
||||
c_dtype: Type[cutlass.Numeric],
|
||||
a_major: str,
|
||||
b_major: str,
|
||||
c_major: str,
|
||||
init_random: bool = True,
|
||||
):
|
||||
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")
|
||||
elif a_major == "m":
|
||||
a = 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")
|
||||
elif b_major == "k":
|
||||
b = 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")
|
||||
elif c_major == "m":
|
||||
c = 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)
|
||||
|
||||
return (
|
||||
a_tensor,
|
||||
b_tensor,
|
||||
c_tensor,
|
||||
a_torch_cpu,
|
||||
b_torch_cpu,
|
||||
c_torch_cpu,
|
||||
c_torch_gpu,
|
||||
a.to(dtype=torch_dtype(ab_dtype)),
|
||||
b.to(dtype=torch_dtype(ab_dtype)),
|
||||
c.to(dtype=torch_dtype(c_dtype)),
|
||||
)
|
||||
|
||||
|
||||
def compare(a_torch_cpu, b_torch_cpu, c_torch_gpu, c_dtype, tolerance):
|
||||
# Copy gpu result back
|
||||
kernel_result = c_torch_gpu.cpu()
|
||||
|
||||
# Compute reference result
|
||||
ref = torch.einsum(
|
||||
"mkl,nkl->mnl",
|
||||
a_torch_cpu.to(dtype=torch.float32),
|
||||
b_torch_cpu.to(dtype=torch.float32),
|
||||
)
|
||||
|
||||
# Convert ref to c_dtype
|
||||
_, ref_torch_gpu = cutlass_torch.cute_tensor_like(
|
||||
ref, c_dtype, is_dynamic_layout=True, assumed_align=16
|
||||
)
|
||||
ref_result = ref_torch_gpu.cpu()
|
||||
|
||||
# Assert close results
|
||||
torch.testing.assert_close(kernel_result, ref_result, atol=tolerance, rtol=1e-05)
|
||||
|
||||
|
||||
def run(
|
||||
mnkl: Tuple[int, int, int, int],
|
||||
ab_dtype: Type[cutlass.Numeric],
|
||||
@@ -1725,48 +1811,55 @@ def run(
|
||||
iterations: int = 1,
|
||||
skip_ref_check: bool = False,
|
||||
use_cold_l2: bool = False,
|
||||
use_tvm_ffi: bool = False,
|
||||
benchmark: bool = False,
|
||||
**kwargs,
|
||||
):
|
||||
"""Execute a persistent batched dense GEMM operation on Blackwell architecture with performance benchmarking.
|
||||
"""
|
||||
Execute a persistent batched dense GEMM operation on Blackwell architecture with performance benchmarking.
|
||||
|
||||
This function prepares input tensors, configures and launches the persistent GEMM kernel,
|
||||
optionally performs reference validation, and benchmarks the execution performance.
|
||||
Prepares input tensors, configures and launches the persistent GEMM kernel,
|
||||
optionally performs reference validation, and benchmarks execution.
|
||||
|
||||
:param mnkl: Problem size (M, N, K, L)
|
||||
: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
|
||||
:param ab_dtype: Data type for input tensors A and B.
|
||||
:type ab_dtype: Type[cutlass.Numeric]
|
||||
:param c_dtype: Data type for output tensor C
|
||||
:param c_dtype: Data type for output tensor C.
|
||||
:type c_dtype: Type[cutlass.Numeric]
|
||||
:param acc_dtype: Data type for accumulation during matrix multiplication
|
||||
:param acc_dtype: Accumulator data type for the matrix multiplication.
|
||||
:type acc_dtype: Type[cutlass.Numeric]
|
||||
:param a_major/b_major/c_major: Memory layout of tensor A/B/C
|
||||
:type a_major/b_major/c_major: str
|
||||
:param mma_tiler_mn: MMA tiling size. If not specified in the decorator parameters, the autotuner will use the
|
||||
default value of (256, 256). Otherwise, the autotuner will use the value specified in the decorator parameters.
|
||||
:param a_major: Memory layout of tensor A.
|
||||
:type a_major: str
|
||||
:param b_major: Memory layout of tensor B.
|
||||
:type b_major: str
|
||||
:param c_major: Memory layout of tensor C.
|
||||
:type c_major: str
|
||||
:param mma_tiler_mn: MMA tiling size (M, N), defaults to (256, 256).
|
||||
:type mma_tiler_mn: Tuple[int, int], optional
|
||||
:param cluster_shape_mn: Cluster shape. If not specified in the decorator parameters, the autotuner will use the
|
||||
default value of (2, 1). Otherwise, the autotuner will use the value specified in the decorator parameters.
|
||||
:param cluster_shape_mn: Cluster shape (M, N), defaults to (2, 1).
|
||||
:type cluster_shape_mn: Tuple[int, int], optional
|
||||
:param use_2cta_instrs: Whether to use 2CTA instructions. If not specified in the decorator parameters, the autotuner
|
||||
will use the default value of True. Otherwise, the autotuner will use the value specified in the decorator parameters.
|
||||
:param use_2cta_instrs: Whether to use 2CTA MMA instructions, defaults to True.
|
||||
:type use_2cta_instrs: bool, optional
|
||||
:param use_tma_store: Whether to use TMA store. If not specified in the decorator parameters, the autotuner will use
|
||||
the default value of True. Otherwise, the autotuner will use the value specified in the decorator parameters.
|
||||
:param use_tma_store: Whether to use TMA store, defaults to True.
|
||||
:type use_tma_store: bool, optional
|
||||
:param tolerance: Tolerance value for reference validation comparison, defaults to 1e-01
|
||||
:param tolerance: Tolerance for reference validation, defaults to 1e-01.
|
||||
:type tolerance: float, optional
|
||||
:param warmup_iterations: Number of warmup iterations before benchmarking, defaults to 0
|
||||
:param warmup_iterations: Number of warmup iterations before benchmarking, defaults to 0.
|
||||
:type warmup_iterations: int, optional
|
||||
:param iterations: Number of benchmark iterations to run, defaults to 1
|
||||
:param iterations: Number of benchmark iterations to run, defaults to 1.
|
||||
:type iterations: int, optional
|
||||
:param skip_ref_check: Whether to skip reference result validation, defaults to False
|
||||
:param skip_ref_check: Whether to skip reference result validation, defaults to False.
|
||||
:type skip_ref_check: bool, optional
|
||||
:param use_cold_l2: Whether to use circular buffer strategy to ensure cold L2 cache, defaults to False
|
||||
:param use_cold_l2: Whether to use circular buffer strategy to ensure cold L2 cache, defaults to False.
|
||||
:type use_cold_l2: bool, optional
|
||||
:raises RuntimeError: If CUDA GPU is not available
|
||||
:raises ValueError: If the configuration is invalid or unsupported by the kernel
|
||||
:return: Execution time of the GEMM kernel
|
||||
:param use_tvm_ffi: Whether to use TVM FFI for the kernel, defaults to False.
|
||||
:type use_tvm_ffi: bool, optional
|
||||
:param benchmark: Whether to only benchmark the kernel, defaults to False.
|
||||
:type benchmark: bool, optional
|
||||
:raises RuntimeError: If CUDA GPU is not available.
|
||||
:raises ValueError: If the configuration is invalid or unsupported by the kernel.
|
||||
:return: Execution time of the GEMM kernel.
|
||||
:rtype: float
|
||||
"""
|
||||
print("Running Blackwell Persistent Dense GEMM test with:")
|
||||
@@ -1781,9 +1874,24 @@ def run(
|
||||
print(f"Iterations: {iterations}")
|
||||
print(f"Skip reference checking: {skip_ref_check}")
|
||||
print(f"Use cold L2: {'True' if use_cold_l2 else 'False'}")
|
||||
print(f"Use TVM FFI: {'True' if use_tvm_ffi else 'False'}")
|
||||
|
||||
# Unpack parameters
|
||||
m, n, k, l = mnkl
|
||||
import torch
|
||||
from cutlass.torch import dtype as torch_dtype
|
||||
|
||||
# Build GEMM object
|
||||
gemm = PersistentDenseGemmKernel(
|
||||
acc_dtype, use_2cta_instrs, mma_tiler_mn, cluster_shape_mn, use_tma_store
|
||||
)
|
||||
can_implement = gemm.can_implement(
|
||||
mnkl, ab_dtype, c_dtype, a_major, b_major, c_major
|
||||
)
|
||||
if not can_implement:
|
||||
raise testing.CantImplementError(
|
||||
f"The current config which is invalid/unsupported: use_2cta_instrs = {use_2cta_instrs}, "
|
||||
f"mma_tiler_mn = {mma_tiler_mn}, cluster_shape_mn = {cluster_shape_mn}, "
|
||||
f"use_tma_store = {use_tma_store}"
|
||||
)
|
||||
|
||||
if not torch.cuda.is_available():
|
||||
raise RuntimeError("GPU is required to run this example!")
|
||||
@@ -1793,59 +1901,75 @@ def run(
|
||||
# Get the raw stream pointer as a CUstream
|
||||
current_stream = cuda.CUstream(torch_stream.cuda_stream)
|
||||
|
||||
a_tensor, b_tensor, c_tensor, a_torch_cpu, b_torch_cpu, c_torch_cpu, c_torch_gpu = (
|
||||
create_tensors(l, m, n, k, a_major, b_major, c_major, ab_dtype, c_dtype)
|
||||
)
|
||||
|
||||
# Build GEMM object
|
||||
gemm = PersistentDenseGemmKernel(
|
||||
acc_dtype, use_2cta_instrs, mma_tiler_mn, cluster_shape_mn, use_tma_store
|
||||
)
|
||||
|
||||
# Check if configuration can be implemented
|
||||
can_implement = gemm.can_implement(a_tensor, b_tensor, c_tensor)
|
||||
if not can_implement:
|
||||
raise ValueError(
|
||||
f"The current config which is invalid/unsupported: use_2cta_instrs = {use_2cta_instrs}, "
|
||||
f"mma_tiler_mn = {mma_tiler_mn}, cluster_shape_mn = {cluster_shape_mn}, "
|
||||
f"use_tma_store = {use_tma_store}"
|
||||
)
|
||||
max_active_clusters = utils.HardwareInfo().get_max_active_clusters(
|
||||
cluster_shape_mn[0] * cluster_shape_mn[1]
|
||||
)
|
||||
compiled_gemm = cute.compile(
|
||||
gemm, a_tensor, b_tensor, c_tensor, max_active_clusters, current_stream
|
||||
|
||||
options = []
|
||||
if use_tvm_ffi:
|
||||
options.append("--enable-tvm-ffi")
|
||||
|
||||
compiled_fn = compile_bmm(
|
||||
gemm,
|
||||
ab_dtype,
|
||||
ab_dtype,
|
||||
c_dtype,
|
||||
a_major,
|
||||
b_major,
|
||||
c_major,
|
||||
max_active_clusters,
|
||||
current_stream,
|
||||
options=",".join(options),
|
||||
)
|
||||
|
||||
# Run and verify BMM with torch
|
||||
a, b, c = prepare_tensors(mnkl, ab_dtype, c_dtype, a_major, b_major, c_major)
|
||||
|
||||
if not skip_ref_check:
|
||||
compiled_gemm(a_tensor, b_tensor, c_tensor, current_stream)
|
||||
compare(a_torch_cpu, b_torch_cpu, c_torch_gpu, c_dtype, tolerance)
|
||||
# Use small random number for deterministic result for reference check
|
||||
compiled_fn(a, b, c, torch_stream)
|
||||
|
||||
# Manually quantize to be comparable
|
||||
ref = (
|
||||
torch.bmm(a.to(dtype=torch.float32), b.to(dtype=torch.float32))
|
||||
.to(dtype=torch_dtype(c_dtype))
|
||||
.to(dtype=torch.float32)
|
||||
)
|
||||
torch.testing.assert_close(
|
||||
c.to(dtype=torch.float32), ref, atol=tolerance, rtol=1e-03
|
||||
)
|
||||
|
||||
if not benchmark:
|
||||
return 0
|
||||
|
||||
def generate_tensors():
|
||||
a_tensor, _ = cutlass_torch.cute_tensor_like(
|
||||
a_torch_cpu, ab_dtype, is_dynamic_layout=True, assumed_align=16
|
||||
init_normal = ab_dtype not in [cutlass.Int8, cutlass.Uint8]
|
||||
a, b, c = prepare_tensors(
|
||||
mnkl,
|
||||
ab_dtype,
|
||||
c_dtype,
|
||||
a_major,
|
||||
b_major,
|
||||
c_major,
|
||||
init_random=not init_normal,
|
||||
)
|
||||
b_tensor, _ = cutlass_torch.cute_tensor_like(
|
||||
b_torch_cpu, ab_dtype, is_dynamic_layout=True, assumed_align=16
|
||||
)
|
||||
c_tensor, _ = cutlass_torch.cute_tensor_like(
|
||||
c_torch_cpu, c_dtype, is_dynamic_layout=True, assumed_align=16
|
||||
)
|
||||
return testing.JitArguments(a_tensor, b_tensor, c_tensor, current_stream)
|
||||
return testing.JitArguments(a, b, c, torch_stream)
|
||||
|
||||
workspace_count = 1
|
||||
if use_cold_l2:
|
||||
one_workspace_bytes = (
|
||||
a_torch_cpu.numel() * a_torch_cpu.element_size()
|
||||
+ b_torch_cpu.numel() * b_torch_cpu.element_size()
|
||||
+ c_torch_cpu.numel() * c_torch_cpu.element_size()
|
||||
a.numel() * a.element_size()
|
||||
+ b.numel() * b.element_size()
|
||||
+ c.numel() * c.element_size()
|
||||
)
|
||||
workspace_count = testing.get_workspace_count(
|
||||
one_workspace_bytes, warmup_iterations, iterations
|
||||
)
|
||||
|
||||
exec_time = testing.benchmark(
|
||||
compiled_gemm,
|
||||
# Return execution time in microseconds
|
||||
return testing.benchmark(
|
||||
compiled_fn,
|
||||
workspace_generator=generate_tensors,
|
||||
workspace_count=workspace_count,
|
||||
stream=current_stream,
|
||||
@@ -1853,18 +1977,17 @@ def run(
|
||||
iterations=iterations,
|
||||
)
|
||||
|
||||
return exec_time # Return execution time in microseconds
|
||||
|
||||
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."
|
||||
)
|
||||
|
||||
|
||||
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."
|
||||
)
|
||||
def prepare_parser():
|
||||
|
||||
parser = argparse.ArgumentParser(
|
||||
description="Example of Dense Persistent GEMM on Blackwell."
|
||||
@@ -1872,19 +1995,13 @@ if __name__ == "__main__":
|
||||
|
||||
parser.add_argument(
|
||||
"--mnkl",
|
||||
type=parse_comma_separated_ints,
|
||||
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,
|
||||
type=_parse_comma_separated_ints,
|
||||
default=(1, 1),
|
||||
help="Cluster shape (comma-separated)",
|
||||
)
|
||||
@@ -1905,6 +2022,9 @@ if __name__ == "__main__":
|
||||
parser.add_argument(
|
||||
"--tolerance", type=float, default=1e-01, help="Tolerance for validation"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--benchmark", action="store_true", help="Only benchmark the kernel"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--warmup_iterations", type=int, default=0, help="Warmup iterations"
|
||||
)
|
||||
@@ -1923,6 +2043,24 @@ if __name__ == "__main__":
|
||||
default=False,
|
||||
help="Use circular buffer tensor sets to ensure L2 cold cache",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--use_tvm_ffi",
|
||||
action="store_true",
|
||||
default=False,
|
||||
help="Enable TVM FFI for the kernel, defaults to False using CuTe DSL's native runtime",
|
||||
)
|
||||
|
||||
return parser
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = prepare_parser()
|
||||
parser.add_argument(
|
||||
"--mma_tiler_mn",
|
||||
type=_parse_comma_separated_ints,
|
||||
default=(128, 128),
|
||||
help="Mma tile shape (comma-separated)",
|
||||
)
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
@@ -1952,5 +2090,7 @@ if __name__ == "__main__":
|
||||
args.iterations,
|
||||
args.skip_ref_check,
|
||||
args.use_cold_l2,
|
||||
args.use_tvm_ffi,
|
||||
args.benchmark,
|
||||
)
|
||||
print("PASS")
|
||||
|
||||
Reference in New Issue
Block a user