Skip to content
KernelIndex
Search⌘K

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
NVFP4 group GEMMsuite of 4 cases
NVIDIA B200
18.4µs
#4 of 145
2026-01-25

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-epilogueself.epi_tile = sm100_utils.compute_epilogue_tile_shape(
mbarrierself.epilog_sync_barrier = pipeline.NamedBarrier(
persistent-kernel) -> Tuple[utils.PersistentTileSchedulerParams, Tuple[int, int, int]]:
shared-memoryself.smem_capacity = utils.get_smem_capacity_in_bytes("sm_100")
tcgen05tcgen05.CtaGroup.TWO if self.use_2cta_instrs else tcgen05.CtaGroup.ONE
warp-specializationab_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