submission 399579
shiyegao · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 6124 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-group-gemm-399579?include=source"interfacepython
Compatibility
measured onNVIDIA B200
declared hardwareNVIDIA B200
architecturessm_100
dtypesfp8_e4m3, nvfp4
Benchmark evidence
1 measurement across 1 GPU, fastest first.
Operation / workload
Hardware
Latency
Rank
Observed
Reported · How evidence levels are derived →
Source and license
sourceavailable
revision digestsha256:f48937dd965d6a6580cee140d301b8695f372ff1a558d7aaac6b86ccaf4901c2
license declaredunknown
license concludedunknown
authorsshiyegao
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fused-epilogue
self.epi_tile = sm100_utils.compute_epilogue_tile_shape(mbarrier
self.epilog_sync_barrier = pipeline.NamedBarrier(persistent-kernel
) -> Tuple[utils.PersistentTileSchedulerParams, Tuple[int, int, int]]:shared-memory
self.smem_capacity = utils.get_smem_capacity_in_bytes("sm_100")tcgen05
tcgen05.CtaGroup.TWO if self.use_2cta_instrs else tcgen05.CtaGroup.ONEwarp-specialization
ab_pipeline_producer_group = pipeline.CooperativeGroup(pipeline.Agent.Thread)Kernel source
submission.py6124 lines
import torch
import atexit
import cutlass
import cutlass.cute as cute
from cutlass.cute.nvgpu import cpasync, tcgen05
import cutlass.pipeline as pipeline
import cutlass.utils as utils
import cutlass.utils.blackwell_helpers as sm100_utils
import cutlass.utils.blockscaled_layout as blockscaled_utils
from cutlass.cute.runtime import make_ptr
from typing import Tuple, Type, Union, NamedTuple
_SF_VEC_SIZE = 16
_ASSUMED_ALIGN_BYTES = 128
_MMA_TILER_MN = (128, 64)
_CLUSTER_SHAPE_MN = (1, 2)
_OCCUPANCY = 2
_TMA_CACHE_EVICT_NORMAL = 0x1000000000000000
_TMA_CACHE_EVICT_FIRST = 0x12F0000000000000
_TMA_CACHE_EVICT_LAST = 0x14F0000000000000
_CACHE_POLICY_SET_DEFAULT = 8
_CACHE_POLICY_SET_G8 = 4
_G8_GROUP_CACHE_POLICY_SET_DEFAULT = 0
_G8_GROUP_CACHE_POLICY_SET_N4096_K7168 = _CACHE_POLICY_SET_G8
_G8_GROUP_CACHE_POLICY_SET_N7168_K2048 = _CACHE_POLICY_SET_G8
_G8_GROUP_CLUSTER_ID_DEFAULT = 0
_G8_GROUP_CLUSTER_ID_N4096_K7168 = 2
_G8_GROUP_CLUSTER_ID_N7168_K2048 = 2
_G8_GROUP_OPT_LEVEL_DEFAULT = 1
_G8_GROUP_OPT_LEVEL_N4096_K7168 = 1
_G8_GROUP_OPT_LEVEL_N7168_K2048 = 1
_G2_GROUP_CACHE_POLICY_SET_DEFAULT = 0
_G2_GROUP_CACHE_POLICY_SET_N3072_K4096 = _CACHE_POLICY_SET_DEFAULT
_G2_GROUP_CACHE_POLICY_SET_N4096_K1536 = _CACHE_POLICY_SET_G8
_G2_GROUP_CLUSTER_ID_DEFAULT = 0
_G2_GROUP_CLUSTER_ID_N3072_K4096 = 1
_G2_GROUP_CLUSTER_ID_N4096_K1536 = 1
_G2_GROUP_OPT_LEVEL_DEFAULT = 1
_G2_GROUP_OPT_LEVEL_N3072_K4096 = 2
_G2_GROUP_OPT_LEVEL_N4096_K1536 = 2
class _GemmCfg(NamedTuple):
mma_tiler_mn: Tuple[int, int]
cluster_shape_mn: Tuple[int, int]
occupancy: int
cache_policy_set: int
_CFG_BASE = _GemmCfg(
_MMA_TILER_MN, _CLUSTER_SHAPE_MN, _OCCUPANCY, _CACHE_POLICY_SET_DEFAULT
)
_CFG_BASE_SET4 = _GemmCfg(_MMA_TILER_MN, _CLUSTER_SHAPE_MN, _OCCUPANCY, _CACHE_POLICY_SET_G8)
_CFG_CLUSTER_M1N4_SET4 = _GemmCfg(_MMA_TILER_MN, (1, 4), _OCCUPANCY, _CACHE_POLICY_SET_G8)
_CFG_CLUSTER_M1N4_SET8 = _GemmCfg(_MMA_TILER_MN, (1, 4), _OCCUPANCY, _CACHE_POLICY_SET_DEFAULT)
_CFG_TILER_N128_SET4 = _GemmCfg((128, 128), _CLUSTER_SHAPE_MN, _OCCUPANCY, _CACHE_POLICY_SET_G8)
_CFG_TILER_N128_SET8 = _GemmCfg((128, 128), _CLUSTER_SHAPE_MN, _OCCUPANCY, _CACHE_POLICY_SET_DEFAULT)
_CFG_TILER_N128_SET4_OCC1 = _GemmCfg((128, 128), _CLUSTER_SHAPE_MN, 1, _CACHE_POLICY_SET_G8)
_CFG_TILER_N128_SET8_OCC1 = _GemmCfg((128, 128), _CLUSTER_SHAPE_MN, 1, _CACHE_POLICY_SET_DEFAULT)
_CFG_TILER_N192 = _GemmCfg((128, 192), _CLUSTER_SHAPE_MN, _OCCUPANCY, _CACHE_POLICY_SET_DEFAULT)
_CFG_TILER_N192_SET4 = _GemmCfg((128, 192), _CLUSTER_SHAPE_MN, _OCCUPANCY, _CACHE_POLICY_SET_G8)
_CFG_TILER_N192_OCC1 = _GemmCfg((128, 192), _CLUSTER_SHAPE_MN, 1, _CACHE_POLICY_SET_DEFAULT)
_CFG_TILER_N192_SET4_OCC1 = _GemmCfg((128, 192), _CLUSTER_SHAPE_MN, 1, _CACHE_POLICY_SET_G8)
_CFG_TILER_N256_SET4 = _GemmCfg((128, 256), _CLUSTER_SHAPE_MN, _OCCUPANCY, _CACHE_POLICY_SET_G8)
_CFG_TILER_N256_SET8 = _GemmCfg((128, 256), _CLUSTER_SHAPE_MN, _OCCUPANCY, _CACHE_POLICY_SET_DEFAULT)
_CFG_TILER_N256_SET4_OCC1 = _GemmCfg((128, 256), _CLUSTER_SHAPE_MN, 1, _CACHE_POLICY_SET_G8)
_DEFAULT_CFG = _CFG_BASE
_G8_CFG_ID_BASE_SET4 = 1
_G8_CFG_ID_BASE_SET8 = 2
_G8_CFG_ID_CLUSTER_M1N4_SET4 = 3
_G8_CFG_ID_CLUSTER_M1N4_SET8 = 4
_G8_CFG_ID_TILER_N128_SET4 = 5
_G8_CFG_ID_TILER_N256_SET4 = 6
_G8_CFG_ID_TILER_N256_SET4_OCC1 = 7
_G8_FORCE_CFG_ID = 0
_G8_CFG_ID_N4096_K7168 = 7
_G8_CFG_ID_N7168_K2048 = 7
_G2_N192_USE_SET4 = 1
_G2_RANKED_OCC1 = 0
_G2_RANKED_OCC1_MMAX_GE_N3072_K4096 = 320
_G2_RANKED_OCC1_MMAX_GE_N4096_K1536 = 448
_G8_RANKED_CACHE_SET_N7168_K2048 = 8
_G8_RANKED_CACHE_SET_N7168_K2048_ALT = 7
_G8_RANKED_CACHE_USE_ALT_N7168_K2048 = 0
_G8_RANKED_CLUSTER_MMAX_GE_N4096_K7168_TO_M1N1 = 1024
_G8_RANKED_CLUSTER_MMAX_GE_N4096_K7168_TO_M1N2 = 160
_G8_RANKED_CLUSTER_MMAX_GE_N7168_K2048_TO_M1N4 = 192
_ENABLE_G8_GROUP = 1
_ENABLE_G2_GROUP = 1
_DIAG_SYNC = 0
_DIAG_LEVEL = 0
_DIAG_CTX_MAX_CHARS = 160
def _cfg_from_g8_id(cfg_id: int) -> _GemmCfg:
if cfg_id == _G8_CFG_ID_BASE_SET4:
return _CFG_BASE_SET4
if cfg_id == _G8_CFG_ID_BASE_SET8:
return _CFG_BASE
if cfg_id == _G8_CFG_ID_CLUSTER_M1N4_SET4:
return _CFG_CLUSTER_M1N4_SET4
if cfg_id == _G8_CFG_ID_CLUSTER_M1N4_SET8:
return _CFG_CLUSTER_M1N4_SET8
if cfg_id == _G8_CFG_ID_TILER_N128_SET4:
return _CFG_TILER_N128_SET4
if cfg_id == _G8_CFG_ID_TILER_N256_SET4:
return _CFG_TILER_N256_SET4
if cfg_id == _G8_CFG_ID_TILER_N256_SET4_OCC1:
return _CFG_TILER_N256_SET4_OCC1
raise ValueError(f"invalid g8 cfg_id={cfg_id} (expected 1..7)")
def _g8_cfg_id_for_nk(n: int, k: int) -> int:
forced = int(_G8_FORCE_CFG_ID)
if forced:
return forced
if n == 4096 and k == 7168:
return int(_G8_CFG_ID_N4096_K7168)
if n == 7168 and k == 2048:
return int(_G8_CFG_ID_N7168_K2048)
return 2
def _cluster_shape_from_id(cluster_id: int) -> Tuple[int, int]:
if cluster_id == 1:
return (1, 2)
if cluster_id == 2:
return (1, 4)
return (1, 1)
def _clamp_cache_set(v: int, fallback: int) -> int:
vv = int(v)
if vv < 1 or vv > 8 or vv == 5 or vv == 6:
return int(fallback)
return vv
def _clamp_opt_level(v: int, fallback: int) -> int:
vv = int(v)
return vv if (vv == 1 or vv == 2) else int(fallback)
def _g8_group_cache_policy_set_for_nk(n: int, k: int, fallback: int) -> int:
if n == 4096 and k == 7168:
v = int(_G8_GROUP_CACHE_POLICY_SET_N4096_K7168)
chosen = int(fallback) if v == 0 else v
return _clamp_cache_set(chosen, fallback)
if n == 7168 and k == 2048:
v = int(_G8_GROUP_CACHE_POLICY_SET_N7168_K2048)
chosen = int(fallback) if v == 0 else v
return _clamp_cache_set(chosen, fallback)
default = int(_G8_GROUP_CACHE_POLICY_SET_DEFAULT)
chosen = int(fallback) if default == 0 else default
return _clamp_cache_set(chosen, fallback)
def _g8_group_cluster_shape_for_nk(n: int, k: int) -> Tuple[int, int]:
cid = int(_G8_GROUP_CLUSTER_ID_DEFAULT)
if n == 4096 and k == 7168:
cid = int(_G8_GROUP_CLUSTER_ID_N4096_K7168)
elif n == 7168 and k == 2048:
cid = int(_G8_GROUP_CLUSTER_ID_N7168_K2048)
return _cluster_shape_from_id(cid)
def _g8_group_opt_level_for_nk(n: int, k: int) -> int:
if n == 4096 and k == 7168:
return _clamp_opt_level(int(_G8_GROUP_OPT_LEVEL_N4096_K7168), 1)
if n == 7168 and k == 2048:
v = int(_G8_GROUP_OPT_LEVEL_N7168_K2048)
v = 1 if v == 3 else v
return _clamp_opt_level(v, 1)
return _clamp_opt_level(int(_G8_GROUP_OPT_LEVEL_DEFAULT), 1)
def _g2_group_cluster_shape_for_nk(n: int, k: int) -> Tuple[int, int]:
cid = int(_G2_GROUP_CLUSTER_ID_DEFAULT)
if n == 3072 and k == 4096:
cid = int(_G2_GROUP_CLUSTER_ID_N3072_K4096)
elif n == 4096 and k == 1536:
cid = int(_G2_GROUP_CLUSTER_ID_N4096_K1536)
return _cluster_shape_from_id(cid)
def _g2_group_cache_policy_set_for_nk(n: int, k: int, fallback: int) -> int:
if n == 3072 and k == 4096:
v = int(_G2_GROUP_CACHE_POLICY_SET_N3072_K4096)
chosen = int(fallback) if v == 0 else v
return _clamp_cache_set(chosen, fallback)
if n == 4096 and k == 1536:
v = int(_G2_GROUP_CACHE_POLICY_SET_N4096_K1536)
chosen = int(fallback) if v == 0 else v
return _clamp_cache_set(chosen, fallback)
default = int(_G2_GROUP_CACHE_POLICY_SET_DEFAULT)
chosen = int(fallback) if default == 0 else default
return _clamp_cache_set(chosen, fallback)
def _g2_group_opt_level_for_nk(n: int, k: int) -> int:
if n == 3072 and k == 4096:
return _clamp_opt_level(int(_G2_GROUP_OPT_LEVEL_N3072_K4096), 1)
if n == 4096 and k == 1536:
return _clamp_opt_level(int(_G2_GROUP_OPT_LEVEL_N4096_K1536), 1)
return _clamp_opt_level(int(_G2_GROUP_OPT_LEVEL_DEFAULT), 1)
class Sm100BlockScaledPersistentDenseGemmKernel:
def __init__(
self,
sf_vec_size: int,
mma_tiler_mn: Tuple[int, int],
cluster_shape_mn: Tuple[int, int],
cache_policy_set: int,
occupancy: int = 1,
):
self.acc_dtype = cutlass.Float32
self.sf_vec_size = sf_vec_size
self.use_2cta_instrs = mma_tiler_mn[0] == 256
self.cluster_shape_mn = cluster_shape_mn
self.cache_policy_set = int(cache_policy_set)
self.mma_tiler = (*mma_tiler_mn, 1)
self.cta_group = (
tcgen05.CtaGroup.TWO if self.use_2cta_instrs else tcgen05.CtaGroup.ONE
)
self.occupancy = int(occupancy)
self.epilog_warp_id = (
0,
1,
2,
3,
)
self.mma_warp_id = 4
self.tma_warp_id = 5
self.threads_per_cta = 32 * len(
(self.mma_warp_id, self.tma_warp_id, *self.epilog_warp_id)
)
self.epilog_sync_barrier = pipeline.NamedBarrier(
barrier_id=1,
num_threads=32 * len(self.epilog_warp_id),
)
self.tmem_alloc_barrier = pipeline.NamedBarrier(
barrier_id=2,
num_threads=32 * len((self.mma_warp_id, *self.epilog_warp_id)),
)
self.smem_capacity = utils.get_smem_capacity_in_bytes("sm_100")
self.num_tmem_alloc_cols = 512
def _setup_attributes(self):
self.mma_inst_shape_mn = (
self.mma_tiler[0],
self.mma_tiler[1],
)
self.mma_inst_shape_mn_sfb = (
self.mma_inst_shape_mn[0] // (2 if self.use_2cta_instrs else 1),
cute.round_up(self.mma_inst_shape_mn[1], 128),
)
tiled_mma = sm100_utils.make_blockscaled_trivial_tiled_mma(
self.a_dtype,
self.a_major_mode,
self.b_major_mode,
self.sf_dtype,
self.sf_vec_size,
self.cta_group,
self.mma_inst_shape_mn,
)
tiled_mma_sfb = sm100_utils.make_blockscaled_trivial_tiled_mma(
self.a_dtype,
self.a_major_mode,
self.b_major_mode,
self.sf_dtype,
self.sf_vec_size,
cute.nvgpu.tcgen05.CtaGroup.ONE,
self.mma_inst_shape_mn_sfb,
)
mma_inst_shape_k = cute.size(tiled_mma.shape_mnk, mode=[2])
mma_inst_tile_k = 4
self.mma_tiler = (
self.mma_inst_shape_mn[0],
self.mma_inst_shape_mn[1],
mma_inst_shape_k * mma_inst_tile_k,
)
self.mma_tiler_sfb = (
self.mma_inst_shape_mn_sfb[0],
self.mma_inst_shape_mn_sfb[1],
mma_inst_shape_k * mma_inst_tile_k,
)
self.cta_tile_shape_mnk = (
self.mma_tiler[0] // cute.size(tiled_mma.thr_id.shape),
self.mma_tiler[1],
self.mma_tiler[2],
)
self.cta_tile_shape_mnk_sfb = (
self.mma_tiler_sfb[0] // cute.size(tiled_mma.thr_id.shape),
self.mma_tiler_sfb[1],
self.mma_tiler_sfb[2],
)
self.cluster_layout_vmnk = cute.tiled_divide(
cute.make_layout((*self.cluster_shape_mn, 1)),
(tiled_mma.thr_id.shape,),
)
self.cluster_layout_sfb_vmnk = cute.tiled_divide(
cute.make_layout((*self.cluster_shape_mn, 1)),
(tiled_mma_sfb.thr_id.shape,),
)
self.num_mcast_ctas_a = cute.size(self.cluster_layout_vmnk.shape[2])
self.num_mcast_ctas_b = cute.size(self.cluster_layout_vmnk.shape[1])
self.num_mcast_ctas_sfb = cute.size(self.cluster_layout_sfb_vmnk.shape[1])
self.is_a_mcast = self.num_mcast_ctas_a > 1
self.is_b_mcast = self.num_mcast_ctas_b > 1
self.is_sfb_mcast = self.num_mcast_ctas_sfb > 1
self.tma_cache_policy_a = (
_TMA_CACHE_EVICT_LAST if self.is_a_mcast else _TMA_CACHE_EVICT_FIRST
)
self.tma_cache_policy_b = (
_TMA_CACHE_EVICT_LAST if self.is_b_mcast else _TMA_CACHE_EVICT_FIRST
)
self.tma_cache_policy_sfa = (
_TMA_CACHE_EVICT_NORMAL if self.is_a_mcast else _TMA_CACHE_EVICT_FIRST
)
self.tma_cache_policy_sfb = (
_TMA_CACHE_EVICT_NORMAL if self.is_sfb_mcast else _TMA_CACHE_EVICT_FIRST
)
self.tma_cache_policy_c = _TMA_CACHE_EVICT_FIRST
if self.cache_policy_set == 1:
self.tma_cache_policy_a = _TMA_CACHE_EVICT_FIRST
self.tma_cache_policy_b = _TMA_CACHE_EVICT_FIRST
self.tma_cache_policy_sfa = _TMA_CACHE_EVICT_FIRST
self.tma_cache_policy_sfb = _TMA_CACHE_EVICT_FIRST
self.tma_cache_policy_c = _TMA_CACHE_EVICT_FIRST
elif self.cache_policy_set == 2:
self.tma_cache_policy_a = _TMA_CACHE_EVICT_FIRST
self.tma_cache_policy_b = _TMA_CACHE_EVICT_LAST
self.tma_cache_policy_sfa = _TMA_CACHE_EVICT_NORMAL
self.tma_cache_policy_sfb = _TMA_CACHE_EVICT_LAST
self.tma_cache_policy_c = _TMA_CACHE_EVICT_FIRST
elif self.cache_policy_set == 3:
self.tma_cache_policy_a = _TMA_CACHE_EVICT_LAST
self.tma_cache_policy_b = _TMA_CACHE_EVICT_FIRST
self.tma_cache_policy_sfa = _TMA_CACHE_EVICT_FIRST
self.tma_cache_policy_sfb = _TMA_CACHE_EVICT_FIRST
self.tma_cache_policy_c = _TMA_CACHE_EVICT_FIRST
elif self.cache_policy_set == 4:
self.tma_cache_policy_a = _TMA_CACHE_EVICT_FIRST
self.tma_cache_policy_b = _TMA_CACHE_EVICT_FIRST
self.tma_cache_policy_sfa = _TMA_CACHE_EVICT_NORMAL
self.tma_cache_policy_sfb = _TMA_CACHE_EVICT_NORMAL
self.tma_cache_policy_c = _TMA_CACHE_EVICT_FIRST
elif self.cache_policy_set == 5:
self.tma_cache_policy_a = _TMA_CACHE_EVICT_LAST
self.tma_cache_policy_b = _TMA_CACHE_EVICT_FIRST
self.tma_cache_policy_sfa = _TMA_CACHE_EVICT_NORMAL
self.tma_cache_policy_sfb = _TMA_CACHE_EVICT_FIRST
self.tma_cache_policy_c = _TMA_CACHE_EVICT_FIRST
elif self.cache_policy_set == 6:
self.tma_cache_policy_a = _TMA_CACHE_EVICT_LAST
self.tma_cache_policy_b = _TMA_CACHE_EVICT_FIRST
self.tma_cache_policy_sfa = _TMA_CACHE_EVICT_FIRST
self.tma_cache_policy_sfb = _TMA_CACHE_EVICT_NORMAL
self.tma_cache_policy_c = _TMA_CACHE_EVICT_FIRST
elif self.cache_policy_set == 7:
self.tma_cache_policy_a = _TMA_CACHE_EVICT_LAST
self.tma_cache_policy_b = _TMA_CACHE_EVICT_FIRST
self.tma_cache_policy_sfa = _TMA_CACHE_EVICT_NORMAL
self.tma_cache_policy_sfb = _TMA_CACHE_EVICT_NORMAL
self.tma_cache_policy_c = _TMA_CACHE_EVICT_FIRST
elif self.cache_policy_set == 8:
self.tma_cache_policy_a = _TMA_CACHE_EVICT_FIRST
self.tma_cache_policy_b = _TMA_CACHE_EVICT_LAST
self.tma_cache_policy_sfa = _TMA_CACHE_EVICT_NORMAL
self.tma_cache_policy_sfb = _TMA_CACHE_EVICT_NORMAL
self.tma_cache_policy_c = _TMA_CACHE_EVICT_FIRST
self.epi_tile = sm100_utils.compute_epilogue_tile_shape(
self.cta_tile_shape_mnk,
self.use_2cta_instrs,
self.c_layout,
self.c_dtype,
)
self.epi_tile_n = cute.size(self.epi_tile[1])
self.num_acc_stage, self.num_ab_stage, self.num_c_stage = self._compute_stages(
tiled_mma,
self.mma_tiler,
self.a_dtype,
self.b_dtype,
self.epi_tile,
self.c_dtype,
self.c_layout,
self.sf_dtype,
self.sf_vec_size,
self.smem_capacity,
self.occupancy,
)
self.a_smem_layout_staged = sm100_utils.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(
tiled_mma,
self.mma_tiler,
self.b_dtype,
self.num_ab_stage,
)
self.sfa_smem_layout_staged = blockscaled_utils.make_smem_layout_sfa(
tiled_mma,
self.mma_tiler,
self.sf_vec_size,
self.num_ab_stage,
)
self.sfb_smem_layout_staged = blockscaled_utils.make_smem_layout_sfb(
tiled_mma,
self.mma_tiler,
self.sf_vec_size,
self.num_ab_stage,
)
self.c_smem_layout_staged = sm100_utils.make_smem_layout_epi(
self.c_dtype,
self.c_layout,
self.epi_tile,
self.num_c_stage,
)
self.overlapping_accum = self.num_acc_stage == 1
sf_atom_mn = 32
self.num_sfa_tmem_cols = (self.cta_tile_shape_mnk[0] // sf_atom_mn) * mma_inst_tile_k
self.num_sfb_tmem_cols = (self.cta_tile_shape_mnk_sfb[1] // sf_atom_mn) * mma_inst_tile_k
self.num_sf_tmem_cols = self.num_sfa_tmem_cols + self.num_sfb_tmem_cols
self.num_accumulator_tmem_cols = self.cta_tile_shape_mnk[1] * self.num_acc_stage if not self.overlapping_accum else self.cta_tile_shape_mnk[1] * 2 - self.num_sf_tmem_cols
self.iter_acc_early_release_in_epilogue = self.num_sf_tmem_cols // self.epi_tile_n
@cute.jit
def __call__(
self,
a_ptr: cute.Pointer,
b_ptr: cute.Pointer,
sfa_ptr: cute.Pointer,
sfb_ptr: cute.Pointer,
c_ptr: cute.Pointer,
problem_size: tuple,
max_active_clusters: cutlass.Constexpr,
epilogue_op: cutlass.Constexpr = lambda x: x,
):
m, n, k, l = problem_size
sf_k = k // self.sf_vec_size
a_tensor = cute.make_tensor(
a_ptr,
cute.make_layout(
(m, cute.assume(k, 32), l),
stride=(cute.assume(k, 32), 1, cute.assume(m * k, 32)),
),
)
b_tensor = cute.make_tensor(
b_ptr,
cute.make_layout(
(n, cute.assume(k, 32), l),
stride=(cute.assume(k, 32), 1, cute.assume(n * k, 32)),
),
)
c_tensor = cute.make_tensor(
c_ptr,
cute.make_layout((m, n, l), stride=(n, 1, m * n)),
)
sfa_tensor = cute.make_tensor(
sfa_ptr,
cute.make_layout((m, sf_k, l), stride=(sf_k, 1, m * sf_k)),
)
sfb_tensor = cute.make_tensor(
sfb_ptr,
cute.make_layout((n, sf_k, l), stride=(sf_k, 1, n * sf_k)),
)
a_tensor.mark_compact_shape_dynamic(mode=1, stride_order=(2, 0, 1), divisibility=2)
b_tensor.mark_compact_shape_dynamic(mode=1, stride_order=(2, 0, 1), divisibility=2)
c_tensor.mark_compact_shape_dynamic(mode=1, stride_order=(2, 0, 1), divisibility=1)
self.a_dtype: Type[cutlass.Numeric] = a_tensor.element_type
self.b_dtype: Type[cutlass.Numeric] = b_tensor.element_type
self.sf_dtype: Type[cutlass.Numeric] = sfa_tensor.element_type
self.c_dtype: Type[cutlass.Numeric] = c_tensor.element_type
self.a_major_mode = utils.LayoutEnum.from_tensor(a_tensor).mma_major_mode()
self.b_major_mode = utils.LayoutEnum.from_tensor(b_tensor).mma_major_mode()
self.c_layout = utils.LayoutEnum.from_tensor(c_tensor)
if cutlass.const_expr(self.a_dtype != self.b_dtype):
raise TypeError(f"Type must match: {self.a_dtype} != {self.b_dtype}")
self._setup_attributes()
sfa_layout = blockscaled_utils.tile_atom_to_shape_SF(
a_tensor.shape, self.sf_vec_size
)
sfa_tensor = cute.make_tensor(sfa_tensor.iterator, sfa_layout)
sfb_layout = blockscaled_utils.tile_atom_to_shape_SF(
b_tensor.shape, self.sf_vec_size
)
sfb_tensor = cute.make_tensor(sfb_tensor.iterator, sfb_layout)
tiled_mma = sm100_utils.make_blockscaled_trivial_tiled_mma(
self.a_dtype,
self.a_major_mode,
self.b_major_mode,
self.sf_dtype,
self.sf_vec_size,
self.cta_group,
self.mma_inst_shape_mn,
)
tiled_mma_sfb = sm100_utils.make_blockscaled_trivial_tiled_mma(
self.a_dtype,
self.a_major_mode,
self.b_major_mode,
self.sf_dtype,
self.sf_vec_size,
cute.nvgpu.tcgen05.CtaGroup.ONE,
self.mma_inst_shape_mn_sfb,
)
atom_thr_size = cute.size(tiled_mma.thr_id.shape)
a_op = sm100_utils.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))
tma_atom_a, tma_tensor_a = cute.nvgpu.make_tiled_tma_atom_A(
a_op,
a_tensor,
a_smem_layout,
self.mma_tiler,
tiled_mma,
self.cluster_layout_vmnk.shape,
)
b_op = sm100_utils.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))
tma_atom_b, tma_tensor_b = cute.nvgpu.make_tiled_tma_atom_B(
b_op,
b_tensor,
b_smem_layout,
self.mma_tiler,
tiled_mma,
self.cluster_layout_vmnk.shape,
)
sfa_op = sm100_utils.cluster_shape_to_tma_atom_A(
self.cluster_shape_mn, tiled_mma.thr_id
)
sfa_smem_layout = cute.slice_(
self.sfa_smem_layout_staged, (None, None, None, 0)
)
tma_atom_sfa, tma_tensor_sfa = cute.nvgpu.make_tiled_tma_atom_A(
sfa_op,
sfa_tensor,
sfa_smem_layout,
self.mma_tiler,
tiled_mma,
self.cluster_layout_vmnk.shape,
internal_type=cutlass.Int16,
)
sfb_op = sm100_utils.cluster_shape_to_tma_atom_SFB(
self.cluster_shape_mn, tiled_mma.thr_id
)
sfb_smem_layout = cute.slice_(
self.sfb_smem_layout_staged, (None, None, None, 0)
)
tma_atom_sfb, tma_tensor_sfb = cute.nvgpu.make_tiled_tma_atom_B(
sfb_op,
sfb_tensor,
sfb_smem_layout,
self.mma_tiler_sfb,
tiled_mma_sfb,
self.cluster_layout_sfb_vmnk.shape,
internal_type=cutlass.Int16,
)
if cutlass.const_expr(self.cta_tile_shape_mnk[1] == 192):
x = tma_tensor_sfb.stride[0][1]
y = cute.ceil_div(tma_tensor_sfb.shape[0][1], 4)
new_shape = (
(
tma_tensor_sfb.shape[0][0],
((2, 2), y)
),
tma_tensor_sfb.shape[1],
tma_tensor_sfb.shape[2]
)
x_times_3 = 3 * x
new_stride = (
(
tma_tensor_sfb.stride[0][0],
((x, x), x_times_3)
),
tma_tensor_sfb.stride[1],
tma_tensor_sfb.stride[2]
)
tma_tensor_sfb_new_layout = cute.make_layout(new_shape, stride=new_stride)
tma_tensor_sfb = cute.make_tensor(tma_tensor_sfb.iterator, tma_tensor_sfb_new_layout)
a_copy_size = cute.size_in_bytes(self.a_dtype, a_smem_layout)
b_copy_size = cute.size_in_bytes(self.b_dtype, b_smem_layout)
sfa_copy_size = cute.size_in_bytes(self.sf_dtype, sfa_smem_layout)
sfb_copy_size = cute.size_in_bytes(self.sf_dtype, sfb_smem_layout)
self.num_tma_load_bytes = (
a_copy_size + b_copy_size + sfa_copy_size + sfb_copy_size
) * atom_thr_size
epi_smem_layout = cute.slice_(self.c_smem_layout_staged, (None, None, 0))
tma_atom_c, tma_tensor_c = cpasync.make_tiled_tma_atom(
cpasync.CopyBulkTensorTileS2GOp(),
c_tensor,
epi_smem_layout,
self.epi_tile,
)
grid_m = cute.ceil_div(m, self.cta_tile_shape_mnk[0]) * cute.size(tiled_mma.thr_id.shape)
grid_n = cute.ceil_div(n, self.cta_tile_shape_mnk[1])
grid = (grid_m, grid_n, l)
self.buffer_align_bytes = 1024
@cute.struct
class SharedStorage:
ab_full_mbar_ptr: cute.struct.MemRange[cutlass.Int64, self.num_ab_stage]
ab_empty_mbar_ptr: cute.struct.MemRange[cutlass.Int64, self.num_ab_stage]
acc_full_mbar_ptr: cute.struct.MemRange[cutlass.Int64, self.num_acc_stage]
acc_empty_mbar_ptr: cute.struct.MemRange[cutlass.Int64, self.num_acc_stage]
tmem_dealloc_mbar_ptr: cutlass.Int64
tmem_holding_buf: cutlass.Int32
sC: cute.struct.Align[
cute.struct.MemRange[
self.c_dtype,
cute.cosize(self.c_smem_layout_staged.outer),
],
self.buffer_align_bytes,
]
sA: cute.struct.Align[
cute.struct.MemRange[
self.a_dtype, cute.cosize(self.a_smem_layout_staged.outer)
],
self.buffer_align_bytes,
]
sB: cute.struct.Align[
cute.struct.MemRange[
self.b_dtype, cute.cosize(self.b_smem_layout_staged.outer)
],
self.buffer_align_bytes,
]
sSFA: cute.struct.Align[
cute.struct.MemRange[
self.sf_dtype, cute.cosize(self.sfa_smem_layout_staged)
],
self.buffer_align_bytes,
]
sSFB: cute.struct.Align[
cute.struct.MemRange[
self.sf_dtype, cute.cosize(self.sfb_smem_layout_staged)
],
self.buffer_align_bytes,
]
self.shared_storage = SharedStorage
self.kernel(
tiled_mma,
tiled_mma_sfb,
tma_atom_a,
tma_tensor_a,
tma_atom_b,
tma_tensor_b,
tma_atom_sfa,
tma_tensor_sfa,
tma_atom_sfb,
tma_tensor_sfb,
tma_atom_c,
tma_tensor_c,
self.cluster_layout_vmnk,
self.cluster_layout_sfb_vmnk,
self.a_smem_layout_staged,
self.b_smem_layout_staged,
self.sfa_smem_layout_staged,
self.sfb_smem_layout_staged,
self.c_smem_layout_staged,
self.epi_tile,
epilogue_op,
).launch(
grid=grid,
block=[self.threads_per_cta, 1, 1],
cluster=(*self.cluster_shape_mn, 1),
min_blocks_per_mp=self.occupancy,
)
return
@cute.kernel
def kernel(
self,
tiled_mma: cute.TiledMma,
tiled_mma_sfb: cute.TiledMma,
tma_atom_a: cute.CopyAtom,
mA_mkl: cute.Tensor,
tma_atom_b: cute.CopyAtom,
mB_nkl: cute.Tensor,
tma_atom_sfa: cute.CopyAtom,
mSFA_mkl: cute.Tensor,
tma_atom_sfb: cute.CopyAtom,
mSFB_nkl: cute.Tensor,
tma_atom_c: cute.CopyAtom,
mC_mnl: cute.Tensor,
cluster_layout_vmnk: cute.Layout,
cluster_layout_sfb_vmnk: cute.Layout,
a_smem_layout_staged: cute.ComposedLayout,
b_smem_layout_staged: cute.ComposedLayout,
sfa_smem_layout_staged: cute.Layout,
sfb_smem_layout_staged: cute.Layout,
c_smem_layout_staged: Union[cute.Layout, cute.ComposedLayout],
epi_tile: cute.Tile,
epilogue_op: cutlass.Constexpr,
):
warp_idx = cute.arch.warp_idx()
warp_idx = cute.arch.make_warp_uniform(warp_idx)
if warp_idx == self.tma_warp_id:
cpasync.prefetch_descriptor(tma_atom_a)
cpasync.prefetch_descriptor(tma_atom_b)
cpasync.prefetch_descriptor(tma_atom_sfa)
cpasync.prefetch_descriptor(tma_atom_sfb)
cpasync.prefetch_descriptor(tma_atom_c)
use_2cta_instrs = cute.size(tiled_mma.thr_id.shape) == 2
bidx, bidy, bidz = cute.arch.block_idx()
mma_tile_coord_v = bidx % cute.size(tiled_mma.thr_id.shape)
is_leader_cta = mma_tile_coord_v == 0
cta_rank_in_cluster = cute.arch.make_warp_uniform(
cute.arch.block_idx_in_cluster()
)
block_in_cluster_coord_vmnk = cluster_layout_vmnk.get_flat_coord(
cta_rank_in_cluster
)
block_in_cluster_coord_sfb_vmnk = cluster_layout_sfb_vmnk.get_flat_coord(
cta_rank_in_cluster
)
tidx, _, _ = cute.arch.thread_idx()
smem = utils.SmemAllocator()
storage = smem.allocate(self.shared_storage)
ab_pipeline_producer_group = pipeline.CooperativeGroup(pipeline.Agent.Thread)
num_tma_producer = self.num_mcast_ctas_a + self.num_mcast_ctas_b - 1
ab_pipeline_consumer_group = pipeline.CooperativeGroup(
pipeline.Agent.Thread, num_tma_producer
)
try:
ab_pipeline = pipeline.PipelineTmaUmma.create(
barrier_storage=storage.ab_full_mbar_ptr.data_ptr(),
num_stages=self.num_ab_stage,
producer_group=ab_pipeline_producer_group,
consumer_group=ab_pipeline_consumer_group,
tx_count=self.num_tma_load_bytes,
cta_layout_vmnk=cluster_layout_vmnk,
defer_sync=True,
)
except TypeError:
ab_pipeline = pipeline.PipelineTmaUmma.create(
barrier_storage=storage.ab_full_mbar_ptr.data_ptr(),
num_stages=self.num_ab_stage,
producer_group=ab_pipeline_producer_group,
consumer_group=ab_pipeline_consumer_group,
tx_count=self.num_tma_load_bytes,
cta_layout_vmnk=cluster_layout_vmnk,
)
acc_pipeline_producer_group = pipeline.CooperativeGroup(pipeline.Agent.Thread)
num_acc_consumer_threads = len(self.epilog_warp_id) * (
2 if use_2cta_instrs else 1
)
acc_pipeline_consumer_group = pipeline.CooperativeGroup(
pipeline.Agent.Thread, num_acc_consumer_threads
)
try:
acc_pipeline = pipeline.PipelineUmmaAsync.create(
barrier_storage=storage.acc_full_mbar_ptr.data_ptr(),
num_stages=self.num_acc_stage,
producer_group=acc_pipeline_producer_group,
consumer_group=acc_pipeline_consumer_group,
cta_layout_vmnk=cluster_layout_vmnk,
defer_sync=True,
)
except TypeError:
acc_pipeline = pipeline.PipelineUmmaAsync.create(
barrier_storage=storage.acc_full_mbar_ptr.data_ptr(),
num_stages=self.num_acc_stage,
producer_group=acc_pipeline_producer_group,
consumer_group=acc_pipeline_consumer_group,
cta_layout_vmnk=cluster_layout_vmnk,
)
tmem = utils.TmemAllocator(
storage.tmem_holding_buf,
barrier_for_retrieve=self.tmem_alloc_barrier,
allocator_warp_id=self.epilog_warp_id[0],
is_two_cta=use_2cta_instrs,
two_cta_tmem_dealloc_mbar_ptr=storage.tmem_dealloc_mbar_ptr,
)
sC = storage.sC.get_tensor(
c_smem_layout_staged.outer, swizzle=c_smem_layout_staged.inner
)
sA = storage.sA.get_tensor(
a_smem_layout_staged.outer, swizzle=a_smem_layout_staged.inner
)
sB = storage.sB.get_tensor(
b_smem_layout_staged.outer, swizzle=b_smem_layout_staged.inner
)
sSFA = storage.sSFA.get_tensor(sfa_smem_layout_staged)
sSFB = storage.sSFB.get_tensor(sfb_smem_layout_staged)
a_full_mcast_mask = None
b_full_mcast_mask = None
sfa_full_mcast_mask = None
sfb_full_mcast_mask = None
if cutlass.const_expr(self.is_a_mcast or self.is_b_mcast or use_2cta_instrs):
a_full_mcast_mask = cpasync.create_tma_multicast_mask(
cluster_layout_vmnk, block_in_cluster_coord_vmnk, mcast_mode=2
)
b_full_mcast_mask = cpasync.create_tma_multicast_mask(
cluster_layout_vmnk, block_in_cluster_coord_vmnk, mcast_mode=1
)
sfa_full_mcast_mask = cpasync.create_tma_multicast_mask(
cluster_layout_vmnk, block_in_cluster_coord_vmnk, mcast_mode=2
)
sfb_full_mcast_mask = cpasync.create_tma_multicast_mask(
cluster_layout_sfb_vmnk, block_in_cluster_coord_sfb_vmnk, mcast_mode=1
)
tma_cache_policy_a = cutlass.Int64(cutlass.Int64(self.tma_cache_policy_a).ir_value())
tma_cache_policy_b = cutlass.Int64(cutlass.Int64(self.tma_cache_policy_b).ir_value())
tma_cache_policy_sfa = cutlass.Int64(cutlass.Int64(self.tma_cache_policy_sfa).ir_value())
tma_cache_policy_sfb = cutlass.Int64(cutlass.Int64(self.tma_cache_policy_sfb).ir_value())
tma_cache_policy_c = cutlass.Int64(cutlass.Int64(self.tma_cache_policy_c).ir_value())
gA_mkl = cute.local_tile(
mA_mkl, cute.slice_(self.mma_tiler, (None, 0, None)), (None, None, None)
)
gB_nkl = cute.local_tile(
mB_nkl, cute.slice_(self.mma_tiler, (0, None, None)), (None, None, None)
)
gSFA_mkl = cute.local_tile(
mSFA_mkl, cute.slice_(self.mma_tiler, (None, 0, None)), (None, None, None)
)
gSFB_nkl = cute.local_tile(
mSFB_nkl,
cute.slice_(self.mma_tiler_sfb, (0, None, None)),
(None, None, None),
)
gC_mnl = cute.local_tile(
mC_mnl, cute.slice_(self.mma_tiler, (None, None, 0)), (None, None, None)
)
k_tile_cnt = cute.size(gA_mkl, mode=[3])
thr_mma = tiled_mma.get_slice(mma_tile_coord_v)
thr_mma_sfb = tiled_mma_sfb.get_slice(mma_tile_coord_v)
tCgA = thr_mma.partition_A(gA_mkl)
tCgB = thr_mma.partition_B(gB_nkl)
tCgSFA = thr_mma.partition_A(gSFA_mkl)
tCgSFB = thr_mma_sfb.partition_B(gSFB_nkl)
tCgC = thr_mma.partition_C(gC_mnl)
a_cta_layout = cute.make_layout(
cute.slice_(cluster_layout_vmnk, (0, 0, None, 0)).shape
)
tAsA, tAgA = cpasync.tma_partition(
tma_atom_a,
block_in_cluster_coord_vmnk[2],
a_cta_layout,
cute.group_modes(sA, 0, 3),
cute.group_modes(tCgA, 0, 3),
)
b_cta_layout = cute.make_layout(
cute.slice_(cluster_layout_vmnk, (0, None, 0, 0)).shape
)
tBsB, tBgB = cpasync.tma_partition(
tma_atom_b,
block_in_cluster_coord_vmnk[1],
b_cta_layout,
cute.group_modes(sB, 0, 3),
cute.group_modes(tCgB, 0, 3),
)
sfa_cta_layout = a_cta_layout
tAsSFA, tAgSFA = cute.nvgpu.cpasync.tma_partition(
tma_atom_sfa,
block_in_cluster_coord_vmnk[2],
sfa_cta_layout,
cute.group_modes(sSFA, 0, 3),
cute.group_modes(tCgSFA, 0, 3),
)
tAsSFA = cute.filter_zeros(tAsSFA)
tAgSFA = cute.filter_zeros(tAgSFA)
sfb_cta_layout = cute.make_layout(
cute.slice_(cluster_layout_sfb_vmnk, (0, None, 0, 0)).shape
)
tBsSFB, tBgSFB = cute.nvgpu.cpasync.tma_partition(
tma_atom_sfb,
block_in_cluster_coord_sfb_vmnk[1],
sfb_cta_layout,
cute.group_modes(sSFB, 0, 3),
cute.group_modes(tCgSFB, 0, 3),
)
tBsSFB = cute.filter_zeros(tBsSFB)
tBgSFB = cute.filter_zeros(tBgSFB)
tCrA = tiled_mma.make_fragment_A(sA)
tCrB = tiled_mma.make_fragment_B(sB)
acc_shape = tiled_mma.partition_shape_C(self.mma_tiler[:2])
if cutlass.const_expr(self.overlapping_accum):
num_acc_stage_overlapped = 2
tCtAcc_fake = tiled_mma.make_fragment_C(
cute.append(acc_shape, num_acc_stage_overlapped)
)
tCtAcc_fake = cute.make_tensor(
tCtAcc_fake.iterator,
cute.make_layout(
tCtAcc_fake.shape,
stride = (
tCtAcc_fake.stride[0],
tCtAcc_fake.stride[1],
tCtAcc_fake.stride[2],
(256 - self.num_sf_tmem_cols) * tCtAcc_fake.stride[0][1]
)
)
)
else:
tCtAcc_fake = tiled_mma.make_fragment_C(
cute.append(acc_shape, self.num_acc_stage)
)
if warp_idx == self.tma_warp_id:
ab_producer_state = pipeline.make_pipeline_state(
pipeline.PipelineUserType.Producer, self.num_ab_stage
)
mma_tile_coord_mnl = (
bidx // cute.size(tiled_mma.thr_id.shape),
bidy,
bidz,
)
tAgA_slice = tAgA[(None, mma_tile_coord_mnl[0], None, mma_tile_coord_mnl[2])]
tBgB_slice = tBgB[(None, mma_tile_coord_mnl[1], None, mma_tile_coord_mnl[2])]
tAgSFA_slice = tAgSFA[(None, mma_tile_coord_mnl[0], None, mma_tile_coord_mnl[2])]
slice_n = mma_tile_coord_mnl[1]
if cutlass.const_expr(self.cta_tile_shape_mnk[1] == 64):
slice_n = mma_tile_coord_mnl[1] // 2
tBgSFB_slice = tBgSFB[(None, slice_n, None, mma_tile_coord_mnl[2])]
ab_producer_state.reset_count()
peek_ab_empty_status = cutlass.Boolean(1)
if ab_producer_state.count < k_tile_cnt:
peek_ab_empty_status = ab_pipeline.producer_try_acquire(ab_producer_state)
for k_tile in cutlass.range(0, k_tile_cnt, 1, unroll=1):
ab_pipeline.producer_acquire(ab_producer_state, peek_ab_empty_status)
try:
cute.copy(
tma_atom_a,
tAgA_slice[(None, ab_producer_state.count)],
tAsA[(None, ab_producer_state.index)],
tma_bar_ptr=ab_pipeline.producer_get_barrier(ab_producer_state),
mcast_mask=a_full_mcast_mask,
cache_policy=tma_cache_policy_a,
)
except TypeError:
cute.copy(
tma_atom_a,
tAgA_slice[(None, ab_producer_state.count)],
tAsA[(None, ab_producer_state.index)],
tma_bar_ptr=ab_pipeline.producer_get_barrier(ab_producer_state),
mcast_mask=a_full_mcast_mask,
)
try:
cute.copy(
tma_atom_b,
tBgB_slice[(None, ab_producer_state.count)],
tBsB[(None, ab_producer_state.index)],
tma_bar_ptr=ab_pipeline.producer_get_barrier(ab_producer_state),
mcast_mask=b_full_mcast_mask,
cache_policy=tma_cache_policy_b,
)
except TypeError:
cute.copy(
tma_atom_b,
tBgB_slice[(None, ab_producer_state.count)],
tBsB[(None, ab_producer_state.index)],
tma_bar_ptr=ab_pipeline.producer_get_barrier(ab_producer_state),
mcast_mask=b_full_mcast_mask,
)
try:
cute.copy(
tma_atom_sfa,
tAgSFA_slice[(None, ab_producer_state.count)],
tAsSFA[(None, ab_producer_state.index)],
tma_bar_ptr=ab_pipeline.producer_get_barrier(ab_producer_state),
mcast_mask=sfa_full_mcast_mask,
cache_policy=tma_cache_policy_sfa,
)
except TypeError:
cute.copy(
tma_atom_sfa,
tAgSFA_slice[(None, ab_producer_state.count)],
tAsSFA[(None, ab_producer_state.index)],
tma_bar_ptr=ab_pipeline.producer_get_barrier(ab_producer_state),
mcast_mask=sfa_full_mcast_mask,
)
try:
cute.copy(
tma_atom_sfb,
tBgSFB_slice[(None, ab_producer_state.count)],
tBsSFB[(None, ab_producer_state.index)],
tma_bar_ptr=ab_pipeline.producer_get_barrier(ab_producer_state),
mcast_mask=sfb_full_mcast_mask,
cache_policy=tma_cache_policy_sfb,
)
except TypeError:
cute.copy(
tma_atom_sfb,
tBgSFB_slice[(None, ab_producer_state.count)],
tBsSFB[(None, ab_producer_state.index)],
tma_bar_ptr=ab_pipeline.producer_get_barrier(ab_producer_state),
mcast_mask=sfb_full_mcast_mask,
)
ab_producer_state.advance()
peek_ab_empty_status = cutlass.Boolean(1)
if ab_producer_state.count < k_tile_cnt:
peek_ab_empty_status = ab_pipeline.producer_try_acquire(ab_producer_state)
ab_pipeline.producer_tail(ab_producer_state)
if warp_idx == self.mma_warp_id:
tmem.wait_for_alloc()
acc_tmem_ptr = tmem.retrieve_ptr(self.acc_dtype)
tCtAcc_base = cute.make_tensor(acc_tmem_ptr, tCtAcc_fake.layout)
sfa_tmem_ptr = cute.recast_ptr(
acc_tmem_ptr + self.num_accumulator_tmem_cols,
dtype=self.sf_dtype,
)
tCtSFA_layout = blockscaled_utils.make_tmem_layout_sfa(
tiled_mma,
self.mma_tiler,
self.sf_vec_size,
cute.slice_(sfa_smem_layout_staged, (None, None, None, 0)),
)
tCtSFA = cute.make_tensor(sfa_tmem_ptr, tCtSFA_layout)
sfb_tmem_ptr = cute.recast_ptr(
acc_tmem_ptr + self.num_accumulator_tmem_cols + self.num_sfa_tmem_cols,
dtype=self.sf_dtype,
)
tCtSFB_layout = blockscaled_utils.make_tmem_layout_sfb(
tiled_mma,
self.mma_tiler,
self.sf_vec_size,
cute.slice_(sfb_smem_layout_staged, (None, None, None, 0)),
)
tCtSFB = cute.make_tensor(sfb_tmem_ptr, tCtSFB_layout)
(
tiled_copy_s2t_sfa,
tCsSFA_compact_s2t,
tCtSFA_compact_s2t,
) = self.mainloop_s2t_copy_and_partition(sSFA, tCtSFA)
(
tiled_copy_s2t_sfb,
tCsSFB_compact_s2t,
tCtSFB_compact_s2t,
) = self.mainloop_s2t_copy_and_partition(sSFB, tCtSFB)
ab_consumer_state = pipeline.make_pipeline_state(
pipeline.PipelineUserType.Consumer, self.num_ab_stage
)
acc_producer_state = pipeline.make_pipeline_state(
pipeline.PipelineUserType.Producer, self.num_acc_stage
)
mma_tile_coord_mnl = (
bidx // cute.size(tiled_mma.thr_id.shape),
bidy,
bidz,
)
if cutlass.const_expr(self.overlapping_accum):
acc_stage_index = acc_producer_state.phase ^ 1
else:
acc_stage_index = acc_producer_state.index
tCtAcc = tCtAcc_base[(None, None, None, acc_stage_index)]
ab_consumer_state.reset_count()
peek_ab_full_status = cutlass.Boolean(1)
if ab_consumer_state.count < k_tile_cnt and is_leader_cta:
peek_ab_full_status = ab_pipeline.consumer_try_wait(
ab_consumer_state
)
if is_leader_cta:
acc_pipeline.producer_acquire(acc_producer_state)
tCtSFB_mma = tCtSFB
if cutlass.const_expr(self.cta_tile_shape_mnk[1] == 192):
offset = cutlass.Int32(2) if mma_tile_coord_mnl[1] % 2 == 1 else cutlass.Int32(0)
shifted_ptr = cute.recast_ptr(
acc_tmem_ptr
+ self.num_accumulator_tmem_cols
+ self.num_sfa_tmem_cols
+ offset,
dtype=self.sf_dtype,
)
tCtSFB_mma = cute.make_tensor(shifted_ptr, tCtSFB_layout)
elif cutlass.const_expr(self.cta_tile_shape_mnk[1] == 64):
offset = cutlass.Int32((mma_tile_coord_mnl[1] % 2) * 2)
shifted_ptr = cute.recast_ptr(
acc_tmem_ptr
+ self.num_accumulator_tmem_cols
+ self.num_sfa_tmem_cols
+ offset,
dtype=self.sf_dtype,
)
tCtSFB_mma = cute.make_tensor(shifted_ptr, tCtSFB_layout)
tiled_mma.set(tcgen05.Field.ACCUMULATE, False)
for k_tile in range(k_tile_cnt):
if is_leader_cta:
ab_pipeline.consumer_wait(
ab_consumer_state, peek_ab_full_status
)
s2t_stage_coord = (
None,
None,
None,
None,
ab_consumer_state.index,
)
tCsSFA_compact_s2t_staged = tCsSFA_compact_s2t[s2t_stage_coord]
tCsSFB_compact_s2t_staged = tCsSFB_compact_s2t[s2t_stage_coord]
cute.copy(
tiled_copy_s2t_sfa,
tCsSFA_compact_s2t_staged,
tCtSFA_compact_s2t,
)
cute.copy(
tiled_copy_s2t_sfb,
tCsSFB_compact_s2t_staged,
tCtSFB_compact_s2t,
)
num_kblocks = cute.size(tCrA, mode=[2])
for kblock_idx in cutlass.range(num_kblocks, unroll_full=True):
kblock_coord = (
None,
None,
kblock_idx,
ab_consumer_state.index,
)
sf_kblock_coord = (None, None, kblock_idx)
tiled_mma.set(
tcgen05.Field.SFA,
tCtSFA[sf_kblock_coord].iterator,
)
tiled_mma.set(
tcgen05.Field.SFB,
tCtSFB_mma[sf_kblock_coord].iterator,
)
cute.gemm(
tiled_mma,
tCtAcc,
tCrA[kblock_coord],
tCrB[kblock_coord],
tCtAcc,
)
tiled_mma.set(tcgen05.Field.ACCUMULATE, True)
ab_pipeline.consumer_release(ab_consumer_state)
ab_consumer_state.advance()
peek_ab_full_status = cutlass.Boolean(1)
if ab_consumer_state.count < k_tile_cnt:
if is_leader_cta:
peek_ab_full_status = ab_pipeline.consumer_try_wait(
ab_consumer_state
)
if is_leader_cta:
acc_pipeline.producer_commit(acc_producer_state)
acc_producer_state.advance()
acc_pipeline.producer_tail(acc_producer_state)
if warp_idx < self.mma_warp_id:
tmem.allocate(self.num_tmem_alloc_cols)
tmem.wait_for_alloc()
acc_tmem_ptr = tmem.retrieve_ptr(self.acc_dtype)
tCtAcc_base = cute.make_tensor(acc_tmem_ptr, tCtAcc_fake.layout)
epi_tidx = tidx
(
tiled_copy_t2r,
tTR_tAcc_base,
tTR_rAcc,
) = self.epilog_tmem_copy_and_partition(
epi_tidx, tCtAcc_base, tCgC, epi_tile, use_2cta_instrs
)
tTR_rC = cute.make_rmem_tensor(tTR_rAcc.shape, self.c_dtype)
tiled_copy_r2s, tRS_rC, tRS_sC = self.epilog_smem_copy_and_partition(
tiled_copy_t2r, tTR_rC, epi_tidx, sC
)
(
tma_atom_c,
bSG_sC,
bSG_gC_partitioned,
) = self.epilog_gmem_copy_and_partition(
epi_tidx, tma_atom_c, tCgC, epi_tile, sC
)
acc_consumer_state = pipeline.make_pipeline_state(
pipeline.PipelineUserType.Consumer, self.num_acc_stage
)
c_producer_group = pipeline.CooperativeGroup(
pipeline.Agent.Thread,
32 * len(self.epilog_warp_id),
)
c_pipeline = pipeline.PipelineTmaStore.create(
num_stages=self.num_c_stage,
producer_group=c_producer_group,
)
mma_tile_coord_mnl = (
bidx // cute.size(tiled_mma.thr_id.shape),
bidy,
bidz,
)
bSG_gC = bSG_gC_partitioned[
(
None,
None,
None,
*mma_tile_coord_mnl,
)
]
if cutlass.const_expr(self.overlapping_accum):
acc_stage_index = acc_consumer_state.phase
reverse_subtile = cutlass.Boolean(True) if acc_stage_index == 0 else cutlass.Boolean(False)
else:
acc_stage_index = acc_consumer_state.index
tTR_tAcc = tTR_tAcc_base[
(None, None, None, None, None, acc_stage_index)
]
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))
subtile_cnt = cute.size(tTR_tAcc.shape, mode=[3])
num_prev_subtiles = cutlass.Int32(0)
for subtile_idx in cutlass.range(subtile_cnt):
real_subtile_idx = subtile_idx
if cutlass.const_expr(self.overlapping_accum):
if reverse_subtile:
real_subtile_idx = self.cta_tile_shape_mnk[1] // self.epi_tile_n - 1 - subtile_idx
tTR_tAcc_mn = tTR_tAcc[(None, None, None, real_subtile_idx)]
cute.copy(tiled_copy_t2r, tTR_tAcc_mn, tTR_rAcc)
if cutlass.const_expr(self.overlapping_accum):
if subtile_idx == self.iter_acc_early_release_in_epilogue:
cute.arch.fence_view_async_tmem_load()
with cute.arch.elect_one():
acc_pipeline.consumer_release(acc_consumer_state)
acc_consumer_state.advance()
acc_vec = tiled_copy_r2s.retile(tTR_rAcc).load()
acc_vec = epilogue_op(acc_vec.to(self.c_dtype))
tRS_rC.store(acc_vec)
c_buffer = (num_prev_subtiles + real_subtile_idx) % self.num_c_stage
cute.copy(
tiled_copy_r2s,
tRS_rC,
tRS_sC[(None, None, None, c_buffer)],
)
cute.arch.fence_proxy(
cute.arch.ProxyKind.async_shared,
space=cute.arch.SharedSpace.shared_cta,
)
self.epilog_sync_barrier.arrive_and_wait()
if warp_idx == self.epilog_warp_id[0]:
try:
cute.copy(
tma_atom_c,
bSG_sC[(None, c_buffer)],
bSG_gC[(None, real_subtile_idx)],
cache_policy=tma_cache_policy_c,
)
except TypeError:
cute.copy(
tma_atom_c,
bSG_sC[(None, c_buffer)],
bSG_gC[(None, real_subtile_idx)],
)
c_pipeline.producer_commit()
c_pipeline.producer_acquire()
self.epilog_sync_barrier.arrive_and_wait()
if cutlass.const_expr(not self.overlapping_accum):
with cute.arch.elect_one():
acc_pipeline.consumer_release(acc_consumer_state)
acc_consumer_state.advance()
tmem.relinquish_alloc_permit()
self.epilog_sync_barrier.arrive_and_wait()
tmem.free(acc_tmem_ptr)
c_pipeline.producer_tail()
def mainloop_s2t_copy_and_partition(
self,
sSF: cute.Tensor,
tSF: cute.Tensor,
) -> Tuple[cute.TiledCopy, cute.Tensor, cute.Tensor]:
tCsSF_compact = cute.filter_zeros(sSF)
tCtSF_compact = cute.filter_zeros(tSF)
copy_atom_s2t = cute.make_copy_atom(
tcgen05.Cp4x32x128bOp(self.cta_group),
self.sf_dtype,
)
tiled_copy_s2t = tcgen05.make_s2t_copy(copy_atom_s2t, tCtSF_compact)
thr_copy_s2t = tiled_copy_s2t.get_slice(0)
tCsSF_compact_s2t_ = thr_copy_s2t.partition_S(tCsSF_compact)
tCsSF_compact_s2t = tcgen05.get_s2t_smem_desc_tensor(
tiled_copy_s2t, tCsSF_compact_s2t_
)
tCtSF_compact_s2t = thr_copy_s2t.partition_D(tCtSF_compact)
return tiled_copy_s2t, tCsSF_compact_s2t, tCtSF_compact_s2t
def epilog_tmem_copy_and_partition(
self,
tidx: cutlass.Int32,
tAcc: cute.Tensor,
gC_mnl: cute.Tensor,
epi_tile: cute.Tile,
use_2cta_instrs: Union[cutlass.Boolean, bool],
) -> Tuple[cute.TiledCopy, cute.Tensor, cute.Tensor]:
copy_atom_t2r = sm100_utils.get_tmem_load_op(
self.cta_tile_shape_mnk,
self.c_layout,
self.c_dtype,
self.acc_dtype,
epi_tile,
use_2cta_instrs,
)
tAcc_epi = cute.flat_divide(
tAcc[((None, None), 0, 0, None)],
epi_tile,
)
tiled_copy_t2r = tcgen05.make_tmem_copy(
copy_atom_t2r, tAcc_epi[(None, None, 0, 0, 0)]
)
thr_copy_t2r = tiled_copy_t2r.get_slice(tidx)
tTR_tAcc = thr_copy_t2r.partition_S(tAcc_epi)
gC_mnl_epi = cute.flat_divide(
gC_mnl[((None, None), 0, 0, None, None, None)], epi_tile
)
tTR_gC = thr_copy_t2r.partition_D(gC_mnl_epi)
tTR_rAcc = cute.make_rmem_tensor(
tTR_gC[(None, None, None, 0, 0, 0, 0, 0)].shape, self.acc_dtype
)
return tiled_copy_t2r, tTR_tAcc, tTR_rAcc
def epilog_smem_copy_and_partition(
self,
tiled_copy_t2r: cute.TiledCopy,
tTR_rC: cute.Tensor,
tidx: cutlass.Int32,
sC: cute.Tensor,
) -> Tuple[cute.TiledCopy, cute.Tensor, cute.Tensor]:
copy_atom_r2s = sm100_utils.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)
thr_copy_r2s = tiled_copy_r2s.get_slice(tidx)
tRS_sC = thr_copy_r2s.partition_D(sC)
tRS_rC = tiled_copy_r2s.retile(tTR_rC)
return tiled_copy_r2s, tRS_rC, tRS_sC
def epilog_gmem_copy_and_partition(
self,
tidx: cutlass.Int32,
atom: Union[cute.CopyAtom, cute.TiledCopy],
gC_mnl: cute.Tensor,
epi_tile: cute.Tile,
sC: cute.Tensor,
) -> Tuple[cute.CopyAtom, cute.Tensor, cute.Tensor]:
gC_epi = cute.flat_divide(
gC_mnl[((None, None), 0, 0, None, None, None)], epi_tile
)
tma_atom_c = atom
sC_for_tma_partition = cute.group_modes(sC, 0, 2)
gC_for_tma_partition = cute.group_modes(gC_epi, 0, 2)
bSG_sC, bSG_gC = cpasync.tma_partition(
tma_atom_c,
0,
cute.make_layout(1),
sC_for_tma_partition,
gC_for_tma_partition,
)
return tma_atom_c, bSG_sC, bSG_gC
@staticmethod
def _compute_stages(
tiled_mma: cute.TiledMma,
mma_tiler_mnk: Tuple[int, int, int],
a_dtype: Type[cutlass.Numeric],
b_dtype: Type[cutlass.Numeric],
epi_tile: cute.Tile,
c_dtype: Type[cutlass.Numeric],
c_layout: utils.LayoutEnum,
sf_dtype: Type[cutlass.Numeric],
sf_vec_size: int,
smem_capacity: int,
occupancy: int,
) -> Tuple[int, int, int]:
num_acc_stage = 1 if mma_tiler_mnk[1] == 256 else 2
a_smem_layout_stage_one = sm100_utils.make_smem_layout_a(
tiled_mma,
mma_tiler_mnk,
a_dtype,
1,
)
b_smem_layout_staged_one = sm100_utils.make_smem_layout_b(
tiled_mma,
mma_tiler_mnk,
b_dtype,
1,
)
sfa_smem_layout_staged_one = blockscaled_utils.make_smem_layout_sfa(
tiled_mma,
mma_tiler_mnk,
sf_vec_size,
1,
)
sfb_smem_layout_staged_one = blockscaled_utils.make_smem_layout_sfb(
tiled_mma,
mma_tiler_mnk,
sf_vec_size,
1,
)
c_smem_layout_staged_one = sm100_utils.make_smem_layout_epi(
c_dtype,
c_layout,
epi_tile,
1,
)
ab_bytes_per_stage = (
cute.size_in_bytes(a_dtype, a_smem_layout_stage_one)
+ cute.size_in_bytes(b_dtype, b_smem_layout_staged_one)
+ cute.size_in_bytes(sf_dtype, sfa_smem_layout_staged_one)
+ cute.size_in_bytes(sf_dtype, sfb_smem_layout_staged_one)
)
mbar_helpers_bytes = 1024
c_bytes_per_stage = cute.size_in_bytes(c_dtype, c_smem_layout_staged_one)
hard_max_ab_stage = 6
hard_max_c_stage = 4
if occupancy >= 2:
max_ab_stage = 3
if mma_tiler_mnk[1] <= 128:
max_ab_stage = 4
if mma_tiler_mnk[1] <= 64:
max_ab_stage = 5
if mma_tiler_mnk[1] == 256:
max_ab_stage = 3
if ab_bytes_per_stage <= 48 * 1024:
max_ab_stage += 1
max_c_stage = 2
if c_bytes_per_stage <= 32 * 1024:
max_c_stage = 3
else:
max_ab_stage = 5
if mma_tiler_mnk[1] == 256:
max_ab_stage = 4
max_c_stage = 3
if max_ab_stage > hard_max_ab_stage:
max_ab_stage = hard_max_ab_stage
if max_c_stage > hard_max_c_stage:
max_c_stage = hard_max_c_stage
smem_per_block_max = smem_capacity // max(1, occupancy)
num_c_stage = 2
avail_for_ab = smem_per_block_max - (
mbar_helpers_bytes + c_bytes_per_stage * num_c_stage
)
if avail_for_ab < 0:
avail_for_ab = 0
num_ab_stage = avail_for_ab // ab_bytes_per_stage
if num_ab_stage < 1:
num_ab_stage = 1
if num_ab_stage > max_ab_stage:
num_ab_stage = max_ab_stage
remain_for_c = smem_per_block_max - (
mbar_helpers_bytes + ab_bytes_per_stage * num_ab_stage
)
max_c_possible = remain_for_c // c_bytes_per_stage if c_bytes_per_stage > 0 else 1
if max_c_possible < 1:
max_c_possible = 1
if max_c_possible < 2:
num_c_stage = max_c_possible
else:
num_c_stage = max_c_possible
if num_c_stage > max_c_stage:
num_c_stage = max_c_stage
return num_acc_stage, num_ab_stage, num_c_stage
@staticmethod
def _compute_grid(
c: cute.Tensor,
cta_tile_shape_mnk: Tuple[int, int, int],
cluster_shape_mn: Tuple[int, int],
max_active_clusters: cutlass.Constexpr,
) -> Tuple[utils.PersistentTileSchedulerParams, Tuple[int, int, int]]:
c_shape = cute.slice_(cta_tile_shape_mnk, (None, None, 0))
gc = cute.zipped_divide(c, tiler=c_shape)
num_ctas_mnl = gc[(0, (None, None, None))].shape
cluster_shape_mnl = (*cluster_shape_mn, 1)
tile_sched_params = utils.PersistentTileSchedulerParams(
num_ctas_mnl, cluster_shape_mnl
)
grid = utils.StaticPersistentTileScheduler.get_grid_shape(
tile_sched_params, max_active_clusters
)
return tile_sched_params, grid
@staticmethod
def is_valid_dtypes_and_scale_factor_vec_size(
ab_dtype: Type[cutlass.Numeric],
sf_dtype: Type[cutlass.Numeric],
sf_vec_size: int,
c_dtype: Type[cutlass.Numeric],
) -> bool:
is_valid = True
if ab_dtype not in {
cutlass.Float4E2M1FN,
cutlass.Float8E5M2,
cutlass.Float8E4M3FN,
}:
is_valid = False
if sf_vec_size not in {16, 32}:
is_valid = False
if sf_dtype not in {cutlass.Float8E8M0FNU, cutlass.Float8E4M3FN}:
is_valid = False
if sf_dtype == cutlass.Float8E4M3FN and sf_vec_size == 32:
is_valid = False
if ab_dtype in {cutlass.Float8E5M2, cutlass.Float8E4M3FN} and sf_vec_size == 16:
is_valid = False
if c_dtype not in {
cutlass.Float32,
cutlass.Float16,
cutlass.BFloat16,
cutlass.Float8E5M2,
cutlass.Float8E4M3FN,
}:
is_valid = False
return is_valid
@staticmethod
def is_valid_layouts(
ab_dtype: Type[cutlass.Numeric],
c_dtype: Type[cutlass.Numeric],
a_major: str,
b_major: str,
c_major: str,
) -> bool:
is_valid = True
if ab_dtype is cutlass.Float4E2M1FN and not (a_major == "k" and b_major == "k"):
is_valid = False
return is_valid
@staticmethod
def is_valid_mma_tiler_and_cluster_shape(
mma_tiler_mn: Tuple[int, int],
cluster_shape_mn: Tuple[int, int],
) -> bool:
is_valid = True
if mma_tiler_mn[0] not in [128, 256]:
is_valid = False
if mma_tiler_mn[1] not in [64, 128, 192, 256]:
is_valid = False
if cluster_shape_mn[0] % (2 if mma_tiler_mn[0] == 256 else 1) != 0:
is_valid = False
is_power_of_2 = lambda x: x > 0 and (x & (x - 1)) == 0
if (
cluster_shape_mn[0] * cluster_shape_mn[1] > 16
or cluster_shape_mn[0] <= 0
or cluster_shape_mn[1] <= 0
or cluster_shape_mn[0] > 4
or cluster_shape_mn[1] > 4
or not is_power_of_2(cluster_shape_mn[0])
or not is_power_of_2(cluster_shape_mn[1])
):
is_valid = False
return is_valid
@staticmethod
def is_valid_tensor_alignment(
m: int,
n: int,
k: int,
l: int,
ab_dtype: Type[cutlass.Numeric],
c_dtype: Type[cutlass.Numeric],
a_major: str,
b_major: str,
c_major: str,
) -> bool:
is_valid = True
def check_contigous_16B_alignment(dtype, is_mode0_major, tensor_shape):
major_mode_idx = 0 if is_mode0_major else 1
num_major_elements = tensor_shape[major_mode_idx]
num_contiguous_elements = 16 * 8 // dtype.width
return num_major_elements % num_contiguous_elements == 0
if (
not check_contigous_16B_alignment(ab_dtype, a_major == "m", (m, k, l))
or not check_contigous_16B_alignment(ab_dtype, b_major == "n", (n, k, l))
or not check_contigous_16B_alignment(c_dtype, c_major == "m", (m, n, l))
):
is_valid = False
return is_valid
@staticmethod
def can_implement(
ab_dtype: Type[cutlass.Numeric],
sf_dtype: Type[cutlass.Numeric],
sf_vec_size: int,
c_dtype: Type[cutlass.Numeric],
mma_tiler_mn: Tuple[int, int],
cluster_shape_mn: Tuple[int, int],
m: int,
n: int,
k: int,
l: int,
a_major: str,
b_major: str,
c_major: str,
) -> bool:
can_implement = True
if not Sm100BlockScaledPersistentDenseGemmKernel.is_valid_dtypes_and_scale_factor_vec_size(
ab_dtype, sf_dtype, sf_vec_size, c_dtype
):
can_implement = False
if not Sm100BlockScaledPersistentDenseGemmKernel.is_valid_layouts(
ab_dtype, c_dtype, a_major, b_major, c_major
):
can_implement = False
if not Sm100BlockScaledPersistentDenseGemmKernel.is_valid_mma_tiler_and_cluster_shape(
mma_tiler_mn, cluster_shape_mn
):
can_implement = False
if not Sm100BlockScaledPersistentDenseGemmKernel.is_valid_tensor_alignment(
m, n, k, l, ab_dtype, c_dtype, a_major, b_major, c_major
):
can_implement = False
return can_implement
class Sm100BlockScaledPersistentDenseGemmKernelGroupG8(Sm100BlockScaledPersistentDenseGemmKernel):
@cute.jit
def __call__(
self,
a0_ptr: cute.Pointer,
b0_ptr: cute.Pointer,
sfa0_ptr: cute.Pointer,
sfb0_ptr: cute.Pointer,
c0_ptr: cute.Pointer,
a1_ptr: cute.Pointer,
b1_ptr: cute.Pointer,
sfa1_ptr: cute.Pointer,
sfb1_ptr: cute.Pointer,
c1_ptr: cute.Pointer,
a2_ptr: cute.Pointer,
b2_ptr: cute.Pointer,
sfa2_ptr: cute.Pointer,
sfb2_ptr: cute.Pointer,
c2_ptr: cute.Pointer,
a3_ptr: cute.Pointer,
b3_ptr: cute.Pointer,
sfa3_ptr: cute.Pointer,
sfb3_ptr: cute.Pointer,
c3_ptr: cute.Pointer,
a4_ptr: cute.Pointer,
b4_ptr: cute.Pointer,
sfa4_ptr: cute.Pointer,
sfb4_ptr: cute.Pointer,
c4_ptr: cute.Pointer,
a5_ptr: cute.Pointer,
b5_ptr: cute.Pointer,
sfa5_ptr: cute.Pointer,
sfb5_ptr: cute.Pointer,
c5_ptr: cute.Pointer,
a6_ptr: cute.Pointer,
b6_ptr: cute.Pointer,
sfa6_ptr: cute.Pointer,
sfb6_ptr: cute.Pointer,
c6_ptr: cute.Pointer,
a7_ptr: cute.Pointer,
b7_ptr: cute.Pointer,
sfa7_ptr: cute.Pointer,
sfb7_ptr: cute.Pointer,
c7_ptr: cute.Pointer,
problem_size: tuple,
max_active_clusters: cutlass.Constexpr,
epilogue_op: cutlass.Constexpr = lambda x: x,
):
m0, m1, m2, m3, m4, m5, m6, m7, m_max, n, k, l = problem_size
sf_k = k // self.sf_vec_size
a_tensor0 = cute.make_tensor(
a0_ptr,
cute.make_layout(
(m0, cute.assume(k, 32), l),
stride=(cute.assume(k, 32), 1, cute.assume(m0 * k, 32)),
),
)
b_tensor0 = cute.make_tensor(
b0_ptr,
cute.make_layout(
(n, cute.assume(k, 32), l),
stride=(cute.assume(k, 32), 1, cute.assume(n * k, 32)),
),
)
c_tensor0 = cute.make_tensor(c0_ptr, cute.make_layout((m0, n, l), stride=(n, 1, m0 * n)))
sfa_tensor0 = cute.make_tensor(
sfa0_ptr,
cute.make_layout((m0, sf_k, l), stride=(sf_k, 1, m0 * sf_k)),
)
sfb_tensor0 = cute.make_tensor(
sfb0_ptr,
cute.make_layout((n, sf_k, l), stride=(sf_k, 1, n * sf_k)),
)
a_tensor1 = cute.make_tensor(
a1_ptr,
cute.make_layout(
(m1, cute.assume(k, 32), l),
stride=(cute.assume(k, 32), 1, cute.assume(m1 * k, 32)),
),
)
b_tensor1 = cute.make_tensor(
b1_ptr,
cute.make_layout(
(n, cute.assume(k, 32), l),
stride=(cute.assume(k, 32), 1, cute.assume(n * k, 32)),
),
)
c_tensor1 = cute.make_tensor(c1_ptr, cute.make_layout((m1, n, l), stride=(n, 1, m1 * n)))
sfa_tensor1 = cute.make_tensor(
sfa1_ptr,
cute.make_layout((m1, sf_k, l), stride=(sf_k, 1, m1 * sf_k)),
)
sfb_tensor1 = cute.make_tensor(
sfb1_ptr,
cute.make_layout((n, sf_k, l), stride=(sf_k, 1, n * sf_k)),
)
a_tensor2 = cute.make_tensor(
a2_ptr,
cute.make_layout(
(m2, cute.assume(k, 32), l),
stride=(cute.assume(k, 32), 1, cute.assume(m2 * k, 32)),
),
)
b_tensor2 = cute.make_tensor(
b2_ptr,
cute.make_layout(
(n, cute.assume(k, 32), l),
stride=(cute.assume(k, 32), 1, cute.assume(n * k, 32)),
),
)
c_tensor2 = cute.make_tensor(c2_ptr, cute.make_layout((m2, n, l), stride=(n, 1, m2 * n)))
sfa_tensor2 = cute.make_tensor(
sfa2_ptr,
cute.make_layout((m2, sf_k, l), stride=(sf_k, 1, m2 * sf_k)),
)
sfb_tensor2 = cute.make_tensor(
sfb2_ptr,
cute.make_layout((n, sf_k, l), stride=(sf_k, 1, n * sf_k)),
)
a_tensor3 = cute.make_tensor(
a3_ptr,
cute.make_layout(
(m3, cute.assume(k, 32), l),
stride=(cute.assume(k, 32), 1, cute.assume(m3 * k, 32)),
),
)
b_tensor3 = cute.make_tensor(
b3_ptr,
cute.make_layout(
(n, cute.assume(k, 32), l),
stride=(cute.assume(k, 32), 1, cute.assume(n * k, 32)),
),
)
c_tensor3 = cute.make_tensor(c3_ptr, cute.make_layout((m3, n, l), stride=(n, 1, m3 * n)))
sfa_tensor3 = cute.make_tensor(
sfa3_ptr,
cute.make_layout((m3, sf_k, l), stride=(sf_k, 1, m3 * sf_k)),
)
sfb_tensor3 = cute.make_tensor(
sfb3_ptr,
cute.make_layout((n, sf_k, l), stride=(sf_k, 1, n * sf_k)),
)
a_tensor4 = cute.make_tensor(
a4_ptr,
cute.make_layout(
(m4, cute.assume(k, 32), l),
stride=(cute.assume(k, 32), 1, cute.assume(m4 * k, 32)),
),
)
b_tensor4 = cute.make_tensor(
b4_ptr,
cute.make_layout(
(n, cute.assume(k, 32), l),
stride=(cute.assume(k, 32), 1, cute.assume(n * k, 32)),
),
)
c_tensor4 = cute.make_tensor(c4_ptr, cute.make_layout((m4, n, l), stride=(n, 1, m4 * n)))
sfa_tensor4 = cute.make_tensor(
sfa4_ptr,
cute.make_layout((m4, sf_k, l), stride=(sf_k, 1, m4 * sf_k)),
)
sfb_tensor4 = cute.make_tensor(
sfb4_ptr,
cute.make_layout((n, sf_k, l), stride=(sf_k, 1, n * sf_k)),
)
a_tensor5 = cute.make_tensor(
a5_ptr,
cute.make_layout(
(m5, cute.assume(k, 32), l),
stride=(cute.assume(k, 32), 1, cute.assume(m5 * k, 32)),
),
)
b_tensor5 = cute.make_tensor(
b5_ptr,
cute.make_layout(
(n, cute.assume(k, 32), l),
stride=(cute.assume(k, 32), 1, cute.assume(n * k, 32)),
),
)
c_tensor5 = cute.make_tensor(c5_ptr, cute.make_layout((m5, n, l), stride=(n, 1, m5 * n)))
sfa_tensor5 = cute.make_tensor(
sfa5_ptr,
cute.make_layout((m5, sf_k, l), stride=(sf_k, 1, m5 * sf_k)),
)
sfb_tensor5 = cute.make_tensor(
sfb5_ptr,
cute.make_layout((n, sf_k, l), stride=(sf_k, 1, n * sf_k)),
)
a_tensor6 = cute.make_tensor(
a6_ptr,
cute.make_layout(
(m6, cute.assume(k, 32), l),
stride=(cute.assume(k, 32), 1, cute.assume(m6 * k, 32)),
),
)
b_tensor6 = cute.make_tensor(
b6_ptr,
cute.make_layout(
(n, cute.assume(k, 32), l),
stride=(cute.assume(k, 32), 1, cute.assume(n * k, 32)),
),
)
c_tensor6 = cute.make_tensor(c6_ptr, cute.make_layout((m6, n, l), stride=(n, 1, m6 * n)))
sfa_tensor6 = cute.make_tensor(
sfa6_ptr,
cute.make_layout((m6, sf_k, l), stride=(sf_k, 1, m6 * sf_k)),
)
sfb_tensor6 = cute.make_tensor(
sfb6_ptr,
cute.make_layout((n, sf_k, l), stride=(sf_k, 1, n * sf_k)),
)
a_tensor7 = cute.make_tensor(
a7_ptr,
cute.make_layout(
(m7, cute.assume(k, 32), l),
stride=(cute.assume(k, 32), 1, cute.assume(m7 * k, 32)),
),
)
b_tensor7 = cute.make_tensor(
b7_ptr,
cute.make_layout(
(n, cute.assume(k, 32), l),
stride=(cute.assume(k, 32), 1, cute.assume(n * k, 32)),
),
)
c_tensor7 = cute.make_tensor(c7_ptr, cute.make_layout((m7, n, l), stride=(n, 1, m7 * n)))
sfa_tensor7 = cute.make_tensor(
sfa7_ptr,
cute.make_layout((m7, sf_k, l), stride=(sf_k, 1, m7 * sf_k)),
)
sfb_tensor7 = cute.make_tensor(
sfb7_ptr,
cute.make_layout((n, sf_k, l), stride=(sf_k, 1, n * sf_k)),
)
a_tensor0.mark_compact_shape_dynamic(mode=1, stride_order=(2, 0, 1), divisibility=2)
b_tensor0.mark_compact_shape_dynamic(mode=1, stride_order=(2, 0, 1), divisibility=2)
c_tensor0.mark_compact_shape_dynamic(mode=1, stride_order=(2, 0, 1), divisibility=1)
a_tensor1.mark_compact_shape_dynamic(mode=1, stride_order=(2, 0, 1), divisibility=2)
b_tensor1.mark_compact_shape_dynamic(mode=1, stride_order=(2, 0, 1), divisibility=2)
c_tensor1.mark_compact_shape_dynamic(mode=1, stride_order=(2, 0, 1), divisibility=1)
a_tensor2.mark_compact_shape_dynamic(mode=1, stride_order=(2, 0, 1), divisibility=2)
b_tensor2.mark_compact_shape_dynamic(mode=1, stride_order=(2, 0, 1), divisibility=2)
c_tensor2.mark_compact_shape_dynamic(mode=1, stride_order=(2, 0, 1), divisibility=1)
a_tensor3.mark_compact_shape_dynamic(mode=1, stride_order=(2, 0, 1), divisibility=2)
b_tensor3.mark_compact_shape_dynamic(mode=1, stride_order=(2, 0, 1), divisibility=2)
c_tensor3.mark_compact_shape_dynamic(mode=1, stride_order=(2, 0, 1), divisibility=1)
a_tensor4.mark_compact_shape_dynamic(mode=1, stride_order=(2, 0, 1), divisibility=2)
b_tensor4.mark_compact_shape_dynamic(mode=1, stride_order=(2, 0, 1), divisibility=2)
c_tensor4.mark_compact_shape_dynamic(mode=1, stride_order=(2, 0, 1), divisibility=1)
a_tensor5.mark_compact_shape_dynamic(mode=1, stride_order=(2, 0, 1), divisibility=2)
b_tensor5.mark_compact_shape_dynamic(mode=1, stride_order=(2, 0, 1), divisibility=2)
c_tensor5.mark_compact_shape_dynamic(mode=1, stride_order=(2, 0, 1), divisibility=1)
a_tensor6.mark_compact_shape_dynamic(mode=1, stride_order=(2, 0, 1), divisibility=2)
b_tensor6.mark_compact_shape_dynamic(mode=1, stride_order=(2, 0, 1), divisibility=2)
c_tensor6.mark_compact_shape_dynamic(mode=1, stride_order=(2, 0, 1), divisibility=1)
a_tensor7.mark_compact_shape_dynamic(mode=1, stride_order=(2, 0, 1), divisibility=2)
b_tensor7.mark_compact_shape_dynamic(mode=1, stride_order=(2, 0, 1), divisibility=2)
c_tensor7.mark_compact_shape_dynamic(mode=1, stride_order=(2, 0, 1), divisibility=1)
self.a_dtype: Type[cutlass.Numeric] = a_tensor0.element_type
self.b_dtype: Type[cutlass.Numeric] = b_tensor0.element_type
self.sf_dtype: Type[cutlass.Numeric] = sfa_tensor0.element_type
self.c_dtype: Type[cutlass.Numeric] = c_tensor0.element_type
self.a_major_mode = utils.LayoutEnum.from_tensor(a_tensor0).mma_major_mode()
self.b_major_mode = utils.LayoutEnum.from_tensor(b_tensor0).mma_major_mode()
self.c_layout = utils.LayoutEnum.from_tensor(c_tensor0)
if cutlass.const_expr(self.a_dtype != self.b_dtype):
raise TypeError(f"Type must match: {self.a_dtype} != {self.b_dtype}")
self._setup_attributes()
sfa_layout0 = blockscaled_utils.tile_atom_to_shape_SF(a_tensor0.shape, self.sf_vec_size)
sfa_tensor0 = cute.make_tensor(sfa_tensor0.iterator, sfa_layout0)
sfb_layout0 = blockscaled_utils.tile_atom_to_shape_SF(b_tensor0.shape, self.sf_vec_size)
sfb_tensor0 = cute.make_tensor(sfb_tensor0.iterator, sfb_layout0)
sfa_layout1 = blockscaled_utils.tile_atom_to_shape_SF(a_tensor1.shape, self.sf_vec_size)
sfa_tensor1 = cute.make_tensor(sfa_tensor1.iterator, sfa_layout1)
sfb_layout1 = blockscaled_utils.tile_atom_to_shape_SF(b_tensor1.shape, self.sf_vec_size)
sfb_tensor1 = cute.make_tensor(sfb_tensor1.iterator, sfb_layout1)
sfa_layout2 = blockscaled_utils.tile_atom_to_shape_SF(a_tensor2.shape, self.sf_vec_size)
sfa_tensor2 = cute.make_tensor(sfa_tensor2.iterator, sfa_layout2)
sfb_layout2 = blockscaled_utils.tile_atom_to_shape_SF(b_tensor2.shape, self.sf_vec_size)
sfb_tensor2 = cute.make_tensor(sfb_tensor2.iterator, sfb_layout2)
sfa_layout3 = blockscaled_utils.tile_atom_to_shape_SF(a_tensor3.shape, self.sf_vec_size)
sfa_tensor3 = cute.make_tensor(sfa_tensor3.iterator, sfa_layout3)
sfb_layout3 = blockscaled_utils.tile_atom_to_shape_SF(b_tensor3.shape, self.sf_vec_size)
sfb_tensor3 = cute.make_tensor(sfb_tensor3.iterator, sfb_layout3)
sfa_layout4 = blockscaled_utils.tile_atom_to_shape_SF(a_tensor4.shape, self.sf_vec_size)
sfa_tensor4 = cute.make_tensor(sfa_tensor4.iterator, sfa_layout4)
sfb_layout4 = blockscaled_utils.tile_atom_to_shape_SF(b_tensor4.shape, self.sf_vec_size)
sfb_tensor4 = cute.make_tensor(sfb_tensor4.iterator, sfb_layout4)
sfa_layout5 = blockscaled_utils.tile_atom_to_shape_SF(a_tensor5.shape, self.sf_vec_size)
sfa_tensor5 = cute.make_tensor(sfa_tensor5.iterator, sfa_layout5)
sfb_layout5 = blockscaled_utils.tile_atom_to_shape_SF(b_tensor5.shape, self.sf_vec_size)
sfb_tensor5 = cute.make_tensor(sfb_tensor5.iterator, sfb_layout5)
sfa_layout6 = blockscaled_utils.tile_atom_to_shape_SF(a_tensor6.shape, self.sf_vec_size)
sfa_tensor6 = cute.make_tensor(sfa_tensor6.iterator, sfa_layout6)
sfb_layout6 = blockscaled_utils.tile_atom_to_shape_SF(b_tensor6.shape, self.sf_vec_size)
sfb_tensor6 = cute.make_tensor(sfb_tensor6.iterator, sfb_layout6)
sfa_layout7 = blockscaled_utils.tile_atom_to_shape_SF(a_tensor7.shape, self.sf_vec_size)
sfa_tensor7 = cute.make_tensor(sfa_tensor7.iterator, sfa_layout7)
sfb_layout7 = blockscaled_utils.tile_atom_to_shape_SF(b_tensor7.shape, self.sf_vec_size)
sfb_tensor7 = cute.make_tensor(sfb_tensor7.iterator, sfb_layout7)
tiled_mma = sm100_utils.make_blockscaled_trivial_tiled_mma(
self.a_dtype,
self.a_major_mode,
self.b_major_mode,
self.sf_dtype,
self.sf_vec_size,
self.cta_group,
self.mma_inst_shape_mn,
)
tiled_mma_sfb = sm100_utils.make_blockscaled_trivial_tiled_mma(
self.a_dtype,
self.a_major_mode,
self.b_major_mode,
self.sf_dtype,
self.sf_vec_size,
cute.nvgpu.tcgen05.CtaGroup.ONE,
self.mma_inst_shape_mn_sfb,
)
atom_thr_size = cute.size(tiled_mma.thr_id.shape)
a_op = sm100_utils.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))
b_op = sm100_utils.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))
sfa_op = sm100_utils.cluster_shape_to_tma_atom_A(self.cluster_shape_mn, tiled_mma.thr_id)
sfa_smem_layout = cute.slice_(self.sfa_smem_layout_staged, (None, None, None, 0))
sfb_op = sm100_utils.cluster_shape_to_tma_atom_SFB(self.cluster_shape_mn, tiled_mma.thr_id)
sfb_smem_layout = cute.slice_(self.sfb_smem_layout_staged, (None, None, None, 0))
tma_atom_a0, tma_tensor_a0 = cute.nvgpu.make_tiled_tma_atom_A(
a_op,
a_tensor0,
a_smem_layout,
self.mma_tiler,
tiled_mma,
self.cluster_layout_vmnk.shape,
)
tma_atom_b0, tma_tensor_b0 = cute.nvgpu.make_tiled_tma_atom_B(
b_op,
b_tensor0,
b_smem_layout,
self.mma_tiler,
tiled_mma,
self.cluster_layout_vmnk.shape,
)
tma_atom_sfa0, tma_tensor_sfa0 = cute.nvgpu.make_tiled_tma_atom_A(
sfa_op,
sfa_tensor0,
sfa_smem_layout,
self.mma_tiler,
tiled_mma,
self.cluster_layout_vmnk.shape,
internal_type=cutlass.Int16,
)
tma_atom_sfb0, tma_tensor_sfb0 = cute.nvgpu.make_tiled_tma_atom_B(
sfb_op,
sfb_tensor0,
sfb_smem_layout,
self.mma_tiler_sfb,
tiled_mma_sfb,
self.cluster_layout_sfb_vmnk.shape,
internal_type=cutlass.Int16,
)
tma_atom_a1, tma_tensor_a1 = cute.nvgpu.make_tiled_tma_atom_A(
a_op,
a_tensor1,
a_smem_layout,
self.mma_tiler,
tiled_mma,
self.cluster_layout_vmnk.shape,
)
tma_atom_b1, tma_tensor_b1 = cute.nvgpu.make_tiled_tma_atom_B(
b_op,
b_tensor1,
b_smem_layout,
self.mma_tiler,
tiled_mma,
self.cluster_layout_vmnk.shape,
)
tma_atom_sfa1, tma_tensor_sfa1 = cute.nvgpu.make_tiled_tma_atom_A(
sfa_op,
sfa_tensor1,
sfa_smem_layout,
self.mma_tiler,
tiled_mma,
self.cluster_layout_vmnk.shape,
internal_type=cutlass.Int16,
)
tma_atom_sfb1, tma_tensor_sfb1 = cute.nvgpu.make_tiled_tma_atom_B(
sfb_op,
sfb_tensor1,
sfb_smem_layout,
self.mma_tiler_sfb,
tiled_mma_sfb,
self.cluster_layout_sfb_vmnk.shape,
internal_type=cutlass.Int16,
)
tma_atom_a2, tma_tensor_a2 = cute.nvgpu.make_tiled_tma_atom_A(
a_op,
a_tensor2,
a_smem_layout,
self.mma_tiler,
tiled_mma,
self.cluster_layout_vmnk.shape,
)
tma_atom_b2, tma_tensor_b2 = cute.nvgpu.make_tiled_tma_atom_B(
b_op,
b_tensor2,
b_smem_layout,
self.mma_tiler,
tiled_mma,
self.cluster_layout_vmnk.shape,
)
tma_atom_sfa2, tma_tensor_sfa2 = cute.nvgpu.make_tiled_tma_atom_A(
sfa_op,
sfa_tensor2,
sfa_smem_layout,
self.mma_tiler,
tiled_mma,
self.cluster_layout_vmnk.shape,
internal_type=cutlass.Int16,
)
tma_atom_sfb2, tma_tensor_sfb2 = cute.nvgpu.make_tiled_tma_atom_B(
sfb_op,
sfb_tensor2,
sfb_smem_layout,
self.mma_tiler_sfb,
tiled_mma_sfb,
self.cluster_layout_sfb_vmnk.shape,
internal_type=cutlass.Int16,
)
tma_atom_a3, tma_tensor_a3 = cute.nvgpu.make_tiled_tma_atom_A(
a_op,
a_tensor3,
a_smem_layout,
self.mma_tiler,
tiled_mma,
self.cluster_layout_vmnk.shape,
)
tma_atom_b3, tma_tensor_b3 = cute.nvgpu.make_tiled_tma_atom_B(
b_op,
b_tensor3,
b_smem_layout,
self.mma_tiler,
tiled_mma,
self.cluster_layout_vmnk.shape,
)
tma_atom_sfa3, tma_tensor_sfa3 = cute.nvgpu.make_tiled_tma_atom_A(
sfa_op,
sfa_tensor3,
sfa_smem_layout,
self.mma_tiler,
tiled_mma,
self.cluster_layout_vmnk.shape,
internal_type=cutlass.Int16,
)
tma_atom_sfb3, tma_tensor_sfb3 = cute.nvgpu.make_tiled_tma_atom_B(
sfb_op,
sfb_tensor3,
sfb_smem_layout,
self.mma_tiler_sfb,
tiled_mma_sfb,
self.cluster_layout_sfb_vmnk.shape,
internal_type=cutlass.Int16,
)
tma_atom_a4, tma_tensor_a4 = cute.nvgpu.make_tiled_tma_atom_A(
a_op,
a_tensor4,
a_smem_layout,
self.mma_tiler,
tiled_mma,
self.cluster_layout_vmnk.shape,
)
tma_atom_b4, tma_tensor_b4 = cute.nvgpu.make_tiled_tma_atom_B(
b_op,
b_tensor4,
b_smem_layout,
self.mma_tiler,
tiled_mma,
self.cluster_layout_vmnk.shape,
)
tma_atom_sfa4, tma_tensor_sfa4 = cute.nvgpu.make_tiled_tma_atom_A(
sfa_op,
sfa_tensor4,
sfa_smem_layout,
self.mma_tiler,
tiled_mma,
self.cluster_layout_vmnk.shape,
internal_type=cutlass.Int16,
)
tma_atom_sfb4, tma_tensor_sfb4 = cute.nvgpu.make_tiled_tma_atom_B(
sfb_op,
sfb_tensor4,
sfb_smem_layout,
self.mma_tiler_sfb,
tiled_mma_sfb,
self.cluster_layout_sfb_vmnk.shape,
internal_type=cutlass.Int16,
)
tma_atom_a5, tma_tensor_a5 = cute.nvgpu.make_tiled_tma_atom_A(
a_op,
a_tensor5,
a_smem_layout,
self.mma_tiler,
tiled_mma,
self.cluster_layout_vmnk.shape,
)
tma_atom_b5, tma_tensor_b5 = cute.nvgpu.make_tiled_tma_atom_B(
b_op,
b_tensor5,
b_smem_layout,
self.mma_tiler,
tiled_mma,
self.cluster_layout_vmnk.shape,
)
tma_atom_sfa5, tma_tensor_sfa5 = cute.nvgpu.make_tiled_tma_atom_A(
sfa_op,
sfa_tensor5,
sfa_smem_layout,
self.mma_tiler,
tiled_mma,
self.cluster_layout_vmnk.shape,
internal_type=cutlass.Int16,
)
tma_atom_sfb5, tma_tensor_sfb5 = cute.nvgpu.make_tiled_tma_atom_B(
sfb_op,
sfb_tensor5,
sfb_smem_layout,
self.mma_tiler_sfb,
tiled_mma_sfb,
self.cluster_layout_sfb_vmnk.shape,
internal_type=cutlass.Int16,
)
tma_atom_a6, tma_tensor_a6 = cute.nvgpu.make_tiled_tma_atom_A(
a_op,
a_tensor6,
a_smem_layout,
self.mma_tiler,
tiled_mma,
self.cluster_layout_vmnk.shape,
)
tma_atom_b6, tma_tensor_b6 = cute.nvgpu.make_tiled_tma_atom_B(
b_op,
b_tensor6,
b_smem_layout,
self.mma_tiler,
tiled_mma,
self.cluster_layout_vmnk.shape,
)
tma_atom_sfa6, tma_tensor_sfa6 = cute.nvgpu.make_tiled_tma_atom_A(
sfa_op,
sfa_tensor6,
sfa_smem_layout,
self.mma_tiler,
tiled_mma,
self.cluster_layout_vmnk.shape,
internal_type=cutlass.Int16,
)
tma_atom_sfb6, tma_tensor_sfb6 = cute.nvgpu.make_tiled_tma_atom_B(
sfb_op,
sfb_tensor6,
sfb_smem_layout,
self.mma_tiler_sfb,
tiled_mma_sfb,
self.cluster_layout_sfb_vmnk.shape,
internal_type=cutlass.Int16,
)
tma_atom_a7, tma_tensor_a7 = cute.nvgpu.make_tiled_tma_atom_A(
a_op,
a_tensor7,
a_smem_layout,
self.mma_tiler,
tiled_mma,
self.cluster_layout_vmnk.shape,
)
tma_atom_b7, tma_tensor_b7 = cute.nvgpu.make_tiled_tma_atom_B(
b_op,
b_tensor7,
b_smem_layout,
self.mma_tiler,
tiled_mma,
self.cluster_layout_vmnk.shape,
)
tma_atom_sfa7, tma_tensor_sfa7 = cute.nvgpu.make_tiled_tma_atom_A(
sfa_op,
sfa_tensor7,
sfa_smem_layout,
self.mma_tiler,
tiled_mma,
self.cluster_layout_vmnk.shape,
internal_type=cutlass.Int16,
)
tma_atom_sfb7, tma_tensor_sfb7 = cute.nvgpu.make_tiled_tma_atom_B(
sfb_op,
sfb_tensor7,
sfb_smem_layout,
self.mma_tiler_sfb,
tiled_mma_sfb,
self.cluster_layout_sfb_vmnk.shape,
internal_type=cutlass.Int16,
)
if cutlass.const_expr(self.cta_tile_shape_mnk[1] == 192):
x0 = tma_tensor_sfb0.stride[0][1]
y0 = cute.ceil_div(tma_tensor_sfb0.shape[0][1], 4)
new_shape0 = ((tma_tensor_sfb0.shape[0][0], ((2, 2), y0)), tma_tensor_sfb0.shape[1], tma_tensor_sfb0.shape[2])
x0_times_3 = 3 * x0
new_stride0 = ((tma_tensor_sfb0.stride[0][0], ((x0, x0), x0_times_3)), tma_tensor_sfb0.stride[1], tma_tensor_sfb0.stride[2])
tma_tensor_sfb0 = cute.make_tensor(tma_tensor_sfb0.iterator, cute.make_layout(new_shape0, stride=new_stride0))
x1 = tma_tensor_sfb1.stride[0][1]
y1 = cute.ceil_div(tma_tensor_sfb1.shape[0][1], 4)
new_shape1 = ((tma_tensor_sfb1.shape[0][0], ((2, 2), y1)), tma_tensor_sfb1.shape[1], tma_tensor_sfb1.shape[2])
x1_times_3 = 3 * x1
new_stride1 = ((tma_tensor_sfb1.stride[0][0], ((x1, x1), x1_times_3)), tma_tensor_sfb1.stride[1], tma_tensor_sfb1.stride[2])
tma_tensor_sfb1 = cute.make_tensor(tma_tensor_sfb1.iterator, cute.make_layout(new_shape1, stride=new_stride1))
x2 = tma_tensor_sfb2.stride[0][1]
y2 = cute.ceil_div(tma_tensor_sfb2.shape[0][1], 4)
new_shape2 = ((tma_tensor_sfb2.shape[0][0], ((2, 2), y2)), tma_tensor_sfb2.shape[1], tma_tensor_sfb2.shape[2])
x2_times_3 = 3 * x2
new_stride2 = ((tma_tensor_sfb2.stride[0][0], ((x2, x2), x2_times_3)), tma_tensor_sfb2.stride[1], tma_tensor_sfb2.stride[2])
tma_tensor_sfb2 = cute.make_tensor(tma_tensor_sfb2.iterator, cute.make_layout(new_shape2, stride=new_stride2))
x3 = tma_tensor_sfb3.stride[0][1]
y3 = cute.ceil_div(tma_tensor_sfb3.shape[0][1], 4)
new_shape3 = ((tma_tensor_sfb3.shape[0][0], ((2, 2), y3)), tma_tensor_sfb3.shape[1], tma_tensor_sfb3.shape[2])
x3_times_3 = 3 * x3
new_stride3 = ((tma_tensor_sfb3.stride[0][0], ((x3, x3), x3_times_3)), tma_tensor_sfb3.stride[1], tma_tensor_sfb3.stride[2])
tma_tensor_sfb3 = cute.make_tensor(tma_tensor_sfb3.iterator, cute.make_layout(new_shape3, stride=new_stride3))
x4 = tma_tensor_sfb4.stride[0][1]
y4 = cute.ceil_div(tma_tensor_sfb4.shape[0][1], 4)
new_shape4 = ((tma_tensor_sfb4.shape[0][0], ((2, 2), y4)), tma_tensor_sfb4.shape[1], tma_tensor_sfb4.shape[2])
x4_times_3 = 3 * x4
new_stride4 = ((tma_tensor_sfb4.stride[0][0], ((x4, x4), x4_times_3)), tma_tensor_sfb4.stride[1], tma_tensor_sfb4.stride[2])
tma_tensor_sfb4 = cute.make_tensor(tma_tensor_sfb4.iterator, cute.make_layout(new_shape4, stride=new_stride4))
x5 = tma_tensor_sfb5.stride[0][1]
y5 = cute.ceil_div(tma_tensor_sfb5.shape[0][1], 4)
new_shape5 = ((tma_tensor_sfb5.shape[0][0], ((2, 2), y5)), tma_tensor_sfb5.shape[1], tma_tensor_sfb5.shape[2])
x5_times_3 = 3 * x5
new_stride5 = ((tma_tensor_sfb5.stride[0][0], ((x5, x5), x5_times_3)), tma_tensor_sfb5.stride[1], tma_tensor_sfb5.stride[2])
tma_tensor_sfb5 = cute.make_tensor(tma_tensor_sfb5.iterator, cute.make_layout(new_shape5, stride=new_stride5))
x6 = tma_tensor_sfb6.stride[0][1]
y6 = cute.ceil_div(tma_tensor_sfb6.shape[0][1], 4)
new_shape6 = ((tma_tensor_sfb6.shape[0][0], ((2, 2), y6)), tma_tensor_sfb6.shape[1], tma_tensor_sfb6.shape[2])
x6_times_3 = 3 * x6
new_stride6 = ((tma_tensor_sfb6.stride[0][0], ((x6, x6), x6_times_3)), tma_tensor_sfb6.stride[1], tma_tensor_sfb6.stride[2])
tma_tensor_sfb6 = cute.make_tensor(tma_tensor_sfb6.iterator, cute.make_layout(new_shape6, stride=new_stride6))
x7 = tma_tensor_sfb7.stride[0][1]
y7 = cute.ceil_div(tma_tensor_sfb7.shape[0][1], 4)
new_shape7 = ((tma_tensor_sfb7.shape[0][0], ((2, 2), y7)), tma_tensor_sfb7.shape[1], tma_tensor_sfb7.shape[2])
x7_times_3 = 3 * x7
new_stride7 = ((tma_tensor_sfb7.stride[0][0], ((x7, x7), x7_times_3)), tma_tensor_sfb7.stride[1], tma_tensor_sfb7.stride[2])
tma_tensor_sfb7 = cute.make_tensor(tma_tensor_sfb7.iterator, cute.make_layout(new_shape7, stride=new_stride7))
a_copy_size = cute.size_in_bytes(self.a_dtype, a_smem_layout)
b_copy_size = cute.size_in_bytes(self.b_dtype, b_smem_layout)
sfa_copy_size = cute.size_in_bytes(self.sf_dtype, sfa_smem_layout)
sfb_copy_size = cute.size_in_bytes(self.sf_dtype, sfb_smem_layout)
self.num_tma_load_bytes = (a_copy_size + b_copy_size + sfa_copy_size + sfb_copy_size) * atom_thr_size
epi_smem_layout = cute.slice_(self.c_smem_layout_staged, (None, None, 0))
tma_atom_c0, tma_tensor_c0 = cpasync.make_tiled_tma_atom(
cpasync.CopyBulkTensorTileS2GOp(),
c_tensor0,
epi_smem_layout,
self.epi_tile,
)
tma_atom_c1, tma_tensor_c1 = cpasync.make_tiled_tma_atom(
cpasync.CopyBulkTensorTileS2GOp(),
c_tensor1,
epi_smem_layout,
self.epi_tile,
)
tma_atom_c2, tma_tensor_c2 = cpasync.make_tiled_tma_atom(
cpasync.CopyBulkTensorTileS2GOp(),
c_tensor2,
epi_smem_layout,
self.epi_tile,
)
tma_atom_c3, tma_tensor_c3 = cpasync.make_tiled_tma_atom(
cpasync.CopyBulkTensorTileS2GOp(),
c_tensor3,
epi_smem_layout,
self.epi_tile,
)
tma_atom_c4, tma_tensor_c4 = cpasync.make_tiled_tma_atom(
cpasync.CopyBulkTensorTileS2GOp(),
c_tensor4,
epi_smem_layout,
self.epi_tile,
)
tma_atom_c5, tma_tensor_c5 = cpasync.make_tiled_tma_atom(
cpasync.CopyBulkTensorTileS2GOp(),
c_tensor5,
epi_smem_layout,
self.epi_tile,
)
tma_atom_c6, tma_tensor_c6 = cpasync.make_tiled_tma_atom(
cpasync.CopyBulkTensorTileS2GOp(),
c_tensor6,
epi_smem_layout,
self.epi_tile,
)
tma_atom_c7, tma_tensor_c7 = cpasync.make_tiled_tma_atom(
cpasync.CopyBulkTensorTileS2GOp(),
c_tensor7,
epi_smem_layout,
self.epi_tile,
)
tile_m = self.cta_tile_shape_mnk[0]
blk0 = cute.ceil_div(m0, tile_m) * atom_thr_size
blk1 = cute.ceil_div(m1, tile_m) * atom_thr_size
blk2 = cute.ceil_div(m2, tile_m) * atom_thr_size
blk3 = cute.ceil_div(m3, tile_m) * atom_thr_size
blk4 = cute.ceil_div(m4, tile_m) * atom_thr_size
blk5 = cute.ceil_div(m5, tile_m) * atom_thr_size
blk6 = cute.ceil_div(m6, tile_m) * atom_thr_size
blk7 = cute.ceil_div(m7, tile_m) * atom_thr_size
p0 = blk0
p1 = p0 + blk1
p2 = p1 + blk2
p3 = p2 + blk3
p4 = p3 + blk4
p5 = p4 + blk5
p6 = p5 + blk6
grid_m = p6 + blk7
grid_n = cute.ceil_div(n, self.cta_tile_shape_mnk[1])
grid = (grid_m, grid_n, l)
self.buffer_align_bytes = 1024
@cute.struct
class SharedStorage:
ab_full_mbar_ptr: cute.struct.MemRange[cutlass.Int64, self.num_ab_stage]
ab_empty_mbar_ptr: cute.struct.MemRange[cutlass.Int64, self.num_ab_stage]
acc_full_mbar_ptr: cute.struct.MemRange[cutlass.Int64, self.num_acc_stage]
acc_empty_mbar_ptr: cute.struct.MemRange[cutlass.Int64, self.num_acc_stage]
tmem_dealloc_mbar_ptr: cutlass.Int64
tmem_holding_buf: cutlass.Int32
sC: cute.struct.Align[
cute.struct.MemRange[
self.c_dtype,
cute.cosize(self.c_smem_layout_staged.outer),
],
self.buffer_align_bytes,
]
sA: cute.struct.Align[
cute.struct.MemRange[
self.a_dtype,
cute.cosize(self.a_smem_layout_staged.outer),
],
self.buffer_align_bytes,
]
sB: cute.struct.Align[
cute.struct.MemRange[
self.b_dtype,
cute.cosize(self.b_smem_layout_staged.outer),
],
self.buffer_align_bytes,
]
sSFA: cute.struct.Align[
cute.struct.MemRange[
self.sf_dtype,
cute.cosize(self.sfa_smem_layout_staged),
],
self.buffer_align_bytes,
]
sSFB: cute.struct.Align[
cute.struct.MemRange[
self.sf_dtype,
cute.cosize(self.sfb_smem_layout_staged),
],
self.buffer_align_bytes,
]
self.shared_storage = SharedStorage
self.kernel_group8(
tiled_mma,
tiled_mma_sfb,
tma_atom_a0,
tma_tensor_a0,
tma_atom_b0,
tma_tensor_b0,
tma_atom_sfa0,
tma_tensor_sfa0,
tma_atom_sfb0,
tma_tensor_sfb0,
tma_atom_c0,
tma_tensor_c0,
tma_atom_a1,
tma_tensor_a1,
tma_atom_b1,
tma_tensor_b1,
tma_atom_sfa1,
tma_tensor_sfa1,
tma_atom_sfb1,
tma_tensor_sfb1,
tma_atom_c1,
tma_tensor_c1,
tma_atom_a2,
tma_tensor_a2,
tma_atom_b2,
tma_tensor_b2,
tma_atom_sfa2,
tma_tensor_sfa2,
tma_atom_sfb2,
tma_tensor_sfb2,
tma_atom_c2,
tma_tensor_c2,
tma_atom_a3,
tma_tensor_a3,
tma_atom_b3,
tma_tensor_b3,
tma_atom_sfa3,
tma_tensor_sfa3,
tma_atom_sfb3,
tma_tensor_sfb3,
tma_atom_c3,
tma_tensor_c3,
tma_atom_a4,
tma_tensor_a4,
tma_atom_b4,
tma_tensor_b4,
tma_atom_sfa4,
tma_tensor_sfa4,
tma_atom_sfb4,
tma_tensor_sfb4,
tma_atom_c4,
tma_tensor_c4,
tma_atom_a5,
tma_tensor_a5,
tma_atom_b5,
tma_tensor_b5,
tma_atom_sfa5,
tma_tensor_sfa5,
tma_atom_sfb5,
tma_tensor_sfb5,
tma_atom_c5,
tma_tensor_c5,
tma_atom_a6,
tma_tensor_a6,
tma_atom_b6,
tma_tensor_b6,
tma_atom_sfa6,
tma_tensor_sfa6,
tma_atom_sfb6,
tma_tensor_sfb6,
tma_atom_c6,
tma_tensor_c6,
tma_atom_a7,
tma_tensor_a7,
tma_atom_b7,
tma_tensor_b7,
tma_atom_sfa7,
tma_tensor_sfa7,
tma_atom_sfb7,
tma_tensor_sfb7,
tma_atom_c7,
tma_tensor_c7,
self.cluster_layout_vmnk,
self.cluster_layout_sfb_vmnk,
self.a_smem_layout_staged,
self.b_smem_layout_staged,
self.sfa_smem_layout_staged,
self.sfb_smem_layout_staged,
self.c_smem_layout_staged,
self.epi_tile,
epilogue_op,
blk0,
blk1,
blk2,
blk3,
blk4,
blk5,
blk6,
).launch(
grid=grid,
block=[self.threads_per_cta, 1, 1],
cluster=(*self.cluster_shape_mn, 1),
min_blocks_per_mp=self.occupancy,
)
return
@cute.kernel
def kernel_group8(
self,
tiled_mma: cute.TiledMma,
tiled_mma_sfb: cute.TiledMma,
tma_atom_a0: cute.CopyAtom,
mA0_mkl: cute.Tensor,
tma_atom_b0: cute.CopyAtom,
mB0_nkl: cute.Tensor,
tma_atom_sfa0: cute.CopyAtom,
mSFA0_mkl: cute.Tensor,
tma_atom_sfb0: cute.CopyAtom,
mSFB0_nkl: cute.Tensor,
tma_atom_c0: cute.CopyAtom,
mC0_mnl: cute.Tensor,
tma_atom_a1: cute.CopyAtom,
mA1_mkl: cute.Tensor,
tma_atom_b1: cute.CopyAtom,
mB1_nkl: cute.Tensor,
tma_atom_sfa1: cute.CopyAtom,
mSFA1_mkl: cute.Tensor,
tma_atom_sfb1: cute.CopyAtom,
mSFB1_nkl: cute.Tensor,
tma_atom_c1: cute.CopyAtom,
mC1_mnl: cute.Tensor,
tma_atom_a2: cute.CopyAtom,
mA2_mkl: cute.Tensor,
tma_atom_b2: cute.CopyAtom,
mB2_nkl: cute.Tensor,
tma_atom_sfa2: cute.CopyAtom,
mSFA2_mkl: cute.Tensor,
tma_atom_sfb2: cute.CopyAtom,
mSFB2_nkl: cute.Tensor,
tma_atom_c2: cute.CopyAtom,
mC2_mnl: cute.Tensor,
tma_atom_a3: cute.CopyAtom,
mA3_mkl: cute.Tensor,
tma_atom_b3: cute.CopyAtom,
mB3_nkl: cute.Tensor,
tma_atom_sfa3: cute.CopyAtom,
mSFA3_mkl: cute.Tensor,
tma_atom_sfb3: cute.CopyAtom,
mSFB3_nkl: cute.Tensor,
tma_atom_c3: cute.CopyAtom,
mC3_mnl: cute.Tensor,
tma_atom_a4: cute.CopyAtom,
mA4_mkl: cute.Tensor,
tma_atom_b4: cute.CopyAtom,
mB4_nkl: cute.Tensor,
tma_atom_sfa4: cute.CopyAtom,
mSFA4_mkl: cute.Tensor,
tma_atom_sfb4: cute.CopyAtom,
mSFB4_nkl: cute.Tensor,
tma_atom_c4: cute.CopyAtom,
mC4_mnl: cute.Tensor,
tma_atom_a5: cute.CopyAtom,
mA5_mkl: cute.Tensor,
tma_atom_b5: cute.CopyAtom,
mB5_nkl: cute.Tensor,
tma_atom_sfa5: cute.CopyAtom,
mSFA5_mkl: cute.Tensor,
tma_atom_sfb5: cute.CopyAtom,
mSFB5_nkl: cute.Tensor,
tma_atom_c5: cute.CopyAtom,
mC5_mnl: cute.Tensor,
tma_atom_a6: cute.CopyAtom,
mA6_mkl: cute.Tensor,
tma_atom_b6: cute.CopyAtom,
mB6_nkl: cute.Tensor,
tma_atom_sfa6: cute.CopyAtom,
mSFA6_mkl: cute.Tensor,
tma_atom_sfb6: cute.CopyAtom,
mSFB6_nkl: cute.Tensor,
tma_atom_c6: cute.CopyAtom,
mC6_mnl: cute.Tensor,
tma_atom_a7: cute.CopyAtom,
mA7_mkl: cute.Tensor,
tma_atom_b7: cute.CopyAtom,
mB7_nkl: cute.Tensor,
tma_atom_sfa7: cute.CopyAtom,
mSFA7_mkl: cute.Tensor,
tma_atom_sfb7: cute.CopyAtom,
mSFB7_nkl: cute.Tensor,
tma_atom_c7: cute.CopyAtom,
mC7_mnl: cute.Tensor,
cluster_layout_vmnk: cute.Layout,
cluster_layout_sfb_vmnk: cute.Layout,
a_smem_layout_staged: cute.ComposedLayout,
b_smem_layout_staged: cute.ComposedLayout,
sfa_smem_layout_staged: cute.Layout,
sfb_smem_layout_staged: cute.Layout,
c_smem_layout_staged: Union[cute.Layout, cute.ComposedLayout],
epi_tile: cute.Tile,
epilogue_op: cutlass.Constexpr,
blk0: cutlass.Int32,
blk1: cutlass.Int32,
blk2: cutlass.Int32,
blk3: cutlass.Int32,
blk4: cutlass.Int32,
blk5: cutlass.Int32,
blk6: cutlass.Int32,
):
warp_idx = cute.arch.warp_idx()
warp_idx = cute.arch.make_warp_uniform(warp_idx)
use_2cta_instrs = cute.size(tiled_mma.thr_id.shape) == 2
bidx, bidy, bidz = cute.arch.block_idx()
l = mC0_mnl.shape[2]
group_idx = bidz // l
group_idx = cute.arch.make_warp_uniform(group_idx)
l_idx = bidz - group_idx * l
l_idx = cute.arch.make_warp_uniform(l_idx)
tma_atom_a = tma_atom_a0
mA_mkl = mA0_mkl
tma_atom_b = tma_atom_b0
mB_nkl = mB0_nkl
tma_atom_sfa = tma_atom_sfa0
mSFA_mkl = mSFA0_mkl
tma_atom_sfb = tma_atom_sfb0
mSFB_nkl = mSFB0_nkl
tma_atom_c = tma_atom_c0
mC_mnl = mC0_mnl
if group_idx == 1:
tma_atom_a = tma_atom_a1
mA_mkl = mA1_mkl
tma_atom_b = tma_atom_b1
mB_nkl = mB1_nkl
tma_atom_sfa = tma_atom_sfa1
mSFA_mkl = mSFA1_mkl
tma_atom_sfb = tma_atom_sfb1
mSFB_nkl = mSFB1_nkl
tma_atom_c = tma_atom_c1
mC_mnl = mC1_mnl
elif group_idx == 2:
tma_atom_a = tma_atom_a2
mA_mkl = mA2_mkl
tma_atom_b = tma_atom_b2
mB_nkl = mB2_nkl
tma_atom_sfa = tma_atom_sfa2
mSFA_mkl = mSFA2_mkl
tma_atom_sfb = tma_atom_sfb2
mSFB_nkl = mSFB2_nkl
tma_atom_c = tma_atom_c2
mC_mnl = mC2_mnl
elif group_idx == 3:
tma_atom_a = tma_atom_a3
mA_mkl = mA3_mkl
tma_atom_b = tma_atom_b3
mB_nkl = mB3_nkl
tma_atom_sfa = tma_atom_sfa3
mSFA_mkl = mSFA3_mkl
tma_atom_sfb = tma_atom_sfb3
mSFB_nkl = mSFB3_nkl
tma_atom_c = tma_atom_c3
mC_mnl = mC3_mnl
elif group_idx == 4:
tma_atom_a = tma_atom_a4
mA_mkl = mA4_mkl
tma_atom_b = tma_atom_b4
mB_nkl = mB4_nkl
tma_atom_sfa = tma_atom_sfa4
mSFA_mkl = mSFA4_mkl
tma_atom_sfb = tma_atom_sfb4
mSFB_nkl = mSFB4_nkl
tma_atom_c = tma_atom_c4
mC_mnl = mC4_mnl
elif group_idx == 5:
tma_atom_a = tma_atom_a5
mA_mkl = mA5_mkl
tma_atom_b = tma_atom_b5
mB_nkl = mB5_nkl
tma_atom_sfa = tma_atom_sfa5
mSFA_mkl = mSFA5_mkl
tma_atom_sfb = tma_atom_sfb5
mSFB_nkl = mSFB5_nkl
tma_atom_c = tma_atom_c5
mC_mnl = mC5_mnl
elif group_idx == 6:
tma_atom_a = tma_atom_a6
mA_mkl = mA6_mkl
tma_atom_b = tma_atom_b6
mB_nkl = mB6_nkl
tma_atom_sfa = tma_atom_sfa6
mSFA_mkl = mSFA6_mkl
tma_atom_sfb = tma_atom_sfb6
mSFB_nkl = mSFB6_nkl
tma_atom_c = tma_atom_c6
mC_mnl = mC6_mnl
elif group_idx == 7:
tma_atom_a = tma_atom_a7
mA_mkl = mA7_mkl
tma_atom_b = tma_atom_b7
mB_nkl = mB7_nkl
tma_atom_sfa = tma_atom_sfa7
mSFA_mkl = mSFA7_mkl
tma_atom_sfb = tma_atom_sfb7
mSFB_nkl = mSFB7_nkl
tma_atom_c = tma_atom_c7
mC_mnl = mC7_mnl
tile_m_coord = bidx // cute.size(tiled_mma.thr_id.shape)
tile_active = tile_m_coord * self.cta_tile_shape_mnk[0] < mC_mnl.shape[0]
tile_active = cute.arch.make_warp_uniform(tile_active)
mma_tile_coord_v = bidx % cute.size(tiled_mma.thr_id.shape)
is_leader_cta = mma_tile_coord_v == 0
cta_rank_in_cluster = cute.arch.make_warp_uniform(
cute.arch.block_idx_in_cluster()
)
block_in_cluster_coord_vmnk = cluster_layout_vmnk.get_flat_coord(
cta_rank_in_cluster
)
block_in_cluster_coord_sfb_vmnk = cluster_layout_sfb_vmnk.get_flat_coord(
cta_rank_in_cluster
)
tidx, _, _ = cute.arch.thread_idx()
smem = utils.SmemAllocator()
storage = smem.allocate(self.shared_storage)
ab_pipeline_producer_group = pipeline.CooperativeGroup(pipeline.Agent.Thread)
num_tma_producer = self.num_mcast_ctas_a + self.num_mcast_ctas_b - 1
ab_pipeline_consumer_group = pipeline.CooperativeGroup(
pipeline.Agent.Thread, num_tma_producer
)
try:
ab_pipeline = pipeline.PipelineTmaUmma.create(
barrier_storage=storage.ab_full_mbar_ptr.data_ptr(),
num_stages=self.num_ab_stage,
producer_group=ab_pipeline_producer_group,
consumer_group=ab_pipeline_consumer_group,
tx_count=self.num_tma_load_bytes,
cta_layout_vmnk=cluster_layout_vmnk,
defer_sync=True,
)
except TypeError:
ab_pipeline = pipeline.PipelineTmaUmma.create(
barrier_storage=storage.ab_full_mbar_ptr.data_ptr(),
num_stages=self.num_ab_stage,
producer_group=ab_pipeline_producer_group,
consumer_group=ab_pipeline_consumer_group,
tx_count=self.num_tma_load_bytes,
cta_layout_vmnk=cluster_layout_vmnk,
)
acc_pipeline_producer_group = pipeline.CooperativeGroup(pipeline.Agent.Thread)
num_acc_consumer_threads = len(self.epilog_warp_id) * (
2 if use_2cta_instrs else 1
)
acc_pipeline_consumer_group = pipeline.CooperativeGroup(
pipeline.Agent.Thread, num_acc_consumer_threads
)
try:
acc_pipeline = pipeline.PipelineUmmaAsync.create(
barrier_storage=storage.acc_full_mbar_ptr.data_ptr(),
num_stages=self.num_acc_stage,
producer_group=acc_pipeline_producer_group,
consumer_group=acc_pipeline_consumer_group,
cta_layout_vmnk=cluster_layout_vmnk,
defer_sync=True,
)
except TypeError:
acc_pipeline = pipeline.PipelineUmmaAsync.create(
barrier_storage=storage.acc_full_mbar_ptr.data_ptr(),
num_stages=self.num_acc_stage,
producer_group=acc_pipeline_producer_group,
consumer_group=acc_pipeline_consumer_group,
cta_layout_vmnk=cluster_layout_vmnk,
)
tmem = utils.TmemAllocator(
storage.tmem_holding_buf,
barrier_for_retrieve=self.tmem_alloc_barrier,
allocator_warp_id=self.epilog_warp_id[0],
is_two_cta=use_2cta_instrs,
two_cta_tmem_dealloc_mbar_ptr=storage.tmem_dealloc_mbar_ptr,
)
sC = storage.sC.get_tensor(
c_smem_layout_staged.outer, swizzle=c_smem_layout_staged.inner
)
sA = storage.sA.get_tensor(
a_smem_layout_staged.outer, swizzle=a_smem_layout_staged.inner
)
sB = storage.sB.get_tensor(
b_smem_layout_staged.outer, swizzle=b_smem_layout_staged.inner
)
sSFA = storage.sSFA.get_tensor(sfa_smem_layout_staged)
sSFB = storage.sSFB.get_tensor(sfb_smem_layout_staged)
a_full_mcast_mask = None
b_full_mcast_mask = None
sfa_full_mcast_mask = None
sfb_full_mcast_mask = None
if cutlass.const_expr(self.is_a_mcast or self.is_b_mcast or use_2cta_instrs):
a_full_mcast_mask = cpasync.create_tma_multicast_mask(
cluster_layout_vmnk, block_in_cluster_coord_vmnk, mcast_mode=2
)
b_full_mcast_mask = cpasync.create_tma_multicast_mask(
cluster_layout_vmnk, block_in_cluster_coord_vmnk, mcast_mode=1
)
sfa_full_mcast_mask = cpasync.create_tma_multicast_mask(
cluster_layout_vmnk, block_in_cluster_coord_vmnk, mcast_mode=2
)
sfb_full_mcast_mask = cpasync.create_tma_multicast_mask(
cluster_layout_sfb_vmnk, block_in_cluster_coord_sfb_vmnk, mcast_mode=1
)
tma_cache_policy_a = cutlass.Int64(cutlass.Int64(self.tma_cache_policy_a).ir_value())
tma_cache_policy_b = cutlass.Int64(cutlass.Int64(self.tma_cache_policy_b).ir_value())
tma_cache_policy_sfa = cutlass.Int64(cutlass.Int64(self.tma_cache_policy_sfa).ir_value())
tma_cache_policy_sfb = cutlass.Int64(cutlass.Int64(self.tma_cache_policy_sfb).ir_value())
tma_cache_policy_c = cutlass.Int64(cutlass.Int64(self.tma_cache_policy_c).ir_value())
gA_mkl = cute.local_tile(
mA_mkl, cute.slice_(self.mma_tiler, (None, 0, None)), (None, None, None)
)
gB_nkl = cute.local_tile(
mB_nkl, cute.slice_(self.mma_tiler, (0, None, None)), (None, None, None)
)
gSFA_mkl = cute.local_tile(
mSFA_mkl, cute.slice_(self.mma_tiler, (None, 0, None)), (None, None, None)
)
gSFB_nkl = cute.local_tile(
mSFB_nkl,
cute.slice_(self.mma_tiler_sfb, (0, None, None)),
(None, None, None),
)
gC_mnl = cute.local_tile(
mC_mnl, cute.slice_(self.mma_tiler, (None, None, 0)), (None, None, None)
)
k_tile_cnt = cute.size(gA_mkl, mode=[3])
thr_mma = tiled_mma.get_slice(mma_tile_coord_v)
thr_mma_sfb = tiled_mma_sfb.get_slice(mma_tile_coord_v)
tCgA = thr_mma.partition_A(gA_mkl)
tCgB = thr_mma.partition_B(gB_nkl)
tCgSFA = thr_mma.partition_A(gSFA_mkl)
tCgSFB = thr_mma_sfb.partition_B(gSFB_nkl)
tCgC = thr_mma.partition_C(gC_mnl)
a_cta_layout = cute.make_layout(
cute.slice_(cluster_layout_vmnk, (0, 0, None, 0)).shape
)
tAsA, tAgA = cpasync.tma_partition(
tma_atom_a,
block_in_cluster_coord_vmnk[2],
a_cta_layout,
cute.group_modes(sA, 0, 3),
cute.group_modes(tCgA, 0, 3),
)
b_cta_layout = cute.make_layout(
cute.slice_(cluster_layout_vmnk, (0, None, 0, 0)).shape
)
tBsB, tBgB = cpasync.tma_partition(
tma_atom_b,
block_in_cluster_coord_vmnk[1],
b_cta_layout,
cute.group_modes(sB, 0, 3),
cute.group_modes(tCgB, 0, 3),
)
sfa_cta_layout = a_cta_layout
tAsSFA, tAgSFA = cute.nvgpu.cpasync.tma_partition(
tma_atom_sfa,
block_in_cluster_coord_vmnk[2],
sfa_cta_layout,
cute.group_modes(sSFA, 0, 3),
cute.group_modes(tCgSFA, 0, 3),
)
tAsSFA = cute.filter_zeros(tAsSFA)
tAgSFA = cute.filter_zeros(tAgSFA)
sfb_cta_layout = cute.make_layout(
cute.slice_(cluster_layout_sfb_vmnk, (0, None, 0, 0)).shape
)
tBsSFB, tBgSFB = cute.nvgpu.cpasync.tma_partition(
tma_atom_sfb,
block_in_cluster_coord_sfb_vmnk[1],
sfb_cta_layout,
cute.group_modes(sSFB, 0, 3),
cute.group_modes(tCgSFB, 0, 3),
)
tBsSFB = cute.filter_zeros(tBsSFB)
tBgSFB = cute.filter_zeros(tBgSFB)
tCrA = tiled_mma.make_fragment_A(sA)
tCrB = tiled_mma.make_fragment_B(sB)
acc_shape = tiled_mma.partition_shape_C(self.mma_tiler[:2])
if cutlass.const_expr(self.overlapping_accum):
num_acc_stage_overlapped = 2
tCtAcc_fake = tiled_mma.make_fragment_C(
cute.append(acc_shape, num_acc_stage_overlapped)
)
tCtAcc_fake = cute.make_tensor(
tCtAcc_fake.iterator,
cute.make_layout(
tCtAcc_fake.shape,
stride = (
tCtAcc_fake.stride[0],
tCtAcc_fake.stride[1],
tCtAcc_fake.stride[2],
(256 - self.num_sf_tmem_cols) * tCtAcc_fake.stride[0][1]
)
)
)
else:
tCtAcc_fake = tiled_mma.make_fragment_C(
cute.append(acc_shape, self.num_acc_stage)
)
if warp_idx == self.tma_warp_id:
ab_producer_state = pipeline.make_pipeline_state(
pipeline.PipelineUserType.Producer, self.num_ab_stage
)
mma_tile_coord_mnl = (
bidx // cute.size(tiled_mma.thr_id.shape),
bidy,
l_idx,
)
tAgA_slice = tAgA[(None, mma_tile_coord_mnl[0], None, mma_tile_coord_mnl[2])]
tBgB_slice = tBgB[(None, mma_tile_coord_mnl[1], None, mma_tile_coord_mnl[2])]
tAgSFA_slice = tAgSFA[(None, mma_tile_coord_mnl[0], None, mma_tile_coord_mnl[2])]
slice_n = mma_tile_coord_mnl[1]
if cutlass.const_expr(self.cta_tile_shape_mnk[1] == 64):
slice_n = mma_tile_coord_mnl[1] // 2
tBgSFB_slice = tBgSFB[(None, slice_n, None, mma_tile_coord_mnl[2])]
ab_producer_state.reset_count()
peek_ab_empty_status = cutlass.Boolean(1)
if ab_producer_state.count < k_tile_cnt:
peek_ab_empty_status = ab_pipeline.producer_try_acquire(ab_producer_state)
for k_tile in cutlass.range(0, k_tile_cnt, 1, unroll=1):
ab_pipeline.producer_acquire(ab_producer_state, peek_ab_empty_status)
try:
cute.copy(
tma_atom_a,
tAgA_slice[(None, ab_producer_state.count)],
tAsA[(None, ab_producer_state.index)],
tma_bar_ptr=ab_pipeline.producer_get_barrier(ab_producer_state),
mcast_mask=a_full_mcast_mask,
cache_policy=tma_cache_policy_a,
)
except TypeError:
cute.copy(
tma_atom_a,
tAgA_slice[(None, ab_producer_state.count)],
tAsA[(None, ab_producer_state.index)],
tma_bar_ptr=ab_pipeline.producer_get_barrier(ab_producer_state),
mcast_mask=a_full_mcast_mask,
)
try:
cute.copy(
tma_atom_b,
tBgB_slice[(None, ab_producer_state.count)],
tBsB[(None, ab_producer_state.index)],
tma_bar_ptr=ab_pipeline.producer_get_barrier(ab_producer_state),
mcast_mask=b_full_mcast_mask,
cache_policy=tma_cache_policy_b,
)
except TypeError:
cute.copy(
tma_atom_b,
tBgB_slice[(None, ab_producer_state.count)],
tBsB[(None, ab_producer_state.index)],
tma_bar_ptr=ab_pipeline.producer_get_barrier(ab_producer_state),
mcast_mask=b_full_mcast_mask,
)
try:
cute.copy(
tma_atom_sfa,
tAgSFA_slice[(None, ab_producer_state.count)],
tAsSFA[(None, ab_producer_state.index)],
tma_bar_ptr=ab_pipeline.producer_get_barrier(ab_producer_state),
mcast_mask=sfa_full_mcast_mask,
cache_policy=tma_cache_policy_sfa,
)
except TypeError:
cute.copy(
tma_atom_sfa,
tAgSFA_slice[(None, ab_producer_state.count)],
tAsSFA[(None, ab_producer_state.index)],
tma_bar_ptr=ab_pipeline.producer_get_barrier(ab_producer_state),
mcast_mask=sfa_full_mcast_mask,
)
try:
cute.copy(
tma_atom_sfb,
tBgSFB_slice[(None, ab_producer_state.count)],
tBsSFB[(None, ab_producer_state.index)],
tma_bar_ptr=ab_pipeline.producer_get_barrier(ab_producer_state),
mcast_mask=sfb_full_mcast_mask,
cache_policy=tma_cache_policy_sfb,
)
except TypeError:
cute.copy(
tma_atom_sfb,
tBgSFB_slice[(None, ab_producer_state.count)],
tBsSFB[(None, ab_producer_state.index)],
tma_bar_ptr=ab_pipeline.producer_get_barrier(ab_producer_state),
mcast_mask=sfb_full_mcast_mask,
)
ab_producer_state.advance()
peek_ab_empty_status = cutlass.Boolean(1)
if ab_producer_state.count < k_tile_cnt:
peek_ab_empty_status = ab_pipeline.producer_try_acquire(
ab_producer_state
)
ab_pipeline.producer_tail(ab_producer_state)
if warp_idx == self.mma_warp_id:
tmem.wait_for_alloc()
acc_tmem_ptr = tmem.retrieve_ptr(self.acc_dtype)
tCtAcc_base = cute.make_tensor(acc_tmem_ptr, tCtAcc_fake.layout)
sfa_tmem_ptr = cute.recast_ptr(
acc_tmem_ptr + self.num_accumulator_tmem_cols,
dtype=self.sf_dtype,
)
tCtSFA_layout = blockscaled_utils.make_tmem_layout_sfa(
tiled_mma,
self.mma_tiler,
self.sf_vec_size,
cute.slice_(sfa_smem_layout_staged, (None, None, None, 0)),
)
tCtSFA = cute.make_tensor(sfa_tmem_ptr, tCtSFA_layout)
sfb_tmem_ptr = cute.recast_ptr(
acc_tmem_ptr + self.num_accumulator_tmem_cols + self.num_sfa_tmem_cols,
dtype=self.sf_dtype,
)
tCtSFB_layout = blockscaled_utils.make_tmem_layout_sfb(
tiled_mma,
self.mma_tiler,
self.sf_vec_size,
cute.slice_(sfb_smem_layout_staged, (None, None, None, 0)),
)
tCtSFB = cute.make_tensor(sfb_tmem_ptr, tCtSFB_layout)
(
tiled_copy_s2t_sfa,
tCsSFA_compact_s2t,
tCtSFA_compact_s2t,
) = self.mainloop_s2t_copy_and_partition(sSFA, tCtSFA)
(
tiled_copy_s2t_sfb,
tCsSFB_compact_s2t,
tCtSFB_compact_s2t,
) = self.mainloop_s2t_copy_and_partition(sSFB, tCtSFB)
ab_consumer_state = pipeline.make_pipeline_state(
pipeline.PipelineUserType.Consumer, self.num_ab_stage
)
acc_producer_state = pipeline.make_pipeline_state(
pipeline.PipelineUserType.Producer, self.num_acc_stage
)
mma_tile_coord_mnl = (
bidx // cute.size(tiled_mma.thr_id.shape),
bidy,
l_idx,
)
if cutlass.const_expr(self.overlapping_accum):
acc_stage_index = acc_producer_state.phase ^ 1
else:
acc_stage_index = acc_producer_state.index
tCtAcc = tCtAcc_base[(None, None, None, acc_stage_index)]
ab_consumer_state.reset_count()
peek_ab_full_status = cutlass.Boolean(1)
if ab_consumer_state.count < k_tile_cnt and is_leader_cta:
peek_ab_full_status = ab_pipeline.consumer_try_wait(ab_consumer_state)
if is_leader_cta:
acc_pipeline.producer_acquire(acc_producer_state)
tCtSFB_mma = tCtSFB
if cutlass.const_expr(self.cta_tile_shape_mnk[1] == 192):
offset = (
cutlass.Int32(2)
if mma_tile_coord_mnl[1] % 2 == 1
else cutlass.Int32(0)
)
shifted_ptr = cute.recast_ptr(
acc_tmem_ptr
+ self.num_accumulator_tmem_cols
+ self.num_sfa_tmem_cols
+ offset,
dtype=self.sf_dtype,
)
tCtSFB_mma = cute.make_tensor(shifted_ptr, tCtSFB_layout)
elif cutlass.const_expr(self.cta_tile_shape_mnk[1] == 64):
offset = cutlass.Int32((mma_tile_coord_mnl[1] % 2) * 2)
shifted_ptr = cute.recast_ptr(
acc_tmem_ptr
+ self.num_accumulator_tmem_cols
+ self.num_sfa_tmem_cols
+ offset,
dtype=self.sf_dtype,
)
tCtSFB_mma = cute.make_tensor(shifted_ptr, tCtSFB_layout)
tiled_mma.set(tcgen05.Field.ACCUMULATE, False)
for k_tile in range(k_tile_cnt):
if is_leader_cta:
ab_pipeline.consumer_wait(ab_consumer_state, peek_ab_full_status)
s2t_stage_coord = (None, None, None, None, ab_consumer_state.index)
tCsSFA_compact_s2t_staged = tCsSFA_compact_s2t[s2t_stage_coord]
tCsSFB_compact_s2t_staged = tCsSFB_compact_s2t[s2t_stage_coord]
cute.copy(
tiled_copy_s2t_sfa,
tCsSFA_compact_s2t_staged,
tCtSFA_compact_s2t,
)
cute.copy(
tiled_copy_s2t_sfb,
tCsSFB_compact_s2t_staged,
tCtSFB_compact_s2t,
)
num_kblocks = cute.size(tCrA, mode=[2])
for kblock_idx in cutlass.range(num_kblocks, unroll_full=True):
kblock_coord = (None, None, kblock_idx, ab_consumer_state.index)
sf_kblock_coord = (None, None, kblock_idx)
tiled_mma.set(tcgen05.Field.SFA, tCtSFA[sf_kblock_coord].iterator)
tiled_mma.set(
tcgen05.Field.SFB, tCtSFB_mma[sf_kblock_coord].iterator
)
cute.gemm(
tiled_mma,
tCtAcc,
tCrA[kblock_coord],
tCrB[kblock_coord],
tCtAcc,
)
tiled_mma.set(tcgen05.Field.ACCUMULATE, True)
ab_pipeline.consumer_release(ab_consumer_state)
ab_consumer_state.advance()
peek_ab_full_status = cutlass.Boolean(1)
if ab_consumer_state.count < k_tile_cnt:
if is_leader_cta:
peek_ab_full_status = ab_pipeline.consumer_try_wait(ab_consumer_state)
if is_leader_cta:
acc_pipeline.producer_commit(acc_producer_state)
acc_producer_state.advance()
acc_pipeline.producer_tail(acc_producer_state)
if warp_idx < self.mma_warp_id:
tmem.allocate(self.num_tmem_alloc_cols)
tmem.wait_for_alloc()
acc_tmem_ptr = tmem.retrieve_ptr(self.acc_dtype)
tCtAcc_base = cute.make_tensor(acc_tmem_ptr, tCtAcc_fake.layout)
epi_tidx = tidx
(
tiled_copy_t2r,
tTR_tAcc_base,
tTR_rAcc,
) = self.epilog_tmem_copy_and_partition(
epi_tidx, tCtAcc_base, tCgC, epi_tile, use_2cta_instrs
)
tTR_rC = cute.make_rmem_tensor(tTR_rAcc.shape, self.c_dtype)
tiled_copy_r2s, tRS_rC, tRS_sC = self.epilog_smem_copy_and_partition(
tiled_copy_t2r, tTR_rC, epi_tidx, sC
)
(
tma_atom_c,
bSG_sC,
bSG_gC_partitioned,
) = self.epilog_gmem_copy_and_partition(
epi_tidx, tma_atom_c, tCgC, epi_tile, sC
)
acc_consumer_state = pipeline.make_pipeline_state(
pipeline.PipelineUserType.Consumer, self.num_acc_stage
)
c_producer_group = pipeline.CooperativeGroup(
pipeline.Agent.Thread,
32 * len(self.epilog_warp_id),
)
c_pipeline = pipeline.PipelineTmaStore.create(
num_stages=self.num_c_stage,
producer_group=c_producer_group,
)
mma_tile_coord_mnl = (
bidx // cute.size(tiled_mma.thr_id.shape),
bidy,
l_idx,
)
bSG_gC = bSG_gC_partitioned[(None, None, None, *mma_tile_coord_mnl)]
if cutlass.const_expr(self.overlapping_accum):
acc_stage_index = acc_consumer_state.phase
reverse_subtile = (
cutlass.Boolean(True)
if acc_stage_index == 0
else cutlass.Boolean(False)
)
else:
acc_stage_index = acc_consumer_state.index
tTR_tAcc = tTR_tAcc_base[(None, None, None, None, None, acc_stage_index)]
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))
subtile_cnt = cute.size(tTR_tAcc.shape, mode=[3])
num_prev_subtiles = cutlass.Int32(0)
for subtile_idx in cutlass.range(subtile_cnt):
real_subtile_idx = subtile_idx
if cutlass.const_expr(self.overlapping_accum):
if reverse_subtile:
real_subtile_idx = (
self.cta_tile_shape_mnk[1] // self.epi_tile_n - 1 - subtile_idx
)
tTR_tAcc_mn = tTR_tAcc[(None, None, None, real_subtile_idx)]
cute.copy(tiled_copy_t2r, tTR_tAcc_mn, tTR_rAcc)
if cutlass.const_expr(self.overlapping_accum):
if subtile_idx == self.iter_acc_early_release_in_epilogue:
cute.arch.fence_view_async_tmem_load()
with cute.arch.elect_one():
acc_pipeline.consumer_release(acc_consumer_state)
acc_consumer_state.advance()
acc_vec = tiled_copy_r2s.retile(tTR_rAcc).load()
acc_vec = epilogue_op(acc_vec.to(self.c_dtype))
tRS_rC.store(acc_vec)
c_buffer = (num_prev_subtiles + real_subtile_idx) % self.num_c_stage
cute.copy(
tiled_copy_r2s,
tRS_rC,
tRS_sC[(None, None, None, c_buffer)],
)
cute.arch.fence_proxy(
cute.arch.ProxyKind.async_shared,
space=cute.arch.SharedSpace.shared_cta,
)
self.epilog_sync_barrier.arrive_and_wait()
if warp_idx == self.epilog_warp_id[0]:
try:
cute.copy(
tma_atom_c,
bSG_sC[(None, c_buffer)],
bSG_gC[(None, real_subtile_idx)],
cache_policy=tma_cache_policy_c,
)
except TypeError:
cute.copy(
tma_atom_c,
bSG_sC[(None, c_buffer)],
bSG_gC[(None, real_subtile_idx)],
)
c_pipeline.producer_commit()
c_pipeline.producer_acquire()
self.epilog_sync_barrier.arrive_and_wait()
if cutlass.const_expr(not self.overlapping_accum):
with cute.arch.elect_one():
acc_pipeline.consumer_release(acc_consumer_state)
acc_consumer_state.advance()
tmem.relinquish_alloc_permit()
self.epilog_sync_barrier.arrive_and_wait()
tmem.free(acc_tmem_ptr)
c_pipeline.producer_tail()
class Sm100BlockScaledPersistentDenseGemmKernelGroupG8V2(
Sm100BlockScaledPersistentDenseGemmKernelGroupG8
):
@cute.jit
def _kernel_group8_one(
self,
tiled_mma: cute.TiledMma,
tiled_mma_sfb: cute.TiledMma,
tma_atom_a: cute.CopyAtom,
mA_mkl: cute.Tensor,
tma_atom_b: cute.CopyAtom,
mB_nkl: cute.Tensor,
tma_atom_sfa: cute.CopyAtom,
mSFA_mkl: cute.Tensor,
tma_atom_sfb: cute.CopyAtom,
mSFB_nkl: cute.Tensor,
tma_atom_c: cute.CopyAtom,
mC_mnl: cute.Tensor,
cluster_layout_vmnk: cute.Layout,
cluster_layout_sfb_vmnk: cute.Layout,
a_smem_layout_staged: cute.ComposedLayout,
b_smem_layout_staged: cute.ComposedLayout,
sfa_smem_layout_staged: cute.Layout,
sfb_smem_layout_staged: cute.Layout,
c_smem_layout_staged: Union[cute.Layout, cute.ComposedLayout],
epi_tile: cute.Tile,
epilogue_op: cutlass.Constexpr,
tile_m_coord: cutlass.Int32,
l_idx: cutlass.Int32,
):
warp_idx = cute.arch.warp_idx()
warp_idx = cute.arch.make_warp_uniform(warp_idx)
use_2cta_instrs = cute.size(tiled_mma.thr_id.shape) == 2
bidx, bidy, _bidz = cute.arch.block_idx()
tile_m_coord = cute.arch.make_warp_uniform(tile_m_coord)
mma_tile_coord_v = bidx % cute.size(tiled_mma.thr_id.shape)
is_leader_cta = mma_tile_coord_v == 0
cta_rank_in_cluster = cute.arch.make_warp_uniform(
cute.arch.block_idx_in_cluster()
)
block_in_cluster_coord_vmnk = cluster_layout_vmnk.get_flat_coord(
cta_rank_in_cluster
)
block_in_cluster_coord_sfb_vmnk = cluster_layout_sfb_vmnk.get_flat_coord(
cta_rank_in_cluster
)
tidx, _, _ = cute.arch.thread_idx()
smem = utils.SmemAllocator()
storage = smem.allocate(self.shared_storage)
ab_pipeline_producer_group = pipeline.CooperativeGroup(pipeline.Agent.Thread)
num_tma_producer = self.num_mcast_ctas_a + self.num_mcast_ctas_b - 1
ab_pipeline_consumer_group = pipeline.CooperativeGroup(
pipeline.Agent.Thread, num_tma_producer
)
try:
ab_pipeline = pipeline.PipelineTmaUmma.create(
barrier_storage=storage.ab_full_mbar_ptr.data_ptr(),
num_stages=self.num_ab_stage,
producer_group=ab_pipeline_producer_group,
consumer_group=ab_pipeline_consumer_group,
tx_count=self.num_tma_load_bytes,
cta_layout_vmnk=cluster_layout_vmnk,
defer_sync=True,
)
except TypeError:
ab_pipeline = pipeline.PipelineTmaUmma.create(
barrier_storage=storage.ab_full_mbar_ptr.data_ptr(),
num_stages=self.num_ab_stage,
producer_group=ab_pipeline_producer_group,
consumer_group=ab_pipeline_consumer_group,
tx_count=self.num_tma_load_bytes,
cta_layout_vmnk=cluster_layout_vmnk,
)
acc_pipeline_producer_group = pipeline.CooperativeGroup(pipeline.Agent.Thread)
num_acc_consumer_threads = len(self.epilog_warp_id) * (
2 if use_2cta_instrs else 1
)
acc_pipeline_consumer_group = pipeline.CooperativeGroup(
pipeline.Agent.Thread, num_acc_consumer_threads
)
try:
acc_pipeline = pipeline.PipelineUmmaAsync.create(
barrier_storage=storage.acc_full_mbar_ptr.data_ptr(),
num_stages=self.num_acc_stage,
producer_group=acc_pipeline_producer_group,
consumer_group=acc_pipeline_consumer_group,
cta_layout_vmnk=cluster_layout_vmnk,
defer_sync=True,
)
except TypeError:
acc_pipeline = pipeline.PipelineUmmaAsync.create(
barrier_storage=storage.acc_full_mbar_ptr.data_ptr(),
num_stages=self.num_acc_stage,
producer_group=acc_pipeline_producer_group,
consumer_group=acc_pipeline_consumer_group,
cta_layout_vmnk=cluster_layout_vmnk,
)
tmem = utils.TmemAllocator(
storage.tmem_holding_buf,
barrier_for_retrieve=self.tmem_alloc_barrier,
allocator_warp_id=self.epilog_warp_id[0],
is_two_cta=use_2cta_instrs,
two_cta_tmem_dealloc_mbar_ptr=storage.tmem_dealloc_mbar_ptr,
)
sC = storage.sC.get_tensor(
c_smem_layout_staged.outer, swizzle=c_smem_layout_staged.inner
)
sA = storage.sA.get_tensor(
a_smem_layout_staged.outer, swizzle=a_smem_layout_staged.inner
)
sB = storage.sB.get_tensor(
b_smem_layout_staged.outer, swizzle=b_smem_layout_staged.inner
)
sSFA = storage.sSFA.get_tensor(sfa_smem_layout_staged)
sSFB = storage.sSFB.get_tensor(sfb_smem_layout_staged)
a_full_mcast_mask = None
b_full_mcast_mask = None
sfa_full_mcast_mask = None
sfb_full_mcast_mask = None
if cutlass.const_expr(self.is_a_mcast or self.is_b_mcast or use_2cta_instrs):
a_full_mcast_mask = cpasync.create_tma_multicast_mask(
cluster_layout_vmnk, block_in_cluster_coord_vmnk, mcast_mode=2
)
b_full_mcast_mask = cpasync.create_tma_multicast_mask(
cluster_layout_vmnk, block_in_cluster_coord_vmnk, mcast_mode=1
)
sfa_full_mcast_mask = cpasync.create_tma_multicast_mask(
cluster_layout_vmnk, block_in_cluster_coord_vmnk, mcast_mode=2
)
sfb_full_mcast_mask = cpasync.create_tma_multicast_mask(
cluster_layout_sfb_vmnk, block_in_cluster_coord_sfb_vmnk, mcast_mode=1
)
tma_cache_policy_a = cutlass.Int64(cutlass.Int64(self.tma_cache_policy_a).ir_value())
tma_cache_policy_b = cutlass.Int64(cutlass.Int64(self.tma_cache_policy_b).ir_value())
tma_cache_policy_sfa = cutlass.Int64(cutlass.Int64(self.tma_cache_policy_sfa).ir_value())
tma_cache_policy_sfb = cutlass.Int64(cutlass.Int64(self.tma_cache_policy_sfb).ir_value())
tma_cache_policy_c = cutlass.Int64(cutlass.Int64(self.tma_cache_policy_c).ir_value())
gA_mkl = cute.local_tile(
mA_mkl, cute.slice_(self.mma_tiler, (None, 0, None)), (None, None, None)
)
gB_nkl = cute.local_tile(
mB_nkl, cute.slice_(self.mma_tiler, (0, None, None)), (None, None, None)
)
gSFA_mkl = cute.local_tile(
mSFA_mkl, cute.slice_(self.mma_tiler, (None, 0, None)), (None, None, None)
)
gSFB_nkl = cute.local_tile(
mSFB_nkl,
cute.slice_(self.mma_tiler_sfb, (0, None, None)),
(None, None, None),
)
gC_mnl = cute.local_tile(
mC_mnl, cute.slice_(self.mma_tiler, (None, None, 0)), (None, None, None)
)
k_tile_cnt = cute.size(gA_mkl, mode=[3])
thr_mma = tiled_mma.get_slice(mma_tile_coord_v)
thr_mma_sfb = tiled_mma_sfb.get_slice(mma_tile_coord_v)
tCgA = thr_mma.partition_A(gA_mkl)
tCgB = thr_mma.partition_B(gB_nkl)
tCgSFA = thr_mma.partition_A(gSFA_mkl)
tCgSFB = thr_mma_sfb.partition_B(gSFB_nkl)
tCgC = thr_mma.partition_C(gC_mnl)
a_cta_layout = cute.make_layout(
cute.slice_(cluster_layout_vmnk, (0, 0, None, 0)).shape
)
tAsA, tAgA = cpasync.tma_partition(
tma_atom_a,
block_in_cluster_coord_vmnk[2],
a_cta_layout,
cute.group_modes(sA, 0, 3),
cute.group_modes(tCgA, 0, 3),
)
b_cta_layout = cute.make_layout(
cute.slice_(cluster_layout_vmnk, (0, None, 0, 0)).shape
)
tBsB, tBgB = cpasync.tma_partition(
tma_atom_b,
block_in_cluster_coord_vmnk[1],
b_cta_layout,
cute.group_modes(sB, 0, 3),
cute.group_modes(tCgB, 0, 3),
)
sfa_cta_layout = a_cta_layout
tAsSFA, tAgSFA = cute.nvgpu.cpasync.tma_partition(
tma_atom_sfa,
block_in_cluster_coord_vmnk[2],
sfa_cta_layout,
cute.group_modes(sSFA, 0, 3),
cute.group_modes(tCgSFA, 0, 3),
)
tAsSFA = cute.filter_zeros(tAsSFA)
tAgSFA = cute.filter_zeros(tAgSFA)
sfb_cta_layout = cute.make_layout(
cute.slice_(cluster_layout_sfb_vmnk, (0, None, 0, 0)).shape
)
tBsSFB, tBgSFB = cute.nvgpu.cpasync.tma_partition(
tma_atom_sfb,
block_in_cluster_coord_sfb_vmnk[1],
sfb_cta_layout,
cute.group_modes(sSFB, 0, 3),
cute.group_modes(tCgSFB, 0, 3),
)
tBsSFB = cute.filter_zeros(tBsSFB)
tBgSFB = cute.filter_zeros(tBgSFB)
tCrA = tiled_mma.make_fragment_A(sA)
tCrB = tiled_mma.make_fragment_B(sB)
acc_shape = tiled_mma.partition_shape_C(self.mma_tiler[:2])
if cutlass.const_expr(self.overlapping_accum):
num_acc_stage_overlapped = 2
tCtAcc_fake = tiled_mma.make_fragment_C(
cute.append(acc_shape, num_acc_stage_overlapped)
)
tCtAcc_fake = cute.make_tensor(
tCtAcc_fake.iterator,
cute.make_layout(
tCtAcc_fake.shape,
stride = (
tCtAcc_fake.stride[0],
tCtAcc_fake.stride[1],
tCtAcc_fake.stride[2],
(256 - self.num_sf_tmem_cols) * tCtAcc_fake.stride[0][1]
)
)
)
else:
tCtAcc_fake = tiled_mma.make_fragment_C(
cute.append(acc_shape, self.num_acc_stage)
)
if warp_idx == self.tma_warp_id:
if cutlass.const_expr(True):
ab_producer_state = pipeline.make_pipeline_state(
pipeline.PipelineUserType.Producer, self.num_ab_stage
)
mma_tile_coord_mnl = (
tile_m_coord,
bidy,
l_idx,
)
tAgA_slice = tAgA[(None, mma_tile_coord_mnl[0], None, mma_tile_coord_mnl[2])]
tBgB_slice = tBgB[(None, mma_tile_coord_mnl[1], None, mma_tile_coord_mnl[2])]
tAgSFA_slice = tAgSFA[(None, mma_tile_coord_mnl[0], None, mma_tile_coord_mnl[2])]
slice_n = mma_tile_coord_mnl[1]
if cutlass.const_expr(self.cta_tile_shape_mnk[1] == 64):
slice_n = mma_tile_coord_mnl[1] // 2
tBgSFB_slice = tBgSFB[(None, slice_n, None, mma_tile_coord_mnl[2])]
ab_producer_state.reset_count()
peek_ab_empty_status = cutlass.Boolean(1)
if ab_producer_state.count < k_tile_cnt:
peek_ab_empty_status = ab_pipeline.producer_try_acquire(ab_producer_state)
for k_tile in cutlass.range(0, k_tile_cnt, 1, unroll=1):
ab_pipeline.producer_acquire(ab_producer_state, peek_ab_empty_status)
try:
cute.copy(
tma_atom_a,
tAgA_slice[(None, ab_producer_state.count)],
tAsA[(None, ab_producer_state.index)],
tma_bar_ptr=ab_pipeline.producer_get_barrier(ab_producer_state),
mcast_mask=a_full_mcast_mask,
cache_policy=tma_cache_policy_a,
)
except TypeError:
cute.copy(
tma_atom_a,
tAgA_slice[(None, ab_producer_state.count)],
tAsA[(None, ab_producer_state.index)],
tma_bar_ptr=ab_pipeline.producer_get_barrier(ab_producer_state),
mcast_mask=a_full_mcast_mask,
)
try:
cute.copy(
tma_atom_b,
tBgB_slice[(None, ab_producer_state.count)],
tBsB[(None, ab_producer_state.index)],
tma_bar_ptr=ab_pipeline.producer_get_barrier(ab_producer_state),
mcast_mask=b_full_mcast_mask,
cache_policy=tma_cache_policy_b,
)
except TypeError:
cute.copy(
tma_atom_b,
tBgB_slice[(None, ab_producer_state.count)],
tBsB[(None, ab_producer_state.index)],
tma_bar_ptr=ab_pipeline.producer_get_barrier(ab_producer_state),
mcast_mask=b_full_mcast_mask,
)
try:
cute.copy(
tma_atom_sfa,
tAgSFA_slice[(None, ab_producer_state.count)],
tAsSFA[(None, ab_producer_state.index)],
tma_bar_ptr=ab_pipeline.producer_get_barrier(ab_producer_state),
mcast_mask=sfa_full_mcast_mask,
cache_policy=tma_cache_policy_sfa,
)
except TypeError:
cute.copy(
tma_atom_sfa,
tAgSFA_slice[(None, ab_producer_state.count)],
tAsSFA[(None, ab_producer_state.index)],
tma_bar_ptr=ab_pipeline.producer_get_barrier(ab_producer_state),
mcast_mask=sfa_full_mcast_mask,
)
try:
cute.copy(
tma_atom_sfb,
tBgSFB_slice[(None, ab_producer_state.count)],
tBsSFB[(None, ab_producer_state.index)],
tma_bar_ptr=ab_pipeline.producer_get_barrier(ab_producer_state),
mcast_mask=sfb_full_mcast_mask,
cache_policy=tma_cache_policy_sfb,
)
except TypeError:
cute.copy(
tma_atom_sfb,
tBgSFB_slice[(None, ab_producer_state.count)],
tBsSFB[(None, ab_producer_state.index)],
tma_bar_ptr=ab_pipeline.producer_get_barrier(ab_producer_state),
mcast_mask=sfb_full_mcast_mask,
)
ab_producer_state.advance()
peek_ab_empty_status = cutlass.Boolean(1)
if ab_producer_state.count < k_tile_cnt:
peek_ab_empty_status = ab_pipeline.producer_try_acquire(
ab_producer_state
)
ab_pipeline.producer_tail(ab_producer_state)
if warp_idx == self.mma_warp_id:
if cutlass.const_expr(True):
tmem.wait_for_alloc()
acc_tmem_ptr = tmem.retrieve_ptr(self.acc_dtype)
tCtAcc_base = cute.make_tensor(acc_tmem_ptr, tCtAcc_fake.layout)
sfa_tmem_ptr = cute.recast_ptr(
acc_tmem_ptr + self.num_accumulator_tmem_cols,
dtype=self.sf_dtype,
)
tCtSFA_layout = blockscaled_utils.make_tmem_layout_sfa(
tiled_mma,
self.mma_tiler,
self.sf_vec_size,
cute.slice_(sfa_smem_layout_staged, (None, None, None, 0)),
)
tCtSFA = cute.make_tensor(sfa_tmem_ptr, tCtSFA_layout)
sfb_tmem_ptr = cute.recast_ptr(
acc_tmem_ptr + self.num_accumulator_tmem_cols + self.num_sfa_tmem_cols,
dtype=self.sf_dtype,
)
tCtSFB_layout = blockscaled_utils.make_tmem_layout_sfb(
tiled_mma,
self.mma_tiler,
self.sf_vec_size,
cute.slice_(sfb_smem_layout_staged, (None, None, None, 0)),
)
tCtSFB = cute.make_tensor(sfb_tmem_ptr, tCtSFB_layout)
(
tiled_copy_s2t_sfa,
tCsSFA_compact_s2t,
tCtSFA_compact_s2t,
) = self.mainloop_s2t_copy_and_partition(sSFA, tCtSFA)
(
tiled_copy_s2t_sfb,
tCsSFB_compact_s2t,
tCtSFB_compact_s2t,
) = self.mainloop_s2t_copy_and_partition(sSFB, tCtSFB)
ab_consumer_state = pipeline.make_pipeline_state(
pipeline.PipelineUserType.Consumer, self.num_ab_stage
)
acc_producer_state = pipeline.make_pipeline_state(
pipeline.PipelineUserType.Producer, self.num_acc_stage
)
mma_tile_coord_mnl = (
tile_m_coord,
bidy,
l_idx,
)
if cutlass.const_expr(self.overlapping_accum):
acc_stage_index = acc_producer_state.phase ^ 1
else:
acc_stage_index = acc_producer_state.index
tCtAcc = tCtAcc_base[(None, None, None, acc_stage_index)]
ab_consumer_state.reset_count()
peek_ab_full_status = cutlass.Boolean(1)
if ab_consumer_state.count < k_tile_cnt and is_leader_cta:
peek_ab_full_status = ab_pipeline.consumer_try_wait(ab_consumer_state)
if is_leader_cta:
acc_pipeline.producer_acquire(acc_producer_state)
tCtSFB_mma = tCtSFB
if cutlass.const_expr(self.cta_tile_shape_mnk[1] == 192):
offset = (
cutlass.Int32(2)
if mma_tile_coord_mnl[1] % 2 == 1
else cutlass.Int32(0)
)
shifted_ptr = cute.recast_ptr(
acc_tmem_ptr
+ self.num_accumulator_tmem_cols
+ self.num_sfa_tmem_cols
+ offset,
dtype=self.sf_dtype,
)
tCtSFB_mma = cute.make_tensor(shifted_ptr, tCtSFB_layout)
elif cutlass.const_expr(self.cta_tile_shape_mnk[1] == 64):
offset = cutlass.Int32((mma_tile_coord_mnl[1] % 2) * 2)
shifted_ptr = cute.recast_ptr(
acc_tmem_ptr
+ self.num_accumulator_tmem_cols
+ self.num_sfa_tmem_cols
+ offset,
dtype=self.sf_dtype,
)
tCtSFB_mma = cute.make_tensor(shifted_ptr, tCtSFB_layout)
tiled_mma.set(tcgen05.Field.ACCUMULATE, False)
for k_tile in range(k_tile_cnt):
if is_leader_cta:
ab_pipeline.consumer_wait(ab_consumer_state, peek_ab_full_status)
s2t_stage_coord = (None, None, None, None, ab_consumer_state.index)
tCsSFA_compact_s2t_staged = tCsSFA_compact_s2t[s2t_stage_coord]
tCsSFB_compact_s2t_staged = tCsSFB_compact_s2t[s2t_stage_coord]
cute.copy(
tiled_copy_s2t_sfa,
tCsSFA_compact_s2t_staged,
tCtSFA_compact_s2t,
)
cute.copy(
tiled_copy_s2t_sfb,
tCsSFB_compact_s2t_staged,
tCtSFB_compact_s2t,
)
num_kblocks = cute.size(tCrA, mode=[2])
for kblock_idx in cutlass.range(num_kblocks, unroll_full=True):
kblock_coord = (None, None, kblock_idx, ab_consumer_state.index)
sf_kblock_coord = (None, None, kblock_idx)
tiled_mma.set(tcgen05.Field.SFA, tCtSFA[sf_kblock_coord].iterator)
tiled_mma.set(
tcgen05.Field.SFB, tCtSFB_mma[sf_kblock_coord].iterator
)
cute.gemm(
tiled_mma,
tCtAcc,
tCrA[kblock_coord],
tCrB[kblock_coord],
tCtAcc,
)
tiled_mma.set(tcgen05.Field.ACCUMULATE, True)
ab_pipeline.consumer_release(ab_consumer_state)
ab_consumer_state.advance()
peek_ab_full_status = cutlass.Boolean(1)
if ab_consumer_state.count < k_tile_cnt:
if is_leader_cta:
peek_ab_full_status = ab_pipeline.consumer_try_wait(ab_consumer_state)
if is_leader_cta:
acc_pipeline.producer_commit(acc_producer_state)
acc_producer_state.advance()
acc_pipeline.producer_tail(acc_producer_state)
if warp_idx < self.mma_warp_id:
if cutlass.const_expr(True):
tmem.allocate(self.num_tmem_alloc_cols)
tmem.wait_for_alloc()
acc_tmem_ptr = tmem.retrieve_ptr(self.acc_dtype)
tCtAcc_base = cute.make_tensor(acc_tmem_ptr, tCtAcc_fake.layout)
epi_tidx = tidx
(
tiled_copy_t2r,
tTR_tAcc_base,
tTR_rAcc,
) = self.epilog_tmem_copy_and_partition(
epi_tidx, tCtAcc_base, tCgC, epi_tile, use_2cta_instrs
)
tTR_rC = cute.make_rmem_tensor(tTR_rAcc.shape, self.c_dtype)
tiled_copy_r2s, tRS_rC, tRS_sC = self.epilog_smem_copy_and_partition(
tiled_copy_t2r, tTR_rC, epi_tidx, sC
)
(
tma_atom_c,
bSG_sC,
bSG_gC_partitioned,
) = self.epilog_gmem_copy_and_partition(
epi_tidx, tma_atom_c, tCgC, epi_tile, sC
)
acc_consumer_state = pipeline.make_pipeline_state(
pipeline.PipelineUserType.Consumer, self.num_acc_stage
)
c_producer_group = pipeline.CooperativeGroup(
pipeline.Agent.Thread,
32 * len(self.epilog_warp_id),
)
c_pipeline = pipeline.PipelineTmaStore.create(
num_stages=self.num_c_stage,
producer_group=c_producer_group,
)
mma_tile_coord_mnl = (
tile_m_coord,
bidy,
l_idx,
)
bSG_gC = bSG_gC_partitioned[(None, None, None, *mma_tile_coord_mnl)]
if cutlass.const_expr(self.overlapping_accum):
acc_stage_index = acc_consumer_state.phase
reverse_subtile = (
cutlass.Boolean(True)
if acc_stage_index == 0
else cutlass.Boolean(False)
)
else:
acc_stage_index = acc_consumer_state.index
tTR_tAcc = tTR_tAcc_base[(None, None, None, None, None, acc_stage_index)]
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))
subtile_cnt = cute.size(tTR_tAcc.shape, mode=[3])
num_prev_subtiles = cutlass.Int32(0)
for subtile_idx in cutlass.range(subtile_cnt):
real_subtile_idx = subtile_idx
if cutlass.const_expr(self.overlapping_accum):
if reverse_subtile:
real_subtile_idx = (
self.cta_tile_shape_mnk[1] // self.epi_tile_n - 1 - subtile_idx
)
tTR_tAcc_mn = tTR_tAcc[(None, None, None, real_subtile_idx)]
cute.copy(tiled_copy_t2r, tTR_tAcc_mn, tTR_rAcc)
if cutlass.const_expr(self.overlapping_accum):
if subtile_idx == self.iter_acc_early_release_in_epilogue:
cute.arch.fence_view_async_tmem_load()
with cute.arch.elect_one():
acc_pipeline.consumer_release(acc_consumer_state)
acc_consumer_state.advance()
acc_vec = tiled_copy_r2s.retile(tTR_rAcc).load()
acc_vec = epilogue_op(acc_vec.to(self.c_dtype))
tRS_rC.store(acc_vec)
c_buffer = (num_prev_subtiles + real_subtile_idx) % self.num_c_stage
cute.copy(
tiled_copy_r2s,
tRS_rC,
tRS_sC[(None, None, None, c_buffer)],
)
cute.arch.fence_proxy(
cute.arch.ProxyKind.async_shared,
space=cute.arch.SharedSpace.shared_cta,
)
self.epilog_sync_barrier.arrive_and_wait()
if warp_idx == self.epilog_warp_id[0]:
try:
cute.copy(
tma_atom_c,
bSG_sC[(None, c_buffer)],
bSG_gC[(None, real_subtile_idx)],
cache_policy=tma_cache_policy_c,
)
except TypeError:
cute.copy(
tma_atom_c,
bSG_sC[(None, c_buffer)],
bSG_gC[(None, real_subtile_idx)],
)
c_pipeline.producer_commit()
c_pipeline.producer_acquire()
self.epilog_sync_barrier.arrive_and_wait()
if cutlass.const_expr(not self.overlapping_accum):
with cute.arch.elect_one():
acc_pipeline.consumer_release(acc_consumer_state)
acc_consumer_state.advance()
tmem.relinquish_alloc_permit()
self.epilog_sync_barrier.arrive_and_wait()
tmem.free(acc_tmem_ptr)
c_pipeline.producer_tail()
@cute.kernel
def kernel_group8(
self,
tiled_mma: cute.TiledMma,
tiled_mma_sfb: cute.TiledMma,
tma_atom_a0: cute.CopyAtom,
mA0_mkl: cute.Tensor,
tma_atom_b0: cute.CopyAtom,
mB0_nkl: cute.Tensor,
tma_atom_sfa0: cute.CopyAtom,
mSFA0_mkl: cute.Tensor,
tma_atom_sfb0: cute.CopyAtom,
mSFB0_nkl: cute.Tensor,
tma_atom_c0: cute.CopyAtom,
mC0_mnl: cute.Tensor,
tma_atom_a1: cute.CopyAtom,
mA1_mkl: cute.Tensor,
tma_atom_b1: cute.CopyAtom,
mB1_nkl: cute.Tensor,
tma_atom_sfa1: cute.CopyAtom,
mSFA1_mkl: cute.Tensor,
tma_atom_sfb1: cute.CopyAtom,
mSFB1_nkl: cute.Tensor,
tma_atom_c1: cute.CopyAtom,
mC1_mnl: cute.Tensor,
tma_atom_a2: cute.CopyAtom,
mA2_mkl: cute.Tensor,
tma_atom_b2: cute.CopyAtom,
mB2_nkl: cute.Tensor,
tma_atom_sfa2: cute.CopyAtom,
mSFA2_mkl: cute.Tensor,
tma_atom_sfb2: cute.CopyAtom,
mSFB2_nkl: cute.Tensor,
tma_atom_c2: cute.CopyAtom,
mC2_mnl: cute.Tensor,
tma_atom_a3: cute.CopyAtom,
mA3_mkl: cute.Tensor,
tma_atom_b3: cute.CopyAtom,
mB3_nkl: cute.Tensor,
tma_atom_sfa3: cute.CopyAtom,
mSFA3_mkl: cute.Tensor,
tma_atom_sfb3: cute.CopyAtom,
mSFB3_nkl: cute.Tensor,
tma_atom_c3: cute.CopyAtom,
mC3_mnl: cute.Tensor,
tma_atom_a4: cute.CopyAtom,
mA4_mkl: cute.Tensor,
tma_atom_b4: cute.CopyAtom,
mB4_nkl: cute.Tensor,
tma_atom_sfa4: cute.CopyAtom,
mSFA4_mkl: cute.Tensor,
tma_atom_sfb4: cute.CopyAtom,
mSFB4_nkl: cute.Tensor,
tma_atom_c4: cute.CopyAtom,
mC4_mnl: cute.Tensor,
tma_atom_a5: cute.CopyAtom,
mA5_mkl: cute.Tensor,
tma_atom_b5: cute.CopyAtom,
mB5_nkl: cute.Tensor,
tma_atom_sfa5: cute.CopyAtom,
mSFA5_mkl: cute.Tensor,
tma_atom_sfb5: cute.CopyAtom,
mSFB5_nkl: cute.Tensor,
tma_atom_c5: cute.CopyAtom,
mC5_mnl: cute.Tensor,
tma_atom_a6: cute.CopyAtom,
mA6_mkl: cute.Tensor,
tma_atom_b6: cute.CopyAtom,
mB6_nkl: cute.Tensor,
tma_atom_sfa6: cute.CopyAtom,
mSFA6_mkl: cute.Tensor,
tma_atom_sfb6: cute.CopyAtom,
mSFB6_nkl: cute.Tensor,
tma_atom_c6: cute.CopyAtom,
mC6_mnl: cute.Tensor,
tma_atom_a7: cute.CopyAtom,
mA7_mkl: cute.Tensor,
tma_atom_b7: cute.CopyAtom,
mB7_nkl: cute.Tensor,
tma_atom_sfa7: cute.CopyAtom,
mSFA7_mkl: cute.Tensor,
tma_atom_sfb7: cute.CopyAtom,
mSFB7_nkl: cute.Tensor,
tma_atom_c7: cute.CopyAtom,
mC7_mnl: cute.Tensor,
cluster_layout_vmnk: cute.Layout,
cluster_layout_sfb_vmnk: cute.Layout,
a_smem_layout_staged: cute.ComposedLayout,
b_smem_layout_staged: cute.ComposedLayout,
sfa_smem_layout_staged: cute.Layout,
sfb_smem_layout_staged: cute.Layout,
c_smem_layout_staged: Union[cute.Layout, cute.ComposedLayout],
epi_tile: cute.Tile,
epilogue_op: cutlass.Constexpr,
blk0: cutlass.Int32,
blk1: cutlass.Int32,
blk2: cutlass.Int32,
blk3: cutlass.Int32,
blk4: cutlass.Int32,
blk5: cutlass.Int32,
blk6: cutlass.Int32,
):
bidx, _, _bidz = cute.arch.block_idx()
atom_thr = cutlass.Int32(cute.size(tiled_mma.thr_id.shape))
b_linear = bidx
p0 = blk0
p1 = p0 + blk1
p2 = p1 + blk2
p3 = p2 + blk3
p4 = p3 + blk4
p5 = p4 + blk5
p6 = p5 + blk6
c0 = cutlass.Int32(1) if b_linear >= p0 else cutlass.Int32(0)
c1 = cutlass.Int32(1) if b_linear >= p1 else cutlass.Int32(0)
c2 = cutlass.Int32(1) if b_linear >= p2 else cutlass.Int32(0)
c3 = cutlass.Int32(1) if b_linear >= p3 else cutlass.Int32(0)
c4 = cutlass.Int32(1) if b_linear >= p4 else cutlass.Int32(0)
c5 = cutlass.Int32(1) if b_linear >= p5 else cutlass.Int32(0)
c6 = cutlass.Int32(1) if b_linear >= p6 else cutlass.Int32(0)
group_idx = c0 + c1 + c2 + c3 + c4 + c5 + c6
base = c0 * blk0 + c1 * blk1 + c2 * blk2 + c3 * blk3 + c4 * blk4 + c5 * blk5 + c6 * blk6
group_idx = cute.arch.make_warp_uniform(group_idx)
base = cute.arch.make_warp_uniform(base)
tile_m_coord = (b_linear - base) // atom_thr
tile_m_coord = cute.arch.make_warp_uniform(tile_m_coord)
l_idx = cutlass.Int32(0)
l_idx = cute.arch.make_warp_uniform(l_idx)
if group_idx == 0:
if cutlass.const_expr(True):
self._kernel_group8_one(
tiled_mma,
tiled_mma_sfb,
tma_atom_a0,
mA0_mkl,
tma_atom_b0,
mB0_nkl,
tma_atom_sfa0,
mSFA0_mkl,
tma_atom_sfb0,
mSFB0_nkl,
tma_atom_c0,
mC0_mnl,
cluster_layout_vmnk,
cluster_layout_sfb_vmnk,
a_smem_layout_staged,
b_smem_layout_staged,
sfa_smem_layout_staged,
sfb_smem_layout_staged,
c_smem_layout_staged,
epi_tile,
epilogue_op,
tile_m_coord,
l_idx,
)
elif group_idx == 1:
if cutlass.const_expr(True):
self._kernel_group8_one(
tiled_mma,
tiled_mma_sfb,
tma_atom_a1,
mA1_mkl,
tma_atom_b1,
mB1_nkl,
tma_atom_sfa1,
mSFA1_mkl,
tma_atom_sfb1,
mSFB1_nkl,
tma_atom_c1,
mC1_mnl,
cluster_layout_vmnk,
cluster_layout_sfb_vmnk,
a_smem_layout_staged,
b_smem_layout_staged,
sfa_smem_layout_staged,
sfb_smem_layout_staged,
c_smem_layout_staged,
epi_tile,
epilogue_op,
tile_m_coord,
l_idx,
)
elif group_idx == 2:
if cutlass.const_expr(True):
self._kernel_group8_one(
tiled_mma,
tiled_mma_sfb,
tma_atom_a2,
mA2_mkl,
tma_atom_b2,
mB2_nkl,
tma_atom_sfa2,
mSFA2_mkl,
tma_atom_sfb2,
mSFB2_nkl,
tma_atom_c2,
mC2_mnl,
cluster_layout_vmnk,
cluster_layout_sfb_vmnk,
a_smem_layout_staged,
b_smem_layout_staged,
sfa_smem_layout_staged,
sfb_smem_layout_staged,
c_smem_layout_staged,
epi_tile,
epilogue_op,
tile_m_coord,
l_idx,
)
elif group_idx == 3:
if cutlass.const_expr(True):
self._kernel_group8_one(
tiled_mma,
tiled_mma_sfb,
tma_atom_a3,
mA3_mkl,
tma_atom_b3,
mB3_nkl,
tma_atom_sfa3,
mSFA3_mkl,
tma_atom_sfb3,
mSFB3_nkl,
tma_atom_c3,
mC3_mnl,
cluster_layout_vmnk,
cluster_layout_sfb_vmnk,
a_smem_layout_staged,
b_smem_layout_staged,
sfa_smem_layout_staged,
sfb_smem_layout_staged,
c_smem_layout_staged,
epi_tile,
epilogue_op,
tile_m_coord,
l_idx,
)
elif group_idx == 4:
if cutlass.const_expr(True):
self._kernel_group8_one(
tiled_mma,
tiled_mma_sfb,
tma_atom_a4,
mA4_mkl,
tma_atom_b4,
mB4_nkl,
tma_atom_sfa4,
mSFA4_mkl,
tma_atom_sfb4,
mSFB4_nkl,
tma_atom_c4,
mC4_mnl,
cluster_layout_vmnk,
cluster_layout_sfb_vmnk,
a_smem_layout_staged,
b_smem_layout_staged,
sfa_smem_layout_staged,
sfb_smem_layout_staged,
c_smem_layout_staged,
epi_tile,
epilogue_op,
tile_m_coord,
l_idx,
)
elif group_idx == 5:
if cutlass.const_expr(True):
self._kernel_group8_one(
tiled_mma,
tiled_mma_sfb,
tma_atom_a5,
mA5_mkl,
tma_atom_b5,
mB5_nkl,
tma_atom_sfa5,
mSFA5_mkl,
tma_atom_sfb5,
mSFB5_nkl,
tma_atom_c5,
mC5_mnl,
cluster_layout_vmnk,
cluster_layout_sfb_vmnk,
a_smem_layout_staged,
b_smem_layout_staged,
sfa_smem_layout_staged,
sfb_smem_layout_staged,
c_smem_layout_staged,
epi_tile,
epilogue_op,
tile_m_coord,
l_idx,
)
elif group_idx == 6:
if cutlass.const_expr(True):
self._kernel_group8_one(
tiled_mma,
tiled_mma_sfb,
tma_atom_a6,
mA6_mkl,
tma_atom_b6,
mB6_nkl,
tma_atom_sfa6,
mSFA6_mkl,
tma_atom_sfb6,
mSFB6_nkl,
tma_atom_c6,
mC6_mnl,
cluster_layout_vmnk,
cluster_layout_sfb_vmnk,
a_smem_layout_staged,
b_smem_layout_staged,
sfa_smem_layout_staged,
sfb_smem_layout_staged,
c_smem_layout_staged,
epi_tile,
epilogue_op,
tile_m_coord,
l_idx,
)
elif group_idx == 7:
if cutlass.const_expr(True):
self._kernel_group8_one(
tiled_mma,
tiled_mma_sfb,
tma_atom_a7,
mA7_mkl,
tma_atom_b7,
mB7_nkl,
tma_atom_sfa7,
mSFA7_mkl,
tma_atom_sfb7,
mSFB7_nkl,
tma_atom_c7,
mC7_mnl,
cluster_layout_vmnk,
cluster_layout_sfb_vmnk,
a_smem_layout_staged,
b_smem_layout_staged,
sfa_smem_layout_staged,
sfb_smem_layout_staged,
c_smem_layout_staged,
epi_tile,
epilogue_op,
tile_m_coord,
l_idx,
)
class Sm100BlockScaledPersistentDenseGemmKernelGroupG2V2(
Sm100BlockScaledPersistentDenseGemmKernelGroupG8V2
):
@cute.jit
def __call__(
self,
a0_ptr: cute.Pointer,
b0_ptr: cute.Pointer,
sfa0_ptr: cute.Pointer,
sfb0_ptr: cute.Pointer,
c0_ptr: cute.Pointer,
a1_ptr: cute.Pointer,
b1_ptr: cute.Pointer,
sfa1_ptr: cute.Pointer,
sfb1_ptr: cute.Pointer,
c1_ptr: cute.Pointer,
problem_size: tuple,
max_active_clusters: cutlass.Constexpr,
epilogue_op: cutlass.Constexpr = lambda x: x,
):
m0, m1, m_max, n, k, l = problem_size
sf_k = k // self.sf_vec_size
a_tensor0 = cute.make_tensor(
a0_ptr,
cute.make_layout(
(m0, cute.assume(k, 32), l),
stride=(cute.assume(k, 32), 1, cute.assume(m0 * k, 32)),
),
)
b_tensor0 = cute.make_tensor(
b0_ptr,
cute.make_layout(
(n, cute.assume(k, 32), l),
stride=(cute.assume(k, 32), 1, cute.assume(n * k, 32)),
),
)
c_tensor0 = cute.make_tensor(c0_ptr, cute.make_layout((m0, n, l), stride=(n, 1, m0 * n)))
sfa_tensor0 = cute.make_tensor(
sfa0_ptr,
cute.make_layout((m0, sf_k, l), stride=(sf_k, 1, m0 * sf_k)),
)
sfb_tensor0 = cute.make_tensor(
sfb0_ptr,
cute.make_layout((n, sf_k, l), stride=(sf_k, 1, n * sf_k)),
)
a_tensor1 = cute.make_tensor(
a1_ptr,
cute.make_layout(
(m1, cute.assume(k, 32), l),
stride=(cute.assume(k, 32), 1, cute.assume(m1 * k, 32)),
),
)
b_tensor1 = cute.make_tensor(
b1_ptr,
cute.make_layout(
(n, cute.assume(k, 32), l),
stride=(cute.assume(k, 32), 1, cute.assume(n * k, 32)),
),
)
c_tensor1 = cute.make_tensor(c1_ptr, cute.make_layout((m1, n, l), stride=(n, 1, m1 * n)))
sfa_tensor1 = cute.make_tensor(
sfa1_ptr,
cute.make_layout((m1, sf_k, l), stride=(sf_k, 1, m1 * sf_k)),
)
sfb_tensor1 = cute.make_tensor(
sfb1_ptr,
cute.make_layout((n, sf_k, l), stride=(sf_k, 1, n * sf_k)),
)
a_tensor0.mark_compact_shape_dynamic(mode=1, stride_order=(2, 0, 1), divisibility=2)
b_tensor0.mark_compact_shape_dynamic(mode=1, stride_order=(2, 0, 1), divisibility=2)
c_tensor0.mark_compact_shape_dynamic(mode=1, stride_order=(2, 0, 1), divisibility=1)
a_tensor1.mark_compact_shape_dynamic(mode=1, stride_order=(2, 0, 1), divisibility=2)
b_tensor1.mark_compact_shape_dynamic(mode=1, stride_order=(2, 0, 1), divisibility=2)
c_tensor1.mark_compact_shape_dynamic(mode=1, stride_order=(2, 0, 1), divisibility=1)
self.a_dtype: Type[cutlass.Numeric] = a_tensor0.element_type
self.b_dtype: Type[cutlass.Numeric] = b_tensor0.element_type
self.sf_dtype: Type[cutlass.Numeric] = sfa_tensor0.element_type
self.c_dtype: Type[cutlass.Numeric] = c_tensor0.element_type
self.a_major_mode = utils.LayoutEnum.from_tensor(a_tensor0).mma_major_mode()
self.b_major_mode = utils.LayoutEnum.from_tensor(b_tensor0).mma_major_mode()
self.c_layout = utils.LayoutEnum.from_tensor(c_tensor0)
if cutlass.const_expr(self.a_dtype != self.b_dtype):
raise TypeError(f"Type must match: {self.a_dtype} != {self.b_dtype}")
self._setup_attributes()
sfa_layout0 = blockscaled_utils.tile_atom_to_shape_SF(a_tensor0.shape, self.sf_vec_size)
sfa_tensor0 = cute.make_tensor(sfa_tensor0.iterator, sfa_layout0)
sfb_layout0 = blockscaled_utils.tile_atom_to_shape_SF(b_tensor0.shape, self.sf_vec_size)
sfb_tensor0 = cute.make_tensor(sfb_tensor0.iterator, sfb_layout0)
sfa_layout1 = blockscaled_utils.tile_atom_to_shape_SF(a_tensor1.shape, self.sf_vec_size)
sfa_tensor1 = cute.make_tensor(sfa_tensor1.iterator, sfa_layout1)
sfb_layout1 = blockscaled_utils.tile_atom_to_shape_SF(b_tensor1.shape, self.sf_vec_size)
sfb_tensor1 = cute.make_tensor(sfb_tensor1.iterator, sfb_layout1)
tiled_mma = sm100_utils.make_blockscaled_trivial_tiled_mma(
self.a_dtype,
self.a_major_mode,
self.b_major_mode,
self.sf_dtype,
self.sf_vec_size,
self.cta_group,
self.mma_inst_shape_mn,
)
tiled_mma_sfb = sm100_utils.make_blockscaled_trivial_tiled_mma(
self.a_dtype,
self.a_major_mode,
self.b_major_mode,
self.sf_dtype,
self.sf_vec_size,
cute.nvgpu.tcgen05.CtaGroup.ONE,
self.mma_inst_shape_mn_sfb,
)
a_op = sm100_utils.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))
tma_atom_a0, tma_tensor_a0 = cute.nvgpu.make_tiled_tma_atom_A(
a_op,
a_tensor0,
a_smem_layout,
self.mma_tiler,
tiled_mma,
self.cluster_layout_vmnk.shape,
)
tma_atom_a1, tma_tensor_a1 = cute.nvgpu.make_tiled_tma_atom_A(
a_op,
a_tensor1,
a_smem_layout,
self.mma_tiler,
tiled_mma,
self.cluster_layout_vmnk.shape,
)
b_op = sm100_utils.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))
tma_atom_b0, tma_tensor_b0 = cute.nvgpu.make_tiled_tma_atom_B(
b_op,
b_tensor0,
b_smem_layout,
self.mma_tiler,
tiled_mma,
self.cluster_layout_vmnk.shape,
)
tma_atom_b1, tma_tensor_b1 = cute.nvgpu.make_tiled_tma_atom_B(
b_op,
b_tensor1,
b_smem_layout,
self.mma_tiler,
tiled_mma,
self.cluster_layout_vmnk.shape,
)
sfa_op = sm100_utils.cluster_shape_to_tma_atom_A(self.cluster_shape_mn, tiled_mma.thr_id)
sfa_smem_layout = cute.slice_(self.sfa_smem_layout_staged, (None, None, None, 0))
tma_atom_sfa0, tma_tensor_sfa0 = cute.nvgpu.make_tiled_tma_atom_A(
sfa_op,
sfa_tensor0,
sfa_smem_layout,
self.mma_tiler,
tiled_mma,
self.cluster_layout_vmnk.shape,
internal_type=cutlass.Int16,
)
tma_atom_sfa1, tma_tensor_sfa1 = cute.nvgpu.make_tiled_tma_atom_A(
sfa_op,
sfa_tensor1,
sfa_smem_layout,
self.mma_tiler,
tiled_mma,
self.cluster_layout_vmnk.shape,
internal_type=cutlass.Int16,
)
sfb_op = sm100_utils.cluster_shape_to_tma_atom_SFB(self.cluster_shape_mn, tiled_mma.thr_id)
sfb_smem_layout = cute.slice_(self.sfb_smem_layout_staged, (None, None, None, 0))
tma_atom_sfb0, tma_tensor_sfb0 = cute.nvgpu.make_tiled_tma_atom_B(
sfb_op,
sfb_tensor0,
sfb_smem_layout,
self.mma_tiler_sfb,
tiled_mma_sfb,
self.cluster_layout_sfb_vmnk.shape,
internal_type=cutlass.Int16,
)
tma_atom_sfb1, tma_tensor_sfb1 = cute.nvgpu.make_tiled_tma_atom_B(
sfb_op,
sfb_tensor1,
sfb_smem_layout,
self.mma_tiler_sfb,
tiled_mma_sfb,
self.cluster_layout_sfb_vmnk.shape,
internal_type=cutlass.Int16,
)
if cutlass.const_expr(self.cta_tile_shape_mnk[1] == 192):
x0 = tma_tensor_sfb0.stride[0][1]
y0 = cute.ceil_div(tma_tensor_sfb0.shape[0][1], 4)
new_shape0 = ((tma_tensor_sfb0.shape[0][0], ((2, 2), y0)), tma_tensor_sfb0.shape[1], tma_tensor_sfb0.shape[2])
x0_times_3 = 3 * x0
new_stride0 = ((tma_tensor_sfb0.stride[0][0], ((x0, x0), x0_times_3)), tma_tensor_sfb0.stride[1], tma_tensor_sfb0.stride[2])
tma_tensor_sfb0 = cute.make_tensor(tma_tensor_sfb0.iterator, cute.make_layout(new_shape0, stride=new_stride0))
x1 = tma_tensor_sfb1.stride[0][1]
y1 = cute.ceil_div(tma_tensor_sfb1.shape[0][1], 4)
new_shape1 = ((tma_tensor_sfb1.shape[0][0], ((2, 2), y1)), tma_tensor_sfb1.shape[1], tma_tensor_sfb1.shape[2])
x1_times_3 = 3 * x1
new_stride1 = ((tma_tensor_sfb1.stride[0][0], ((x1, x1), x1_times_3)), tma_tensor_sfb1.stride[1], tma_tensor_sfb1.stride[2])
tma_tensor_sfb1 = cute.make_tensor(tma_tensor_sfb1.iterator, cute.make_layout(new_shape1, stride=new_stride1))
atom_thr_size = cute.size(tiled_mma.thr_id.shape)
a_copy_size = cute.size_in_bytes(self.a_dtype, a_smem_layout)
b_copy_size = cute.size_in_bytes(self.b_dtype, b_smem_layout)
sfa_copy_size = cute.size_in_bytes(self.sf_dtype, sfa_smem_layout)
sfb_copy_size = cute.size_in_bytes(self.sf_dtype, sfb_smem_layout)
self.num_tma_load_bytes = (a_copy_size + b_copy_size + sfa_copy_size + sfb_copy_size) * atom_thr_size
epi_smem_layout = cute.slice_(self.c_smem_layout_staged, (None, None, 0))
tma_atom_c0, tma_tensor_c0 = cpasync.make_tiled_tma_atom(
cpasync.CopyBulkTensorTileS2GOp(),
c_tensor0,
epi_smem_layout,
self.epi_tile,
)
tma_atom_c1, tma_tensor_c1 = cpasync.make_tiled_tma_atom(
cpasync.CopyBulkTensorTileS2GOp(),
c_tensor1,
epi_smem_layout,
self.epi_tile,
)
tile_m = self.cta_tile_shape_mnk[0]
blk0 = cute.ceil_div(m0, tile_m) * atom_thr_size
blk1 = cute.ceil_div(m1, tile_m) * atom_thr_size
grid_m = blk0 + blk1
grid_n = cute.ceil_div(n, self.cta_tile_shape_mnk[1])
grid = (grid_m, grid_n, l)
self.buffer_align_bytes = 1024
@cute.struct
class SharedStorage:
ab_full_mbar_ptr: cute.struct.MemRange[cutlass.Int64, self.num_ab_stage]
ab_empty_mbar_ptr: cute.struct.MemRange[cutlass.Int64, self.num_ab_stage]
acc_full_mbar_ptr: cute.struct.MemRange[cutlass.Int64, self.num_acc_stage]
acc_empty_mbar_ptr: cute.struct.MemRange[cutlass.Int64, self.num_acc_stage]
tmem_dealloc_mbar_ptr: cutlass.Int64
tmem_holding_buf: cutlass.Int32
sC: cute.struct.Align[
cute.struct.MemRange[
self.c_dtype,
cute.cosize(self.c_smem_layout_staged.outer),
],
self.buffer_align_bytes,
]
sA: cute.struct.Align[
cute.struct.MemRange[
self.a_dtype,
cute.cosize(self.a_smem_layout_staged.outer),
],
self.buffer_align_bytes,
]
sB: cute.struct.Align[
cute.struct.MemRange[
self.b_dtype,
cute.cosize(self.b_smem_layout_staged.outer),
],
self.buffer_align_bytes,
]
sSFA: cute.struct.Align[
cute.struct.MemRange[
self.sf_dtype,
cute.cosize(self.sfa_smem_layout_staged),
],
self.buffer_align_bytes,
]
sSFB: cute.struct.Align[
cute.struct.MemRange[
self.sf_dtype,
cute.cosize(self.sfb_smem_layout_staged),
],
self.buffer_align_bytes,
]
self.shared_storage = SharedStorage
self.kernel_group2(
tiled_mma,
tiled_mma_sfb,
tma_atom_a0,
tma_tensor_a0,
tma_atom_b0,
tma_tensor_b0,
tma_atom_sfa0,
tma_tensor_sfa0,
tma_atom_sfb0,
tma_tensor_sfb0,
tma_atom_c0,
tma_tensor_c0,
tma_atom_a1,
tma_tensor_a1,
tma_atom_b1,
tma_tensor_b1,
tma_atom_sfa1,
tma_tensor_sfa1,
tma_atom_sfb1,
tma_tensor_sfb1,
tma_atom_c1,
tma_tensor_c1,
self.cluster_layout_vmnk,
self.cluster_layout_sfb_vmnk,
self.a_smem_layout_staged,
self.b_smem_layout_staged,
self.sfa_smem_layout_staged,
self.sfb_smem_layout_staged,
self.c_smem_layout_staged,
self.epi_tile,
epilogue_op,
blk0,
).launch(
grid=grid,
block=[self.threads_per_cta, 1, 1],
cluster=(*self.cluster_shape_mn, 1),
min_blocks_per_mp=self.occupancy,
)
return
@cute.kernel
def kernel_group2(
self,
tiled_mma: cute.TiledMma,
tiled_mma_sfb: cute.TiledMma,
tma_atom_a0: cute.CopyAtom,
mA0_mkl: cute.Tensor,
tma_atom_b0: cute.CopyAtom,
mB0_nkl: cute.Tensor,
tma_atom_sfa0: cute.CopyAtom,
mSFA0_mkl: cute.Tensor,
tma_atom_sfb0: cute.CopyAtom,
mSFB0_nkl: cute.Tensor,
tma_atom_c0: cute.CopyAtom,
mC0_mnl: cute.Tensor,
tma_atom_a1: cute.CopyAtom,
mA1_mkl: cute.Tensor,
tma_atom_b1: cute.CopyAtom,
mB1_nkl: cute.Tensor,
tma_atom_sfa1: cute.CopyAtom,
mSFA1_mkl: cute.Tensor,
tma_atom_sfb1: cute.CopyAtom,
mSFB1_nkl: cute.Tensor,
tma_atom_c1: cute.CopyAtom,
mC1_mnl: cute.Tensor,
cluster_layout_vmnk: cute.Layout,
cluster_layout_sfb_vmnk: cute.Layout,
a_smem_layout_staged: cute.ComposedLayout,
b_smem_layout_staged: cute.ComposedLayout,
sfa_smem_layout_staged: cute.Layout,
sfb_smem_layout_staged: cute.Layout,
c_smem_layout_staged: Union[cute.Layout, cute.ComposedLayout],
epi_tile: cute.Tile,
epilogue_op: cutlass.Constexpr,
blk0: cutlass.Int32,
):
bidx, _, _bidz = cute.arch.block_idx()
atom_thr = cutlass.Int32(cute.size(tiled_mma.thr_id.shape))
b_linear = bidx
c0 = cutlass.Int32(1) if b_linear >= blk0 else cutlass.Int32(0)
group_idx = cute.arch.make_warp_uniform(c0)
base = cute.arch.make_warp_uniform(c0 * blk0)
tile_m_coord = (b_linear - base) // atom_thr
tile_m_coord = cute.arch.make_warp_uniform(tile_m_coord)
l_idx = cutlass.Int32(0)
l_idx = cute.arch.make_warp_uniform(l_idx)
if group_idx == 0:
if cutlass.const_expr(True):
self._kernel_group8_one(
tiled_mma,
tiled_mma_sfb,
tma_atom_a0,
mA0_mkl,
tma_atom_b0,
mB0_nkl,
tma_atom_sfa0,
mSFA0_mkl,
tma_atom_sfb0,
mSFB0_nkl,
tma_atom_c0,
mC0_mnl,
cluster_layout_vmnk,
cluster_layout_sfb_vmnk,
a_smem_layout_staged,
b_smem_layout_staged,
sfa_smem_layout_staged,
sfb_smem_layout_staged,
c_smem_layout_staged,
epi_tile,
epilogue_op,
tile_m_coord,
l_idx,
)
elif group_idx == 1:
if cutlass.const_expr(True):
self._kernel_group8_one(
tiled_mma,
tiled_mma_sfb,
tma_atom_a1,
mA1_mkl,
tma_atom_b1,
mB1_nkl,
tma_atom_sfa1,
mSFA1_mkl,
tma_atom_sfb1,
mSFB1_nkl,
tma_atom_c1,
mC1_mnl,
cluster_layout_vmnk,
cluster_layout_sfb_vmnk,
a_smem_layout_staged,
b_smem_layout_staged,
sfa_smem_layout_staged,
sfb_smem_layout_staged,
c_smem_layout_staged,
epi_tile,
epilogue_op,
tile_m_coord,
l_idx,
)
_MAX_ACTIVE_CLUSTERS = 4096
_COMPILED_CACHE = {}
_COMPILED_SUPER_CACHE = {}
_COMPILED_G8_GROUP_CACHE = {}
_COMPILED_G8_GROUP_SUPER_CACHE = {}
_COMPILED_G2_GROUP_CACHE = {}
_COMPILED_G2_GROUP_SUPER_CACHE = {}
_WS_CACHE = {}
_PTR_CACHE = {}
_PTR_CACHE_F4 = {}
_PTR_CACHE_F8 = {}
_PTR_CACHE_F16 = {}
_PTR_CACHE_F16_ALIGN = {}
_PTR_LOOKUP_TOTAL = 0
_PTR_HIT_TOTAL = 0
_PTR_LOOKUP_F4 = 0
_PTR_HIT_F4 = 0
_PTR_LOOKUP_F8 = 0
_PTR_HIT_F8 = 0
_PTR_LOOKUP_F16 = 0
_PTR_HIT_F16 = 0
_CFG_ID_COUNTS = {}
_SUPER_KEY_COUNTS = {}
_G8_CLUSTER_ID_COUNTS = [0, 0, 0]
_G8_CACHE_SET_COUNTS = [0, 0, 0]
_G8_OPT_LEVEL_COUNTS = [0, 0, 0, 0]
_G2_CLUSTER_ID_COUNTS = [0, 0, 0]
_G2_CACHE_SET_COUNTS = [0, 0, 0]
_G2_OPT_LEVEL_COUNTS = [0, 0, 0, 0]
_LAST_CTX = ""
_LAST_KEY_G8 = 0
_LAST_KEY_G2 = 0
_LAST_KEY = 0
_ALIGN_FIX_TOTAL = 0
_ALIGN_FIX_A = 0
_ALIGN_FIX_B = 0
_ALIGN_FIX_SFA = 0
_ALIGN_FIX_SFB = 0
_ALIGN_FIX_C = 0
def _fmt_counts(d: dict, limit: int) -> str:
if not d:
return ""
items = sorted(d.items(), key=lambda x: (-int(x[1]), int(x[0])))
if len(items) > limit:
items = items[:limit]
return ",".join(f"{int(k)}:{int(v)}" for k, v in items)
def _dump_align_diag():
global _ALIGN_FIX_TOTAL, _ALIGN_FIX_A, _ALIGN_FIX_B, _ALIGN_FIX_SFA, _ALIGN_FIX_SFB, _ALIGN_FIX_C, _LAST_CTX, _LAST_KEY, _LAST_KEY_G8, _LAST_KEY_G2
ptr_lookup = int(_PTR_LOOKUP_TOTAL)
ptr_hit = int(_PTR_HIT_TOTAL)
cfg_counts = _fmt_counts(_CFG_ID_COUNTS, 4)
super_counts = _fmt_counts(_SUPER_KEY_COUNTS, 4)
g8_cid = _fmt_counts(
{0: _G8_CLUSTER_ID_COUNTS[0], 1: _G8_CLUSTER_ID_COUNTS[1], 2: _G8_CLUSTER_ID_COUNTS[2]},
3,
)
g8_cs = _fmt_counts(
{0: _G8_CACHE_SET_COUNTS[0], 4: _G8_CACHE_SET_COUNTS[1], 8: _G8_CACHE_SET_COUNTS[2]},
3,
)
g8_opt = _fmt_counts(
{0: _G8_OPT_LEVEL_COUNTS[0], 1: _G8_OPT_LEVEL_COUNTS[1], 2: _G8_OPT_LEVEL_COUNTS[2], 3: _G8_OPT_LEVEL_COUNTS[3]},
3,
)
g2_cid = _fmt_counts(
{0: _G2_CLUSTER_ID_COUNTS[0], 1: _G2_CLUSTER_ID_COUNTS[1], 2: _G2_CLUSTER_ID_COUNTS[2]},
3,
)
g2_cs = _fmt_counts(
{0: _G2_CACHE_SET_COUNTS[0], 4: _G2_CACHE_SET_COUNTS[1], 8: _G2_CACHE_SET_COUNTS[2]},
3,
)
g2_opt = _fmt_counts(
{0: _G2_OPT_LEVEL_COUNTS[0], 1: _G2_OPT_LEVEL_COUNTS[1], 2: _G2_OPT_LEVEL_COUNTS[2], 3: _G2_OPT_LEVEL_COUNTS[3]},
3,
)
last_ctx = _LAST_CTX
last_key = int(_LAST_KEY)
last_key_g8 = int(_LAST_KEY_G8)
last_key_g2 = int(_LAST_KEY_G2)
if last_ctx and len(last_ctx) > _DIAG_CTX_MAX_CHARS:
last_ctx = last_ctx[:_DIAG_CTX_MAX_CHARS]
if _ALIGN_FIX_TOTAL == 0:
print(
f"ALIGNFIX total=0 ws={len(_WS_CACHE)} ptr={ptr_hit}/{ptr_lookup} cfg={cfg_counts} sk={super_counts} g8cid={g8_cid} g8cs={g8_cs} g8opt={g8_opt} g2cid={g2_cid} g2cs={g2_cs} g2opt={g2_opt} lk={last_key} lk8={last_key_g8} lk2={last_key_g2} last={last_ctx}",
flush=True,
)
return
print(
f"ALIGNFIX total={_ALIGN_FIX_TOTAL} a={_ALIGN_FIX_A} b={_ALIGN_FIX_B} sfa={_ALIGN_FIX_SFA} sfb={_ALIGN_FIX_SFB} c={_ALIGN_FIX_C} ws={len(_WS_CACHE)} ptr={ptr_hit}/{ptr_lookup} cfg={cfg_counts} sk={super_counts} g8cid={g8_cid} g8cs={g8_cs} g8opt={g8_opt} g2cid={g2_cid} g2cs={g2_cs} g2opt={g2_opt} lk={last_key} lk8={last_key_g8} lk2={last_key_g2} last={last_ctx}",
flush=True,
)
atexit.register(_dump_align_diag)
def _is_aligned_ptr(ptr_value: int, align: int) -> bool:
return (ptr_value & (align - 1)) == 0
def _alloc_like_aligned(tensor: torch.Tensor, align: int, copy_src: bool) -> torch.Tensor:
out = torch.empty_like(tensor)
if not _is_aligned_ptr(out.data_ptr(), align):
out = torch.empty_like(tensor)
if not _is_aligned_ptr(out.data_ptr(), align):
raise RuntimeError("aligned allocation failed")
if copy_src:
out.copy_(tensor)
return out
def _get_ws_like_aligned(tensor: torch.Tensor, align: int) -> torch.Tensor:
key = (
tensor.device,
tensor.dtype,
tuple(tensor.shape),
tuple(tensor.stride()),
int(align),
)
buf = _WS_CACHE.get(key)
if buf is None or not _is_aligned_ptr(buf.data_ptr(), align):
buf = torch.empty_like(tensor)
if not _is_aligned_ptr(buf.data_ptr(), align):
buf = torch.empty_like(tensor)
if not _is_aligned_ptr(buf.data_ptr(), align):
raise RuntimeError("aligned allocation failed")
_WS_CACHE[key] = buf
return buf
def _get_gmem_ptr(dtype: Type[cutlass.Numeric], ptr_value: int, align: int):
key = (
dtype,
int(ptr_value),
int(align),
)
out = _PTR_CACHE.get(key)
if out is None:
out = make_ptr(dtype, ptr_value, cute.AddressSpace.gmem, assumed_align=align)
_PTR_CACHE[key] = out
return out
def _get_gmem_ptr_f4(ptr_value: int):
out = _PTR_CACHE_F4.get(ptr_value)
if out is None:
out = make_ptr(
cutlass.Float4E2M1FN,
ptr_value,
cute.AddressSpace.gmem,
assumed_align=_ASSUMED_ALIGN_BYTES,
)
_PTR_CACHE_F4[ptr_value] = out
return out
def _get_gmem_ptr_f8(ptr_value: int):
out = _PTR_CACHE_F8.get(ptr_value)
if out is None:
out = make_ptr(
cutlass.Float8E4M3FN,
ptr_value,
cute.AddressSpace.gmem,
assumed_align=_ASSUMED_ALIGN_BYTES,
)
_PTR_CACHE_F8[ptr_value] = out
return out
def _get_gmem_ptr_f16(ptr_value: int):
out = _PTR_CACHE_F16.get(ptr_value)
if out is None:
out = make_ptr(
cutlass.Float16,
ptr_value,
cute.AddressSpace.gmem,
assumed_align=_ASSUMED_ALIGN_BYTES,
)
_PTR_CACHE_F16[ptr_value] = out
return out
def _get_gmem_ptr_f16_align(ptr_value: int, align: int):
key = (int(ptr_value), int(align))
out = _PTR_CACHE_F16_ALIGN.get(key)
if out is None:
out = make_ptr(
cutlass.Float16,
ptr_value,
cute.AddressSpace.gmem,
assumed_align=int(align),
)
_PTR_CACHE_F16_ALIGN[key] = out
return out
def _get_compiled(cfg: _GemmCfg):
key = (
cfg.mma_tiler_mn[0],
cfg.mma_tiler_mn[1],
cfg.cluster_shape_mn[0],
cfg.cluster_shape_mn[1],
int(cfg.occupancy),
int(cfg.cache_policy_set),
)
compiled = _COMPILED_CACHE.get(key)
if compiled is not None:
return compiled
gemm = Sm100BlockScaledPersistentDenseGemmKernel(
_SF_VEC_SIZE,
cfg.mma_tiler_mn,
cfg.cluster_shape_mn,
cfg.cache_policy_set,
occupancy=cfg.occupancy,
)
cutlass.cuda.initialize_cuda_context()
a_ptr = make_ptr(
cutlass.Float4E2M1FN, 0, cute.AddressSpace.gmem, assumed_align=_ASSUMED_ALIGN_BYTES
)
b_ptr = make_ptr(
cutlass.Float4E2M1FN, 0, cute.AddressSpace.gmem, assumed_align=_ASSUMED_ALIGN_BYTES
)
sfa_ptr = make_ptr(
cutlass.Float8E4M3FN, 0, cute.AddressSpace.gmem, assumed_align=_ASSUMED_ALIGN_BYTES
)
sfb_ptr = make_ptr(
cutlass.Float8E4M3FN, 0, cute.AddressSpace.gmem, assumed_align=_ASSUMED_ALIGN_BYTES
)
c_ptr = make_ptr(
cutlass.Float16, 0, cute.AddressSpace.gmem, assumed_align=_ASSUMED_ALIGN_BYTES
)
try:
compiled = cute.compile(
gemm,
a_ptr,
b_ptr,
sfa_ptr,
sfb_ptr,
c_ptr,
(0, 0, 0, 0),
_MAX_ACTIVE_CLUSTERS,
options="--opt-level 3",
)
except Exception as e:
ctx = (
f"mma={cfg.mma_tiler_mn} cluster={tuple(cfg.cluster_shape_mn)} "
f"occ={int(cfg.occupancy)} cache_set={int(cfg.cache_policy_set)} opt=3"
)
global _LAST_CTX
_LAST_CTX = ctx
raise RuntimeError(
f"compile_failed [{ctx}]: {type(e).__name__}: {e}"
) from e
_COMPILED_CACHE[key] = compiled
return compiled
def _get_compiled_g8_group(
cfg: _GemmCfg,
cluster_shape_mn: Tuple[int, int],
cache_policy_set: int,
opt_level: int,
c_align: int = _ASSUMED_ALIGN_BYTES,
):
key = (
cfg.mma_tiler_mn[0],
cfg.mma_tiler_mn[1],
cluster_shape_mn[0],
cluster_shape_mn[1],
int(cfg.occupancy),
int(cache_policy_set),
int(opt_level),
int(c_align),
)
compiled = _COMPILED_G8_GROUP_CACHE.get(key)
if compiled is not None:
return compiled
gemm = Sm100BlockScaledPersistentDenseGemmKernelGroupG8V2(
_SF_VEC_SIZE,
cfg.mma_tiler_mn,
cluster_shape_mn,
int(cache_policy_set),
occupancy=cfg.occupancy,
)
cutlass.cuda.initialize_cuda_context()
a_ptr = make_ptr(
cutlass.Float4E2M1FN, 0, cute.AddressSpace.gmem, assumed_align=_ASSUMED_ALIGN_BYTES
)
b_ptr = make_ptr(
cutlass.Float4E2M1FN, 0, cute.AddressSpace.gmem, assumed_align=_ASSUMED_ALIGN_BYTES
)
sfa_ptr = make_ptr(
cutlass.Float8E4M3FN, 0, cute.AddressSpace.gmem, assumed_align=_ASSUMED_ALIGN_BYTES
)
sfb_ptr = make_ptr(
cutlass.Float8E4M3FN, 0, cute.AddressSpace.gmem, assumed_align=_ASSUMED_ALIGN_BYTES
)
c_ptr = make_ptr(
cutlass.Float16, 0, cute.AddressSpace.gmem, assumed_align=int(c_align)
)
try:
opt = f"--opt-level {int(opt_level)}"
compiled = cute.compile(
gemm,
a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr,
a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr,
a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr,
a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr,
a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr,
a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr,
a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr,
a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr,
(0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0),
_MAX_ACTIVE_CLUSTERS,
options=opt,
)
except Exception as e:
ctx = (
f"mma={cfg.mma_tiler_mn} cluster={tuple(cluster_shape_mn)} "
f"occ={int(cfg.occupancy)} cache_set={int(cache_policy_set)} opt={int(opt_level)} c_align={int(c_align)}"
)
global _LAST_CTX
_LAST_CTX = ctx
raise RuntimeError(
f"g8_group_compile_failed [{ctx}]: {type(e).__name__}: {e}"
) from e
_COMPILED_G8_GROUP_CACHE[key] = compiled
return compiled
def _get_compiled_g8_group_super(cfg_id: int):
compiled = _COMPILED_G8_GROUP_SUPER_CACHE.get(cfg_id)
if compiled is not None:
return compiled
cfg = _cfg_from_g8_id(cfg_id)
compiled = _get_compiled_g8_group(cfg, (1, 1), int(cfg.cache_policy_set), 1)
_COMPILED_G8_GROUP_SUPER_CACHE[cfg_id] = compiled
return compiled
def _g2_super_key_for_nk(n: int, k: int, m_max: int) -> int:
if n == 3072 and k == 4096:
use_occ1 = int(m_max) >= int(_G2_RANKED_OCC1_MMAX_GE_N3072_K4096)
base = 64 if use_occ1 else 16
return base | int(_G2_N192_USE_SET4)
if n == 4096 and k == 1536:
use_occ1 = int(m_max) >= int(_G2_RANKED_OCC1_MMAX_GE_N4096_K1536)
return 81 if use_occ1 else 33
return 0
def _get_compiled_g2_group(
cfg: _GemmCfg,
cluster_shape_mn: Tuple[int, int],
cache_policy_set: int,
opt_level: int,
c_align: int = _ASSUMED_ALIGN_BYTES,
):
key = (
cfg.mma_tiler_mn[0],
cfg.mma_tiler_mn[1],
cluster_shape_mn[0],
cluster_shape_mn[1],
int(cfg.occupancy),
int(cache_policy_set),
int(opt_level),
int(c_align),
)
compiled = _COMPILED_G2_GROUP_CACHE.get(key)
if compiled is not None:
return compiled
gemm = Sm100BlockScaledPersistentDenseGemmKernelGroupG2V2(
_SF_VEC_SIZE,
cfg.mma_tiler_mn,
cluster_shape_mn,
int(cache_policy_set),
occupancy=cfg.occupancy,
)
cutlass.cuda.initialize_cuda_context()
a_ptr = make_ptr(
cutlass.Float4E2M1FN, 0, cute.AddressSpace.gmem, assumed_align=_ASSUMED_ALIGN_BYTES
)
b_ptr = make_ptr(
cutlass.Float4E2M1FN, 0, cute.AddressSpace.gmem, assumed_align=_ASSUMED_ALIGN_BYTES
)
sfa_ptr = make_ptr(
cutlass.Float8E4M3FN, 0, cute.AddressSpace.gmem, assumed_align=_ASSUMED_ALIGN_BYTES
)
sfb_ptr = make_ptr(
cutlass.Float8E4M3FN, 0, cute.AddressSpace.gmem, assumed_align=_ASSUMED_ALIGN_BYTES
)
c_ptr = make_ptr(
cutlass.Float16, 0, cute.AddressSpace.gmem, assumed_align=int(c_align)
)
try:
opt = f"--opt-level {int(opt_level)}"
compiled = cute.compile(
gemm,
a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr,
a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr,
(0, 0, 0, 0, 0, 0),
_MAX_ACTIVE_CLUSTERS,
options=opt,
)
except Exception as e:
ctx = (
f"mma={cfg.mma_tiler_mn} cluster={tuple(cluster_shape_mn)} "
f"occ={int(cfg.occupancy)} cache_set={int(cache_policy_set)} opt={int(opt_level)} c_align={int(c_align)}"
)
global _LAST_CTX
_LAST_CTX = ctx
raise RuntimeError(
f"g2_group_compile_failed [{ctx}]: {type(e).__name__}: {e}"
) from e
_COMPILED_G2_GROUP_CACHE[key] = compiled
return compiled
def _get_compiled_g2_group_super(super_key: int):
compiled = _COMPILED_G2_GROUP_SUPER_CACHE.get(super_key)
if compiled is not None:
return compiled
cfg = _cfg_for_super_key(super_key)
compiled = _get_compiled_g2_group(cfg, (1, 1), int(cfg.cache_policy_set), 1)
_COMPILED_G2_GROUP_SUPER_CACHE[super_key] = compiled
return compiled
def _cfg_for_super_key(super_key: int) -> _GemmCfg:
if super_key == 0:
return _DEFAULT_CFG
hi = super_key >> 8
if hi == 8:
return _cfg_from_g8_id(super_key & 255)
if (super_key & ~1) == 16:
return _CFG_TILER_N192_SET4 if (super_key & 1) else _CFG_TILER_N192
if (super_key & ~1) == 64:
return _CFG_TILER_N192_SET4_OCC1 if (super_key & 1) else _CFG_TILER_N192_OCC1
if (super_key & ~1) == 32:
return _CFG_TILER_N128_SET4 if (super_key & 1) else _CFG_TILER_N128_SET8
if (super_key & ~1) == 80:
return _CFG_TILER_N128_SET4_OCC1 if (super_key & 1) else _CFG_TILER_N128_SET8_OCC1
if (super_key & ~1) == 48:
return _CFG_TILER_N256_SET4 if (super_key & 1) else _CFG_TILER_N256_SET8
raise ValueError(f"invalid super_key={super_key}")
def _get_compiled_super(super_key: int):
compiled = _COMPILED_SUPER_CACHE.get(super_key)
if compiled is not None:
return compiled
compiled = _get_compiled(_cfg_for_super_key(super_key))
_COMPILED_SUPER_CACHE[super_key] = compiled
return compiled
def custom_kernel(data):
abc_tensors, _sfasfb_tensors, sfasfb_reordered_tensors, problem_sizes = data
g = len(problem_sizes)
global _ALIGN_FIX_TOTAL, _ALIGN_FIX_A, _ALIGN_FIX_B, _ALIGN_FIX_SFA, _ALIGN_FIX_SFB, _ALIGN_FIX_C, _LAST_CTX, _LAST_KEY, _LAST_KEY_G8, _LAST_KEY_G2
diag_sync = int(_DIAG_SYNC)
diag_level = int(_DIAG_LEVEL)
if diag_sync and diag_level == 0:
diag_level = 1
if _ENABLE_G8_GROUP and g == 8:
ps0 = problem_sizes[0]
ps1 = problem_sizes[1]
ps2 = problem_sizes[2]
ps3 = problem_sizes[3]
ps4 = problem_sizes[4]
ps5 = problem_sizes[5]
ps6 = problem_sizes[6]
ps7 = problem_sizes[7]
n0 = int(ps0[1])
k0 = int(ps0[2])
l0 = int(ps0[3])
same_nkl = (
int(ps1[1]) == n0 and int(ps1[2]) == k0 and int(ps1[3]) == l0
and int(ps2[1]) == n0 and int(ps2[2]) == k0 and int(ps2[3]) == l0
and int(ps3[1]) == n0 and int(ps3[2]) == k0 and int(ps3[3]) == l0
and int(ps4[1]) == n0 and int(ps4[2]) == k0 and int(ps4[3]) == l0
and int(ps5[1]) == n0 and int(ps5[2]) == k0 and int(ps5[3]) == l0
and int(ps6[1]) == n0 and int(ps6[2]) == k0 and int(ps6[3]) == l0
and int(ps7[1]) == n0 and int(ps7[2]) == k0 and int(ps7[3]) == l0
)
if same_nkl and l0 == 1:
m0 = int(ps0[0])
m1 = int(ps1[0])
m2 = int(ps2[0])
m3 = int(ps3[0])
m4 = int(ps4[0])
m5 = int(ps5[0])
m6 = int(ps6[0])
m7 = int(ps7[0])
m_ge_128 = (
(m0 >= 128)
+ (m1 >= 128)
+ (m2 >= 128)
+ (m3 >= 128)
+ (m4 >= 128)
+ (m5 >= 128)
+ (m6 >= 128)
+ (m7 >= 128)
)
m_max = max(m0, m1, m2, m3, m4, m5, m6, m7)
cfg_id = _g8_cfg_id_for_nk(n0, k0)
if diag_level:
_v = _CFG_ID_COUNTS.get(cfg_id)
_CFG_ID_COUNTS[cfg_id] = 1 if _v is None else (_v + 1)
cfg = _cfg_from_g8_id(cfg_id)
cache_policy_set = _g8_group_cache_policy_set_for_nk(n0, k0, int(cfg.cache_policy_set))
cluster_shape_mn = _g8_group_cluster_shape_for_nk(n0, k0)
opt_level = _g8_group_opt_level_for_nk(n0, k0)
if n0 == 4096 and k0 == 7168:
if m_max >= int(_G8_RANKED_CLUSTER_MMAX_GE_N4096_K7168_TO_M1N1):
cluster_shape_mn = (1, 1)
elif m_max >= int(_G8_RANKED_CLUSTER_MMAX_GE_N4096_K7168_TO_M1N2):
cluster_shape_mn = (1, 2)
else:
cluster_shape_mn = (1, 4)
cache_policy_set = _clamp_cache_set(8 if m_ge_128 >= 5 else 4, int(cfg.cache_policy_set))
elif n0 == 7168 and k0 == 2048:
cluster_shape_mn = (1, 4) if m_max >= int(_G8_RANKED_CLUSTER_MMAX_GE_N7168_K2048_TO_M1N4) else (1, 2)
chosen_cs = int(_G8_RANKED_CACHE_SET_N7168_K2048_ALT) if int(_G8_RANKED_CACHE_USE_ALT_N7168_K2048) else int(_G8_RANKED_CACHE_SET_N7168_K2048)
cache_policy_set = _clamp_cache_set(
chosen_cs if m_ge_128 >= 5 else 4,
int(cfg.cache_policy_set),
)
if diag_level:
_cid = 2 if cluster_shape_mn[1] == 4 else (1 if cluster_shape_mn[1] == 2 else 0)
_G8_CLUSTER_ID_COUNTS[_cid] += 1
_cs = int(cache_policy_set)
_G8_CACHE_SET_COUNTS[2 if _cs == 8 else (1 if _cs == 4 else 0)] += 1
_G8_OPT_LEVEL_COUNTS[int(opt_level)] += 1
assumed_align = _ASSUMED_ALIGN_BYTES
get_ws_like_aligned = _get_ws_like_aligned
get_f4 = _get_gmem_ptr_f4
get_f8 = _get_gmem_ptr_f8
get_f16 = _get_gmem_ptr_f16
a0, b0, c0_out = abc_tensors[0]
a1, b1, c1_out = abc_tensors[1]
a2, b2, c2_out = abc_tensors[2]
a3, b3, c3_out = abc_tensors[3]
a4, b4, c4_out = abc_tensors[4]
a5, b5, c5_out = abc_tensors[5]
a6, b6, c6_out = abc_tensors[6]
a7, b7, c7_out = abc_tensors[7]
sfa0_p, sfb0_p = sfasfb_reordered_tensors[0]
sfa1_p, sfb1_p = sfasfb_reordered_tensors[1]
sfa2_p, sfb2_p = sfasfb_reordered_tensors[2]
sfa3_p, sfb3_p = sfasfb_reordered_tensors[3]
sfa4_p, sfb4_p = sfasfb_reordered_tensors[4]
sfa5_p, sfb5_p = sfasfb_reordered_tensors[5]
sfa6_p, sfb6_p = sfasfb_reordered_tensors[6]
sfa7_p, sfb7_p = sfasfb_reordered_tensors[7]
mask = assumed_align - 1
a0_p = a0.data_ptr()
a1_p = a1.data_ptr()
a2_p = a2.data_ptr()
a3_p = a3.data_ptr()
a4_p = a4.data_ptr()
a5_p = a5.data_ptr()
a6_p = a6.data_ptr()
a7_p = a7.data_ptr()
b0_p = b0.data_ptr()
b1_p = b1.data_ptr()
b2_p = b2.data_ptr()
b3_p = b3.data_ptr()
b4_p = b4.data_ptr()
b5_p = b5.data_ptr()
b6_p = b6.data_ptr()
b7_p = b7.data_ptr()
sfa0_pv = sfa0_p.data_ptr()
sfa1_pv = sfa1_p.data_ptr()
sfa2_pv = sfa2_p.data_ptr()
sfa3_pv = sfa3_p.data_ptr()
sfa4_pv = sfa4_p.data_ptr()
sfa5_pv = sfa5_p.data_ptr()
sfa6_pv = sfa6_p.data_ptr()
sfa7_pv = sfa7_p.data_ptr()
sfb0_pv = sfb0_p.data_ptr()
sfb1_pv = sfb1_p.data_ptr()
sfb2_pv = sfb2_p.data_ptr()
sfb3_pv = sfb3_p.data_ptr()
sfb4_pv = sfb4_p.data_ptr()
sfb5_pv = sfb5_p.data_ptr()
sfb6_pv = sfb6_p.data_ptr()
sfb7_pv = sfb7_p.data_ptr()
c0_out_p = c0_out.data_ptr()
c1_out_p = c1_out.data_ptr()
c2_out_p = c2_out.data_ptr()
c3_out_p = c3_out.data_ptr()
c4_out_p = c4_out.data_ptr()
c5_out_p = c5_out.data_ptr()
c6_out_p = c6_out.data_ptr()
c7_out_p = c7_out.data_ptr()
c_or = (
c0_out_p
| c1_out_p
| c2_out_p
| c3_out_p
| c4_out_p
| c5_out_p
| c6_out_p
| c7_out_p
)
c_align = assumed_align
if (c_or & 63) == 0 and (c_or & mask) != 0:
c_align = 64
compiled_group = _get_compiled_g8_group(cfg, cluster_shape_mn, cache_policy_set, opt_level, c_align)
if c_align == assumed_align:
all_aligned = (
(
a0_p
| a1_p
| a2_p
| a3_p
| a4_p
| a5_p
| a6_p
| a7_p
| b0_p
| b1_p
| b2_p
| b3_p
| b4_p
| b5_p
| b6_p
| b7_p
| sfa0_pv
| sfa1_pv
| sfa2_pv
| sfa3_pv
| sfa4_pv
| sfa5_pv
| sfa6_pv
| sfa7_pv
| sfb0_pv
| sfb1_pv
| sfb2_pv
| sfb3_pv
| sfb4_pv
| sfb5_pv
| sfb6_pv
| sfb7_pv
| c_or
)
& mask
) == 0
else:
all_aligned = (
(
a0_p
| a1_p
| a2_p
| a3_p
| a4_p
| a5_p
| a6_p
| a7_p
| b0_p
| b1_p
| b2_p
| b3_p
| b4_p
| b5_p
| b6_p
| b7_p
| sfa0_pv
| sfa1_pv
| sfa2_pv
| sfa3_pv
| sfa4_pv
| sfa5_pv
| sfa6_pv
| sfa7_pv
| sfb0_pv
| sfb1_pv
| sfb2_pv
| sfb3_pv
| sfb4_pv
| sfb5_pv
| sfb6_pv
| sfb7_pv
)
& mask
) == 0
if not all_aligned and (a0_p & mask):
_ALIGN_FIX_TOTAL += 1
_ALIGN_FIX_A += 1
a0_ws = get_ws_like_aligned(a0, assumed_align)
a0_ws.copy_(a0)
a0 = a0_ws
a0_p = a0_ws.data_ptr()
if not all_aligned and (a1_p & mask):
_ALIGN_FIX_TOTAL += 1
_ALIGN_FIX_A += 1
a1_ws = get_ws_like_aligned(a1, assumed_align)
a1_ws.copy_(a1)
a1 = a1_ws
a1_p = a1_ws.data_ptr()
if not all_aligned and (a2_p & mask):
_ALIGN_FIX_TOTAL += 1
_ALIGN_FIX_A += 1
a2_ws = get_ws_like_aligned(a2, assumed_align)
a2_ws.copy_(a2)
a2 = a2_ws
a2_p = a2_ws.data_ptr()
if not all_aligned and (a3_p & mask):
_ALIGN_FIX_TOTAL += 1
_ALIGN_FIX_A += 1
a3_ws = get_ws_like_aligned(a3, assumed_align)
a3_ws.copy_(a3)
a3 = a3_ws
a3_p = a3_ws.data_ptr()
if not all_aligned and (a4_p & mask):
_ALIGN_FIX_TOTAL += 1
_ALIGN_FIX_A += 1
a4_ws = get_ws_like_aligned(a4, assumed_align)
a4_ws.copy_(a4)
a4 = a4_ws
a4_p = a4_ws.data_ptr()
if not all_aligned and (a5_p & mask):
_ALIGN_FIX_TOTAL += 1
_ALIGN_FIX_A += 1
a5_ws = get_ws_like_aligned(a5, assumed_align)
a5_ws.copy_(a5)
a5 = a5_ws
a5_p = a5_ws.data_ptr()
if not all_aligned and (a6_p & mask):
_ALIGN_FIX_TOTAL += 1
_ALIGN_FIX_A += 1
a6_ws = get_ws_like_aligned(a6, assumed_align)
a6_ws.copy_(a6)
a6 = a6_ws
a6_p = a6_ws.data_ptr()
if not all_aligned and (a7_p & mask):
_ALIGN_FIX_TOTAL += 1
_ALIGN_FIX_A += 1
a7_ws = get_ws_like_aligned(a7, assumed_align)
a7_ws.copy_(a7)
a7 = a7_ws
a7_p = a7_ws.data_ptr()
if not all_aligned and (b0_p & mask):
_ALIGN_FIX_TOTAL += 1
_ALIGN_FIX_B += 1
b0_ws = get_ws_like_aligned(b0, assumed_align)
b0_ws.copy_(b0)
b0 = b0_ws
b0_p = b0_ws.data_ptr()
if not all_aligned and (b1_p & mask):
_ALIGN_FIX_TOTAL += 1
_ALIGN_FIX_B += 1
b1_ws = get_ws_like_aligned(b1, assumed_align)
b1_ws.copy_(b1)
b1 = b1_ws
b1_p = b1_ws.data_ptr()
if not all_aligned and (b2_p & mask):
_ALIGN_FIX_TOTAL += 1
_ALIGN_FIX_B += 1
b2_ws = get_ws_like_aligned(b2, assumed_align)
b2_ws.copy_(b2)
b2 = b2_ws
b2_p = b2_ws.data_ptr()
if not all_aligned and (b3_p & mask):
_ALIGN_FIX_TOTAL += 1
_ALIGN_FIX_B += 1
b3_ws = get_ws_like_aligned(b3, assumed_align)
b3_ws.copy_(b3)
b3 = b3_ws
b3_p = b3_ws.data_ptr()
if not all_aligned and (b4_p & mask):
_ALIGN_FIX_TOTAL += 1
_ALIGN_FIX_B += 1
b4_ws = get_ws_like_aligned(b4, assumed_align)
b4_ws.copy_(b4)
b4 = b4_ws
b4_p = b4_ws.data_ptr()
if not all_aligned and (b5_p & mask):
_ALIGN_FIX_TOTAL += 1
_ALIGN_FIX_B += 1
b5_ws = get_ws_like_aligned(b5, assumed_align)
b5_ws.copy_(b5)
b5 = b5_ws
b5_p = b5_ws.data_ptr()
if not all_aligned and (b6_p & mask):
_ALIGN_FIX_TOTAL += 1
_ALIGN_FIX_B += 1
b6_ws = get_ws_like_aligned(b6, assumed_align)
b6_ws.copy_(b6)
b6 = b6_ws
b6_p = b6_ws.data_ptr()
if not all_aligned and (b7_p & mask):
_ALIGN_FIX_TOTAL += 1
_ALIGN_FIX_B += 1
b7_ws = get_ws_like_aligned(b7, assumed_align)
b7_ws.copy_(b7)
b7 = b7_ws
b7_p = b7_ws.data_ptr()
if not all_aligned and (sfa0_pv & mask):
_ALIGN_FIX_TOTAL += 1
_ALIGN_FIX_SFA += 1
sfa0_ws = get_ws_like_aligned(sfa0_p, assumed_align)
sfa0_ws.copy_(sfa0_p)
sfa0_p = sfa0_ws
sfa0_pv = sfa0_ws.data_ptr()
if not all_aligned and (sfa1_pv & mask):
_ALIGN_FIX_TOTAL += 1
_ALIGN_FIX_SFA += 1
sfa1_ws = get_ws_like_aligned(sfa1_p, assumed_align)
sfa1_ws.copy_(sfa1_p)
sfa1_p = sfa1_ws
sfa1_pv = sfa1_ws.data_ptr()
if not all_aligned and (sfa2_pv & mask):
_ALIGN_FIX_TOTAL += 1
_ALIGN_FIX_SFA += 1
sfa2_ws = get_ws_like_aligned(sfa2_p, assumed_align)
sfa2_ws.copy_(sfa2_p)
sfa2_p = sfa2_ws
sfa2_pv = sfa2_ws.data_ptr()
if not all_aligned and (sfa3_pv & mask):
_ALIGN_FIX_TOTAL += 1
_ALIGN_FIX_SFA += 1
sfa3_ws = get_ws_like_aligned(sfa3_p, assumed_align)
sfa3_ws.copy_(sfa3_p)
sfa3_p = sfa3_ws
sfa3_pv = sfa3_ws.data_ptr()
if not all_aligned and (sfa4_pv & mask):
_ALIGN_FIX_TOTAL += 1
_ALIGN_FIX_SFA += 1
sfa4_ws = get_ws_like_aligned(sfa4_p, assumed_align)
sfa4_ws.copy_(sfa4_p)
sfa4_p = sfa4_ws
sfa4_pv = sfa4_ws.data_ptr()
if not all_aligned and (sfa5_pv & mask):
_ALIGN_FIX_TOTAL += 1
_ALIGN_FIX_SFA += 1
sfa5_ws = get_ws_like_aligned(sfa5_p, assumed_align)
sfa5_ws.copy_(sfa5_p)
sfa5_p = sfa5_ws
sfa5_pv = sfa5_ws.data_ptr()
if not all_aligned and (sfa6_pv & mask):
_ALIGN_FIX_TOTAL += 1
_ALIGN_FIX_SFA += 1
sfa6_ws = get_ws_like_aligned(sfa6_p, assumed_align)
sfa6_ws.copy_(sfa6_p)
sfa6_p = sfa6_ws
sfa6_pv = sfa6_ws.data_ptr()
if not all_aligned and (sfa7_pv & mask):
_ALIGN_FIX_TOTAL += 1
_ALIGN_FIX_SFA += 1
sfa7_ws = get_ws_like_aligned(sfa7_p, assumed_align)
sfa7_ws.copy_(sfa7_p)
sfa7_p = sfa7_ws
sfa7_pv = sfa7_ws.data_ptr()
if not all_aligned and (sfb0_pv & mask):
_ALIGN_FIX_TOTAL += 1
_ALIGN_FIX_SFB += 1
sfb0_ws = get_ws_like_aligned(sfb0_p, assumed_align)
sfb0_ws.copy_(sfb0_p)
sfb0_p = sfb0_ws
sfb0_pv = sfb0_ws.data_ptr()
if not all_aligned and (sfb1_pv & mask):
_ALIGN_FIX_TOTAL += 1
_ALIGN_FIX_SFB += 1
sfb1_ws = get_ws_like_aligned(sfb1_p, assumed_align)
sfb1_ws.copy_(sfb1_p)
sfb1_p = sfb1_ws
sfb1_pv = sfb1_ws.data_ptr()
if not all_aligned and (sfb2_pv & mask):
_ALIGN_FIX_TOTAL += 1
_ALIGN_FIX_SFB += 1
sfb2_ws = get_ws_like_aligned(sfb2_p, assumed_align)
sfb2_ws.copy_(sfb2_p)
sfb2_p = sfb2_ws
sfb2_pv = sfb2_ws.data_ptr()
if not all_aligned and (sfb3_pv & mask):
_ALIGN_FIX_TOTAL += 1
_ALIGN_FIX_SFB += 1
sfb3_ws = get_ws_like_aligned(sfb3_p, assumed_align)
sfb3_ws.copy_(sfb3_p)
sfb3_p = sfb3_ws
sfb3_pv = sfb3_ws.data_ptr()
if not all_aligned and (sfb4_pv & mask):
_ALIGN_FIX_TOTAL += 1
_ALIGN_FIX_SFB += 1
sfb4_ws = get_ws_like_aligned(sfb4_p, assumed_align)
sfb4_ws.copy_(sfb4_p)
sfb4_p = sfb4_ws
sfb4_pv = sfb4_ws.data_ptr()
if not all_aligned and (sfb5_pv & mask):
_ALIGN_FIX_TOTAL += 1
_ALIGN_FIX_SFB += 1
sfb5_ws = get_ws_like_aligned(sfb5_p, assumed_align)
sfb5_ws.copy_(sfb5_p)
sfb5_p = sfb5_ws
sfb5_pv = sfb5_ws.data_ptr()
if not all_aligned and (sfb6_pv & mask):
_ALIGN_FIX_TOTAL += 1
_ALIGN_FIX_SFB += 1
sfb6_ws = get_ws_like_aligned(sfb6_p, assumed_align)
sfb6_ws.copy_(sfb6_p)
sfb6_p = sfb6_ws
sfb6_pv = sfb6_ws.data_ptr()
if not all_aligned and (sfb7_pv & mask):
_ALIGN_FIX_TOTAL += 1
_ALIGN_FIX_SFB += 1
sfb7_ws = get_ws_like_aligned(sfb7_p, assumed_align)
sfb7_ws.copy_(sfb7_p)
sfb7_p = sfb7_ws
sfb7_pv = sfb7_ws.data_ptr()
c0_buf = c0_out
c1_buf = c1_out
c2_buf = c2_out
c3_buf = c3_out
c4_buf = c4_out
c5_buf = c5_out
c6_buf = c6_out
c7_buf = c7_out
c0_buf_p = c0_out_p
c1_buf_p = c1_out_p
c2_buf_p = c2_out_p
c3_buf_p = c3_out_p
c4_buf_p = c4_out_p
c5_buf_p = c5_out_p
c6_buf_p = c6_out_p
c7_buf_p = c7_out_p
if c_align == assumed_align and not all_aligned and (c0_buf_p & mask):
_ALIGN_FIX_TOTAL += 1
_ALIGN_FIX_C += 1
c0_buf = get_ws_like_aligned(c0_buf, assumed_align)
c0_buf_p = c0_buf.data_ptr()
if c_align == assumed_align and not all_aligned and (c1_buf_p & mask):
_ALIGN_FIX_TOTAL += 1
_ALIGN_FIX_C += 1
c1_buf = get_ws_like_aligned(c1_buf, assumed_align)
c1_buf_p = c1_buf.data_ptr()
if c_align == assumed_align and not all_aligned and (c2_buf_p & mask):
_ALIGN_FIX_TOTAL += 1
_ALIGN_FIX_C += 1
c2_buf = get_ws_like_aligned(c2_buf, assumed_align)
c2_buf_p = c2_buf.data_ptr()
if c_align == assumed_align and not all_aligned and (c3_buf_p & mask):
_ALIGN_FIX_TOTAL += 1
_ALIGN_FIX_C += 1
c3_buf = get_ws_like_aligned(c3_buf, assumed_align)
c3_buf_p = c3_buf.data_ptr()
if c_align == assumed_align and not all_aligned and (c4_buf_p & mask):
_ALIGN_FIX_TOTAL += 1
_ALIGN_FIX_C += 1
c4_buf = get_ws_like_aligned(c4_buf, assumed_align)
c4_buf_p = c4_buf.data_ptr()
if c_align == assumed_align and not all_aligned and (c5_buf_p & mask):
_ALIGN_FIX_TOTAL += 1
_ALIGN_FIX_C += 1
c5_buf = get_ws_like_aligned(c5_buf, assumed_align)
c5_buf_p = c5_buf.data_ptr()
if c_align == assumed_align and not all_aligned and (c6_buf_p & mask):
_ALIGN_FIX_TOTAL += 1
_ALIGN_FIX_C += 1
c6_buf = get_ws_like_aligned(c6_buf, assumed_align)
c6_buf_p = c6_buf.data_ptr()
if c_align == assumed_align and not all_aligned and (c7_buf_p & mask):
_ALIGN_FIX_TOTAL += 1
_ALIGN_FIX_C += 1
c7_buf = get_ws_like_aligned(c7_buf, assumed_align)
c7_buf_p = c7_buf.data_ptr()
a0_ptr = get_f4(a0_p)
a1_ptr = get_f4(a1_p)
a2_ptr = get_f4(a2_p)
a3_ptr = get_f4(a3_p)
a4_ptr = get_f4(a4_p)
a5_ptr = get_f4(a5_p)
a6_ptr = get_f4(a6_p)
a7_ptr = get_f4(a7_p)
b0_ptr = get_f4(b0_p)
b1_ptr = get_f4(b1_p)
b2_ptr = get_f4(b2_p)
b3_ptr = get_f4(b3_p)
b4_ptr = get_f4(b4_p)
b5_ptr = get_f4(b5_p)
b6_ptr = get_f4(b6_p)
b7_ptr = get_f4(b7_p)
sfa0_ptr = get_f8(sfa0_pv)
sfa1_ptr = get_f8(sfa1_pv)
sfa2_ptr = get_f8(sfa2_pv)
sfa3_ptr = get_f8(sfa3_pv)
sfa4_ptr = get_f8(sfa4_pv)
sfa5_ptr = get_f8(sfa5_pv)
sfa6_ptr = get_f8(sfa6_pv)
sfa7_ptr = get_f8(sfa7_pv)
sfb0_ptr = get_f8(sfb0_pv)
sfb1_ptr = get_f8(sfb1_pv)
sfb2_ptr = get_f8(sfb2_pv)
sfb3_ptr = get_f8(sfb3_pv)
sfb4_ptr = get_f8(sfb4_pv)
sfb5_ptr = get_f8(sfb5_pv)
sfb6_ptr = get_f8(sfb6_pv)
sfb7_ptr = get_f8(sfb7_pv)
if c_align == assumed_align:
c0_ptr = get_f16(c0_buf_p)
c1_ptr = get_f16(c1_buf_p)
c2_ptr = get_f16(c2_buf_p)
c3_ptr = get_f16(c3_buf_p)
c4_ptr = get_f16(c4_buf_p)
c5_ptr = get_f16(c5_buf_p)
c6_ptr = get_f16(c6_buf_p)
c7_ptr = get_f16(c7_buf_p)
else:
get_f16a = _get_gmem_ptr_f16_align
c0_ptr = get_f16a(c0_buf_p, c_align)
c1_ptr = get_f16a(c1_buf_p, c_align)
c2_ptr = get_f16a(c2_buf_p, c_align)
c3_ptr = get_f16a(c3_buf_p, c_align)
c4_ptr = get_f16a(c4_buf_p, c_align)
c5_ptr = get_f16a(c5_buf_p, c_align)
c6_ptr = get_f16a(c6_buf_p, c_align)
c7_ptr = get_f16a(c7_buf_p, c_align)
_cid = 2 if cluster_shape_mn[1] == 4 else (1 if cluster_shape_mn[1] == 2 else 0)
key_val = (
(8 << 24)
| (int(cfg_id) << 16)
| (int(_cid) << 12)
| (int(cache_policy_set) << 8)
| (int(opt_level) << 4)
| (1 if c_align == 64 else 0)
)
_LAST_KEY_G8 = int(key_val)
_LAST_KEY = int(key_val)
try:
compiled_group(
a0_ptr, b0_ptr, sfa0_ptr, sfb0_ptr, c0_ptr,
a1_ptr, b1_ptr, sfa1_ptr, sfb1_ptr, c1_ptr,
a2_ptr, b2_ptr, sfa2_ptr, sfb2_ptr, c2_ptr,
a3_ptr, b3_ptr, sfa3_ptr, sfb3_ptr, c3_ptr,
a4_ptr, b4_ptr, sfa4_ptr, sfb4_ptr, c4_ptr,
a5_ptr, b5_ptr, sfa5_ptr, sfb5_ptr, c5_ptr,
a6_ptr, b6_ptr, sfa6_ptr, sfb6_ptr, c6_ptr,
a7_ptr, b7_ptr, sfa7_ptr, sfb7_ptr, c7_ptr,
(m0, m1, m2, m3, m4, m5, m6, m7, m_max, n0, k0, l0),
)
if diag_sync:
torch.cuda.synchronize()
except Exception as e:
ctx = (
f"g=8 nk=({n0},{k0}) l={l0} cfg_id={int(cfg_id)} "
f"cluster={tuple(cluster_shape_mn)} cache_set={int(cache_policy_set)} "
f"opt={int(opt_level)} c_align={int(c_align)} m_ge_128={int(m_ge_128)} m_max={int(m_max)}"
)
_LAST_CTX = ctx
raise RuntimeError(
f"g8_group_run_failed [{ctx}]: {type(e).__name__}: {e}"
) from e
try:
if c0_buf is not c0_out:
c0_out.copy_(c0_buf)
if c1_buf is not c1_out:
c1_out.copy_(c1_buf)
if c2_buf is not c2_out:
c2_out.copy_(c2_buf)
if c3_buf is not c3_out:
c3_out.copy_(c3_buf)
if c4_buf is not c4_out:
c4_out.copy_(c4_buf)
if c5_buf is not c5_out:
c5_out.copy_(c5_buf)
if c6_buf is not c6_out:
c6_out.copy_(c6_buf)
if c7_buf is not c7_out:
c7_out.copy_(c7_buf)
except Exception as e:
ctx2 = (
f"g=8 nk=({n0},{k0}) l={l0} cfg_id={int(cfg_id)} "
f"cluster={tuple(cluster_shape_mn)} cache_set={int(cache_policy_set)} opt={int(opt_level)} c_align={int(c_align)}"
)
_LAST_CTX = ctx2
raise RuntimeError(
f"g8_group_copy_failed [{ctx2}]: {type(e).__name__}: {e}"
) from e
return [c0_out, c1_out, c2_out, c3_out, c4_out, c5_out, c6_out, c7_out]
if _ENABLE_G2_GROUP and g == 2:
n0 = int(problem_sizes[0][1])
k0 = int(problem_sizes[0][2])
l0 = int(problem_sizes[0][3])
same_nkl = (
int(problem_sizes[1][1]) == n0
and int(problem_sizes[1][2]) == k0
and int(problem_sizes[1][3]) == l0
)
if same_nkl and l0 == 1:
ps0 = problem_sizes[0]
ps1 = problem_sizes[1]
m0 = int(ps0[0])
m1 = int(ps1[0])
m_max = m0 if m0 > m1 else m1
super_key = _g2_super_key_for_nk(n0, k0, m_max)
if diag_level:
_v = _SUPER_KEY_COUNTS.get(super_key)
_SUPER_KEY_COUNTS[super_key] = 1 if _v is None else (_v + 1)
cfg = _cfg_for_super_key(super_key)
cluster_shape_mn = _g2_group_cluster_shape_for_nk(n0, k0)
opt_level = _g2_group_opt_level_for_nk(n0, k0)
cache_policy_set = _g2_group_cache_policy_set_for_nk(n0, k0, int(cfg.cache_policy_set))
if diag_level:
_cid = 2 if cluster_shape_mn[1] == 4 else (1 if cluster_shape_mn[1] == 2 else 0)
_G2_CLUSTER_ID_COUNTS[_cid] += 1
_cs = int(cache_policy_set)
_G2_CACHE_SET_COUNTS[2 if _cs == 8 else (1 if _cs == 4 else 0)] += 1
_G2_OPT_LEVEL_COUNTS[int(opt_level)] += 1
assumed_align = _ASSUMED_ALIGN_BYTES
get_ws_like_aligned = _get_ws_like_aligned
get_f4 = _get_gmem_ptr_f4
get_f8 = _get_gmem_ptr_f8
get_f16 = _get_gmem_ptr_f16
a0, b0, c0_out = abc_tensors[0]
a1, b1, c1_out = abc_tensors[1]
sfa0_p, sfb0_p = sfasfb_reordered_tensors[0]
sfa1_p, sfb1_p = sfasfb_reordered_tensors[1]
mask = assumed_align - 1
a0_p = a0.data_ptr()
a1_p = a1.data_ptr()
b0_p = b0.data_ptr()
b1_p = b1.data_ptr()
sfa0_pv = sfa0_p.data_ptr()
sfa1_pv = sfa1_p.data_ptr()
sfb0_pv = sfb0_p.data_ptr()
sfb1_pv = sfb1_p.data_ptr()
c0_out_p = c0_out.data_ptr()
c1_out_p = c1_out.data_ptr()
c_or = c0_out_p | c1_out_p
c_align = assumed_align
if (c_or & 63) == 0 and (c_or & mask) != 0:
c_align = 64
compiled_group = _get_compiled_g2_group(cfg, cluster_shape_mn, int(cache_policy_set), opt_level, c_align)
if c_align == assumed_align:
all_aligned = (
(
a0_p
| a1_p
| b0_p
| b1_p
| sfa0_pv
| sfa1_pv
| sfb0_pv
| sfb1_pv
| c_or
)
& mask
) == 0
else:
all_aligned = (
(
a0_p
| a1_p
| b0_p
| b1_p
| sfa0_pv
| sfa1_pv
| sfb0_pv
| sfb1_pv
)
& mask
) == 0
if not all_aligned and (a0_p & mask):
_ALIGN_FIX_TOTAL += 1
_ALIGN_FIX_A += 1
a0_ws = get_ws_like_aligned(a0, assumed_align)
a0_ws.copy_(a0)
a0 = a0_ws
a0_p = a0_ws.data_ptr()
if not all_aligned and (a1_p & mask):
_ALIGN_FIX_TOTAL += 1
_ALIGN_FIX_A += 1
a1_ws = get_ws_like_aligned(a1, assumed_align)
a1_ws.copy_(a1)
a1 = a1_ws
a1_p = a1_ws.data_ptr()
if not all_aligned and (b0_p & mask):
_ALIGN_FIX_TOTAL += 1
_ALIGN_FIX_B += 1
b0_ws = get_ws_like_aligned(b0, assumed_align)
b0_ws.copy_(b0)
b0 = b0_ws
b0_p = b0_ws.data_ptr()
if not all_aligned and (b1_p & mask):
_ALIGN_FIX_TOTAL += 1
_ALIGN_FIX_B += 1
b1_ws = get_ws_like_aligned(b1, assumed_align)
b1_ws.copy_(b1)
b1 = b1_ws
b1_p = b1_ws.data_ptr()
if not all_aligned and (sfa0_pv & mask):
_ALIGN_FIX_TOTAL += 1
_ALIGN_FIX_SFA += 1
sfa0_ws = get_ws_like_aligned(sfa0_p, assumed_align)
sfa0_ws.copy_(sfa0_p)
sfa0_p = sfa0_ws
sfa0_pv = sfa0_ws.data_ptr()
if not all_aligned and (sfa1_pv & mask):
_ALIGN_FIX_TOTAL += 1
_ALIGN_FIX_SFA += 1
sfa1_ws = get_ws_like_aligned(sfa1_p, assumed_align)
sfa1_ws.copy_(sfa1_p)
sfa1_p = sfa1_ws
sfa1_pv = sfa1_ws.data_ptr()
if not all_aligned and (sfb0_pv & mask):
_ALIGN_FIX_TOTAL += 1
_ALIGN_FIX_SFB += 1
sfb0_ws = get_ws_like_aligned(sfb0_p, assumed_align)
sfb0_ws.copy_(sfb0_p)
sfb0_p = sfb0_ws
sfb0_pv = sfb0_ws.data_ptr()
if not all_aligned and (sfb1_pv & mask):
_ALIGN_FIX_TOTAL += 1
_ALIGN_FIX_SFB += 1
sfb1_ws = get_ws_like_aligned(sfb1_p, assumed_align)
sfb1_ws.copy_(sfb1_p)
sfb1_p = sfb1_ws
sfb1_pv = sfb1_ws.data_ptr()
c0_buf = c0_out
c1_buf = c1_out
c0_buf_p = c0_out_p
c1_buf_p = c1_out_p
if c_align == assumed_align and not all_aligned and (c0_buf_p & mask):
_ALIGN_FIX_TOTAL += 1
_ALIGN_FIX_C += 1
c0_buf = get_ws_like_aligned(c0_buf, assumed_align)
c0_buf_p = c0_buf.data_ptr()
if c_align == assumed_align and not all_aligned and (c1_buf_p & mask):
_ALIGN_FIX_TOTAL += 1
_ALIGN_FIX_C += 1
c1_buf = get_ws_like_aligned(c1_buf, assumed_align)
c1_buf_p = c1_buf.data_ptr()
a0_ptr = get_f4(a0_p)
a1_ptr = get_f4(a1_p)
b0_ptr = get_f4(b0_p)
b1_ptr = get_f4(b1_p)
sfa0_ptr = get_f8(sfa0_pv)
sfa1_ptr = get_f8(sfa1_pv)
sfb0_ptr = get_f8(sfb0_pv)
sfb1_ptr = get_f8(sfb1_pv)
if c_align == assumed_align:
c0_ptr = get_f16(c0_buf_p)
c1_ptr = get_f16(c1_buf_p)
else:
get_f16a = _get_gmem_ptr_f16_align
c0_ptr = get_f16a(c0_buf_p, c_align)
c1_ptr = get_f16a(c1_buf_p, c_align)
_cid = 2 if cluster_shape_mn[1] == 4 else (1 if cluster_shape_mn[1] == 2 else 0)
key_val = (
(2 << 24)
| (int(super_key) << 12)
| (int(_cid) << 10)
| (int(cache_policy_set) << 6)
| (int(opt_level) << 4)
| (1 if c_align == 64 else 0)
)
_LAST_KEY_G2 = int(key_val)
_LAST_KEY = int(key_val)
try:
compiled_group(
a0_ptr, b0_ptr, sfa0_ptr, sfb0_ptr, c0_ptr,
a1_ptr, b1_ptr, sfa1_ptr, sfb1_ptr, c1_ptr,
(m0, m1, m_max, n0, k0, l0),
)
if diag_sync:
torch.cuda.synchronize()
except Exception as e:
ctx = (
f"g=2 nk=({n0},{k0}) l={l0} super_key={int(super_key)} "
f"cluster={tuple(cluster_shape_mn)} cache_set={int(cache_policy_set)} "
f"opt={int(opt_level)} c_align={int(c_align)} m_max={int(m_max)}"
)
_LAST_CTX = ctx
raise RuntimeError(
f"g2_group_run_failed [{ctx}]: {type(e).__name__}: {e}"
) from e
try:
if c0_buf is not c0_out:
c0_out.copy_(c0_buf)
if c1_buf is not c1_out:
c1_out.copy_(c1_buf)
except Exception as e:
ctx2 = (
f"g=2 nk=({n0},{k0}) l={l0} super_key={int(super_key)} "
f"cluster={tuple(cluster_shape_mn)} cache_set={int(cache_policy_set)} opt={int(opt_level)} c_align={int(c_align)}"
)
_LAST_CTX = ctx2
raise RuntimeError(
f"g2_group_copy_failed [{ctx2}]: {type(e).__name__}: {e}"
) from e
return [c0_out, c1_out]
compiled_uniform = None
compiled_per_group = None
if g == 8:
n0 = int(problem_sizes[0][1])
k0 = int(problem_sizes[0][2])
same_nk = True
for i in range(1, g):
if int(problem_sizes[i][1]) != n0 or int(problem_sizes[i][2]) != k0:
same_nk = False
break
if same_nk:
cfg_id = _g8_cfg_id_for_nk(n0, k0)
compiled_uniform = _get_compiled_super((8 << 8) | cfg_id)
else:
compiled_per_group = []
for i in range(g):
n = int(problem_sizes[i][1])
k = int(problem_sizes[i][2])
cfg_id = _g8_cfg_id_for_nk(n, k)
compiled_per_group.append(
_get_compiled_super((8 << 8) | cfg_id)
)
else:
compiled_default = _get_compiled_super(0)
maybe_use_n192 = False
for i in range(g):
if int(problem_sizes[i][1]) == 3072 and int(problem_sizes[i][2]) == 4096:
maybe_use_n192 = True
break
if maybe_use_n192:
use_set4 = int(_G2_N192_USE_SET4)
compiled_n192 = _get_compiled_super(16 | use_set4)
all_n192 = True
for i in range(g):
if int(problem_sizes[i][1]) != 3072 or int(problem_sizes[i][2]) != 4096:
all_n192 = False
break
if all_n192:
compiled_uniform = compiled_n192
else:
compiled_per_group = []
for i in range(g):
n = int(problem_sizes[i][1])
k = int(problem_sizes[i][2])
compiled_per_group.append(
compiled_n192 if (n == 3072 and k == 4096) else compiled_default
)
else:
compiled_uniform = compiled_default
assumed_align = _ASSUMED_ALIGN_BYTES
outputs = []
is_aligned_ptr = _is_aligned_ptr
get_ws_like_aligned = _get_ws_like_aligned
get_gmem_ptr = _get_gmem_ptr
f4 = cutlass.Float4E2M1FN
f8 = cutlass.Float8E4M3FN
f16 = cutlass.Float16
if compiled_per_group is None:
compiled = compiled_uniform
for i in range(g):
a, b, c = abc_tensors[i]
sfa_p, sfb_p = sfasfb_reordered_tensors[i]
m, n, k, l = problem_sizes[i]
c_out = c
if not is_aligned_ptr(a.data_ptr(), assumed_align):
_ALIGN_FIX_TOTAL += 1
_ALIGN_FIX_A += 1
a_ws = get_ws_like_aligned(a, assumed_align)
a_ws.copy_(a)
a = a_ws
if not is_aligned_ptr(b.data_ptr(), assumed_align):
_ALIGN_FIX_TOTAL += 1
_ALIGN_FIX_B += 1
b_ws = get_ws_like_aligned(b, assumed_align)
b_ws.copy_(b)
b = b_ws
if not is_aligned_ptr(sfa_p.data_ptr(), assumed_align):
_ALIGN_FIX_TOTAL += 1
_ALIGN_FIX_SFA += 1
sfa_ws = get_ws_like_aligned(sfa_p, assumed_align)
sfa_ws.copy_(sfa_p)
sfa_p = sfa_ws
if not is_aligned_ptr(sfb_p.data_ptr(), assumed_align):
_ALIGN_FIX_TOTAL += 1
_ALIGN_FIX_SFB += 1
sfb_ws = get_ws_like_aligned(sfb_p, assumed_align)
sfb_ws.copy_(sfb_p)
sfb_p = sfb_ws
if not is_aligned_ptr(c.data_ptr(), assumed_align):
_ALIGN_FIX_TOTAL += 1
_ALIGN_FIX_C += 1
c = get_ws_like_aligned(c, assumed_align)
a_ptr = get_gmem_ptr(f4, a.data_ptr(), assumed_align)
b_ptr = get_gmem_ptr(f4, b.data_ptr(), assumed_align)
sfa_ptr = get_gmem_ptr(f8, sfa_p.data_ptr(), assumed_align)
sfb_ptr = get_gmem_ptr(f8, sfb_p.data_ptr(), assumed_align)
c_ptr = get_gmem_ptr(f16, c.data_ptr(), assumed_align)
compiled(a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr, (m, n, k, l))
if c is not c_out:
c_out.copy_(c)
outputs.append(c_out)
else:
for i in range(g):
compiled = compiled_per_group[i]
a, b, c = abc_tensors[i]
sfa_p, sfb_p = sfasfb_reordered_tensors[i]
m, n, k, l = problem_sizes[i]
c_out = c
if not is_aligned_ptr(a.data_ptr(), assumed_align):
_ALIGN_FIX_TOTAL += 1
_ALIGN_FIX_A += 1
a_ws = get_ws_like_aligned(a, assumed_align)
a_ws.copy_(a)
a = a_ws
if not is_aligned_ptr(b.data_ptr(), assumed_align):
_ALIGN_FIX_TOTAL += 1
_ALIGN_FIX_B += 1
b_ws = get_ws_like_aligned(b, assumed_align)
b_ws.copy_(b)
b = b_ws
if not is_aligned_ptr(sfa_p.data_ptr(), assumed_align):
_ALIGN_FIX_TOTAL += 1
_ALIGN_FIX_SFA += 1
sfa_ws = get_ws_like_aligned(sfa_p, assumed_align)
sfa_ws.copy_(sfa_p)
sfa_p = sfa_ws
if not is_aligned_ptr(sfb_p.data_ptr(), assumed_align):
_ALIGN_FIX_TOTAL += 1
_ALIGN_FIX_SFB += 1
sfb_ws = get_ws_like_aligned(sfb_p, assumed_align)
sfb_ws.copy_(sfb_p)
sfb_p = sfb_ws
if not is_aligned_ptr(c.data_ptr(), assumed_align):
_ALIGN_FIX_TOTAL += 1
_ALIGN_FIX_C += 1
c = get_ws_like_aligned(c, assumed_align)
a_ptr = get_gmem_ptr(f4, a.data_ptr(), assumed_align)
b_ptr = get_gmem_ptr(f4, b.data_ptr(), assumed_align)
sfa_ptr = get_gmem_ptr(f8, sfa_p.data_ptr(), assumed_align)
sfb_ptr = get_gmem_ptr(f8, sfb_p.data_ptr(), assumed_align)
c_ptr = get_gmem_ptr(f16, c.data_ptr(), assumed_align)
compiled(a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr, (m, n, k, l))
if c is not c_out:
c_out.copy_(c)
outputs.append(c_out)
return outputs
__all__ = ["custom_kernel"]
scrolls · 6124 lines total
Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0
Changes from previous submission
Against this author's previous submission submission 388567.
⋯ diff truncated: revisions differ almost entirely
Best evidence level for this revision: reported
JSON