Skip to content
KernelIndex
Search⌘K

submission 283829

currybab · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

No package. Vendor the mirrored source: 1704 lines, June 9 Researcher Reciprocity License v1.0.

submission_4.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-dual-gemm-283829?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 dual GEMMsuite of 4 cases
NVIDIA B200
20.2µs
#173 of 420
2026-01-06

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:c96fe4aa69041a33629d2c1bd9d8e850704417eed7dc1fef8c58210c5f2013bb
license declaredunknown
license concludedunknown
authorscurrybab
imported2026-08-26

Techniques

Extracted from the mirrored source by pattern, never inferred. Each row cites its line.

fp4AB_DTYPE = cutlass.Float4E2M1FN # nvfp4
fused-epilogueEpilogue에서 accumulator(TMEM)를 register로 load하기 위한 copy/partition 생성.
mbarrierCTA_SYNC_BARRIER = pipeline.NamedBarrier(barrier_id=1, num_threads=THREADS_PER_CTA)
persistent-kernel_PERSISTENT_ACTIVE_CLUSTERS_CAP = 32
shared-memorySMEM_CAPACITY_BYTES = utils.get_smem_capacity_in_bytes("sm_100")
tcgen05A_MAJOR_MODE = tcgen05.OperandMajorMode.K
tile-k = 4MMA_INST_TILE_K = 4
warp-specializationab_pipeline_producer_group = pipeline.CooperativeGroup(pipeline.Agent.Thread)

Kernel source

submission_4.py1704 lines
from __future__ import annotations

from typing import Dict, Tuple, Union

import torch

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 task import input_t, output_t

# =============================================================================
# 개요
# =============================================================================
# 이 커널은 Blackwell(SM100, B200)에서 "block-scaled NVFP4 GEMM"을 2번 수행하고,
# epilogue에서 아래 연산을 fused로 수행합니다.
#
#   C = silu(A @ B1) * (A @ B2)
#
# - A, B1, B2 : Float4E2M1FN (nvfp4)
# - SFA, SFB1, SFB2 : Float8E4M3FN (scale factor, vec=16)
# - Accumulator : Float32 (TMEM)
# - Output C : Float16
#
# 또한 persistent tile scheduler를 사용하여, 하나의 CTA가 여러 타일을 순차적으로 처리합니다.
# Warp specialization:
#   - warp 5: TMA load 전담 (global -> shared)
#   - warp 4: MMA 전담 (shared -> TMEM(SF), UMMA 2x GEMM)
#   - warp 0~3: epilogue 전담 (TMEM(acc)->shared->global, silu+mul)
# =============================================================================


# =============================================================================
# 0) DTypes / Layout (태스크 고정 가정)
# =============================================================================
AB_DTYPE = cutlass.Float4E2M1FN        # nvfp4
SF_DTYPE = cutlass.Float8E4M3FN        # scale factor
C_DTYPE = cutlass.Float16              # output
ACC_DTYPE = cutlass.Float32            # accumulation

# 입력 A/B는 "K-major" (문제 세팅), 출력은 row-major
A_MAJOR_MODE = tcgen05.OperandMajorMode.K
B_MAJOR_MODE = tcgen05.OperandMajorMode.K
C_LAYOUT = utils.LayoutEnum.ROW_MAJOR


# =============================================================================
# 1) 튜닝 파라미터 (원본 코드에서 실사용하는 설정)
# =============================================================================
SF_VEC_SIZE = 16

# UMMA instruction MN shape (원본: (128, 64))
MMA_INST_SHAPE_MN = (128, 64)

# K 방향으로 instruction을 몇 번 반복해서 CTA 타일 K를 만들지 (tile_k = inst_k * MMA_INST_TILE_K)
# 보통 inst_k는 64이므로 64*4=256 K-타일이 됨.
MMA_INST_TILE_K = 4

# 클러스터(CTA 그룹) 형태: (cluster_m, cluster_n)
# (1,4)이면 N방향으로 4 CTA가 묶여서 A/SFA를 multicast로 공유.
CLUSTER_SHAPE_MN = (1, 4)

# persistent scheduler에서 활성 클러스터 상한
# NOTE: If this is >= total clusters, the kernel degenerates to non-persistent (1 tile per cluster),
# which hides none of the epilogue with the next tile's mainloop.
MAX_ACTIVE_CLUSTERS = 1024

# For B200-sized problems in this task, limiting active clusters can enable real persistence
# (multiple tiles per cluster) without hurting small-M cases too much.
_PERSISTENT_ACTIVE_CLUSTERS_CAP = 32

# compile opt-level
COMPILE_OPT_LEVEL = 3

# 2CTA-instruction 경로는 MMA_INST_SHAPE_MN[0]==256 일 때만 사용 (여기선 False)
USE_2CTA_INSTRS = MMA_INST_SHAPE_MN[0] == 256
CTA_GROUP = tcgen05.CtaGroup.TWO if USE_2CTA_INSTRS else tcgen05.CtaGroup.ONE

# Warp specialization IDs
EPILOG_WARPS = (0, 1, 2, 3)
MMA_WARP_ID = 4
TMA_WARP_ID = 5

# CTA thread 수 (warp 6개 사용)
THREADS_PER_CTA = 32 * len((MMA_WARP_ID, TMA_WARP_ID, *EPILOG_WARPS))

# Named barriers (CTA 내부 동기화)
CTA_SYNC_BARRIER = pipeline.NamedBarrier(barrier_id=1, num_threads=THREADS_PER_CTA)
EPILOG_SYNC_BARRIER = pipeline.NamedBarrier(
    barrier_id=2, num_threads=32 * len(EPILOG_WARPS)
)
TMEM_ALLOC_BARRIER = pipeline.NamedBarrier(
    barrier_id=3, num_threads=32 * len((MMA_WARP_ID, *EPILOG_WARPS))
)

# HW resource info
SMEM_CAPACITY_BYTES = utils.get_smem_capacity_in_bytes("sm_100")
NUM_TMEM_ALLOC_COLS = 512              # TMEM allocator에 요청할 column 수 (원본과 동일)
BUFFER_ALIGN_BYTES = 1024              # shared buffer 정렬


# =============================================================================
# 2) Helper 함수들 (원본 메서드를 "함수형"으로 옮겨둔 것)
# =============================================================================

def _compute_stages(
    tiled_mma: cute.TiledMma,
    mma_tiler_mnk: Tuple[int, int, int],
    a_dtype,
    b_dtype,
    epi_tile: cute.Tile,
    c_dtype,
    c_layout: utils.LayoutEnum,
    sf_dtype,
    sf_vec_size: int,
    smem_capacity: int,
    occupancy: int,
) -> Tuple[int, int, int]:
    """
    shared memory capacity에서 pipeline stage 수(AB stage, C stage) 등을 계산.
    dual-GEMM이라 B/SFB가 2번 로드되는 것까지 고려해서 stage 수를 산정.
    """
    # ACC stage heuristic (원본)
    num_acc_stage = 1 if mma_tiler_mnk[1] == 256 else 2
    num_c_stage = 2

    # 1-stage 기준으로 각 buffer가 얼마나 shared memory를 먹는지 계산
    a_smem_layout_stage_one = sm100_utils.make_smem_layout_a(
        tiled_mma, mma_tiler_mnk, a_dtype, 1
    )
    b_smem_layout_stage_one = sm100_utils.make_smem_layout_b(
        tiled_mma, mma_tiler_mnk, b_dtype, 1
    )
    sfa_smem_layout_stage_one = blockscaled_utils.make_smem_layout_sfa(
        tiled_mma, mma_tiler_mnk, sf_vec_size, 1
    )
    sfb_smem_layout_stage_one = blockscaled_utils.make_smem_layout_sfb(
        tiled_mma, mma_tiler_mnk, sf_vec_size, 1
    )
    c_smem_layout_stage_one = sm100_utils.make_smem_layout_epi(
        c_dtype, c_layout, epi_tile, 1
    )

    # dual GEMM: B와 SFB가 2개(B1,B2 / SFB1,SFB2)라서 2배
    ab_bytes_per_stage = (
        cute.size_in_bytes(a_dtype, a_smem_layout_stage_one)
        + 2 * cute.size_in_bytes(b_dtype, b_smem_layout_stage_one)
        + cute.size_in_bytes(sf_dtype, sfa_smem_layout_stage_one)
        + 2 * cute.size_in_bytes(sf_dtype, sfb_smem_layout_stage_one)
    )

    # barrier/helper 공간(대략)
    mbar_helpers_bytes = 1024

    c_bytes_per_stage = cute.size_in_bytes(c_dtype, c_smem_layout_stage_one)
    c_bytes = c_bytes_per_stage * num_c_stage

    # AB stage 수 계산
    num_ab_stage = (
        smem_capacity // occupancy - (mbar_helpers_bytes + c_bytes)
    ) // ab_bytes_per_stage

    # 0-stage 방지
    if num_ab_stage < 1:
        num_ab_stage = 1

    # 남는 shared로 C stage를 더 늘릴 수 있으면 늘림
    num_c_stage += (
        smem_capacity
        - occupancy * ab_bytes_per_stage * num_ab_stage
        - occupancy * (mbar_helpers_bytes + c_bytes)
    ) // (occupancy * c_bytes_per_stage)

    if num_c_stage < 2:
        num_c_stage = 2

    return num_acc_stage, num_ab_stage, num_c_stage


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]]:
    """
    persistent scheduler를 위한 grid shape 계산.
    - C를 CTA 타일 크기로 쪼갠 뒤, 타일 개수 및 cluster shape를 파라미터로 전달.
    """
    c_shape = cute.slice_(cta_tile_shape_mnk, (None, None, 0))  # (Mtile, Ntile)
    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


def _mainloop_s2t_copy_and_partition(
    sSF: cute.Tensor,
    tSF: cute.Tensor,
    sf_dtype=SF_DTYPE,
    cta_group=CTA_GROUP,
) -> Tuple[cute.TiledCopy, cute.Tensor, cute.Tensor]:
    """
    Scale factor를 SMEM -> TMEM 으로 옮길 때 쓰는 copy 객체/파티션 생성.
    - sSF: shared memory tensor
    - tSF: tensor memory tensor
    반환:
      tiled_copy_s2t, (SMEM desc tensor), (TMEM dest partition)
    """
    # filter_zeros는 "의미없는(0 stride 등) 축"을 제거해서 더 compact한 view를 만듦
    tCsSF_compact = cute.filter_zeros(sSF)
    tCtSF_compact = cute.filter_zeros(tSF)

    copy_atom_s2t = cute.make_copy_atom(
        tcgen05.Cp4x32x128bOp(cta_group),
        sf_dtype,
    )
    tiled_copy_s2t = tcgen05.make_s2t_copy(copy_atom_s2t, tCtSF_compact)
    thr_copy_s2t = tiled_copy_s2t.get_slice(0)  # warp 단위 copy라 slice=0 사용

    # source partition (shared)
    tCsSF_compact_s2t_ = thr_copy_s2t.partition_S(tCsSF_compact)

    # tcgen05 S2T는 "shared descriptor tensor"로 변환이 필요
    tCsSF_compact_s2t = tcgen05.get_s2t_smem_desc_tensor(
        tiled_copy_s2t, tCsSF_compact_s2t_
    )

    # dest partition (tmem)
    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(
    tidx: cutlass.Int32,
    tAcc: cute.Tensor,
    gC_mnl: cute.Tensor,
    epi_tile: cute.Tile,
    use_2cta_instrs: Union[cutlass.Boolean, bool],
    cta_tile_shape_mnk: Tuple[int, int, int],
    c_layout: utils.LayoutEnum,
    c_dtype,
    acc_dtype=ACC_DTYPE,
) -> Tuple[cute.TiledCopy, cute.Tensor, cute.Tensor]:
    """
    Epilogue에서 accumulator(TMEM)를 register로 load하기 위한 copy/partition 생성.
    """
    # 타일/레이아웃에 맞는 TMEM load op 선택
    copy_atom_t2r = sm100_utils.get_tmem_load_op(
        cta_tile_shape_mnk,
        c_layout,
        c_dtype,
        acc_dtype,
        epi_tile,
        use_2cta_instrs,
    )

    # accumulator를 epilogue tile 단위로 나눈 view
    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)]
    )

    # thread slice에 맞게 partition
    thr_copy_t2r = tiled_copy_t2r.get_slice(tidx)
    tTR_tAcc = thr_copy_t2r.partition_S(tAcc_epi)

    # destination은 gC에 맞게 partition (shape만 참조 용도)
    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)

    # 실제 register buffer
    tTR_rAcc = cute.make_rmem_tensor(
        tTR_gC[(None, None, None, 0, 0, 0, 0, 0)].shape, acc_dtype
    )
    return tiled_copy_t2r, tTR_tAcc, tTR_rAcc


def _epilog_smem_copy_and_partition(
    tiled_copy_t2r: cute.TiledCopy,
    tTR_rC: cute.Tensor,
    tidx: cutlass.Int32,
    sC: cute.Tensor,
    c_layout: utils.LayoutEnum,
    c_dtype,
    acc_dtype=ACC_DTYPE,
) -> Tuple[cute.TiledCopy, cute.Tensor, cute.Tensor]:
    """
    Epilogue에서 register -> shared store 파티셔닝 생성.
    """
    copy_atom_r2s = sm100_utils.get_smem_store_op(
        c_layout, c_dtype, 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)         # shared destination partition
    tRS_rC = tiled_copy_r2s.retile(tTR_rC)        # register source retile
    return tiled_copy_r2s, tRS_rC, tRS_sC


def _epilog_gmem_copy_and_partition(
    atom: Union[cute.CopyAtom, cute.TiledCopy],
    gC_mnl: cute.Tensor,
    epi_tile: cute.Tile,
    sC: cute.Tensor,
) -> Tuple[cute.CopyAtom, cute.Tensor, cute.Tensor]:
    """
    shared -> global (TMA store) partition 생성.
    """
    gC_epi = cute.flat_divide(
        gC_mnl[((None, None), 0, 0, None, None, None)], epi_tile
    )

    # TMA partition을 위해 rank를 맞춰준다
    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(
        atom,            # tma store atom
        0,
        cute.make_layout(1),
        sC_for_tma_partition,
        gC_for_tma_partition,
    )
    return atom, bSG_sC, bSG_gC


# =============================================================================
# 3) Device kernel (@cute.kernel)
# =============================================================================
@cute.kernel
def dual_gemm_silu_kernel(
    # MMA objects
    tiled_mma: cute.TiledMma,
    tiled_mma_sfb: cute.TiledMma,

    # TMA atoms + global tensors (A/B/SF)
    tma_atom_a: cute.CopyAtom,
    mA_mkl: cute.Tensor,
    tma_atom_b1: cute.CopyAtom,
    mB1_nkl: cute.Tensor,
    tma_atom_b2: cute.CopyAtom,
    mB2_nkl: cute.Tensor,
    tma_atom_sfa: cute.CopyAtom,
    mSFA_mkl: cute.Tensor,
    tma_atom_sfb1: cute.CopyAtom,
    mSFB1_nkl: cute.Tensor,
    tma_atom_sfb2: cute.CopyAtom,
    mSFB2_nkl: cute.Tensor,

    # TMA store atom + output tensor C
    tma_atom_c: cute.CopyAtom,
    mC_mnl: cute.Tensor,

    # cluster layout (for multicast + persistent scheduling)
    cluster_layout_vmnk: cute.Layout,
    cluster_layout_sfb_vmnk: cute.Layout,

    # staged shared memory layouts
    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],

    # epilogue tile + scheduler params
    epi_tile: cute.Tile,
    tile_sched_params: utils.PersistentTileSchedulerParams,

    # constexpr params (compile-time constants)
    mma_tiler: cutlass.Constexpr[Tuple[int, int, int]],
    mma_tiler_sfb: cutlass.Constexpr[Tuple[int, int, int]],
    cta_tile_shape_mnk: cutlass.Constexpr[Tuple[int, int, int]],
    num_acc_stage: cutlass.Constexpr[int],
    num_ab_stage: cutlass.Constexpr[int],
    num_c_stage: cutlass.Constexpr[int],
    num_tma_load_bytes: cutlass.Constexpr[int],

    # epilogue op: silu(x) = x / (1+exp(-x))
    epilogue_op: cutlass.Constexpr = lambda x: x
    * (1.0 / (1.0 + cute.math.exp(-x, fastmath=True))),
):
    # -------------------------------------------------------------------------
    # Warp ID 확인
    # -------------------------------------------------------------------------
    warp_idx = cute.arch.warp_idx()
    warp_idx = cute.arch.make_warp_uniform(warp_idx)

    # -------------------------------------------------------------------------
    # (TMA warp) TMA descriptor prefetch
    #   - TMA load/store를 빠르게 하기 위해 descriptor를 L1에 prefetch
    # -------------------------------------------------------------------------
    if warp_idx == TMA_WARP_ID:
        cpasync.prefetch_descriptor(tma_atom_a)
        cpasync.prefetch_descriptor(tma_atom_b1)
        cpasync.prefetch_descriptor(tma_atom_b2)
        cpasync.prefetch_descriptor(tma_atom_sfa)
        cpasync.prefetch_descriptor(tma_atom_sfb1)
        cpasync.prefetch_descriptor(tma_atom_sfb2)
        cpasync.prefetch_descriptor(tma_atom_c)

    # 2CTA instruction 여부 (여기서는 거의 False)
    use_2cta_instrs = cute.size(tiled_mma.thr_id.shape) == 2

    # -------------------------------------------------------------------------
    # Block/Cluster 좌표
    # -------------------------------------------------------------------------
    bidx, bidy, bidz = cute.arch.block_idx()

    # cluster 내 v 좌표 (2CTA일 때 v=0/1)
    mma_tile_coord_v = bidx % cute.size(tiled_mma.thr_id.shape)
    is_leader_cta = mma_tile_coord_v == 0  # 2CTA에서 리더 CTA만 MMA/epilogue를 하게 만들 때 사용

    # cluster 내 CTA rank
    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()[0]

    # -------------------------------------------------------------------------
    # Shared storage:
    #   - barrier 및 TMEM allocator 상태만 struct로 잡고,
    #   - 실제 sA/sB/sSF/sC 버퍼는 SmemAllocator로 따로 allocate.
    # -------------------------------------------------------------------------
    @cute.struct
    class SharedStorage:
        # AB pipeline barriers (full/empty)
        ab_full_mbar_ptr: cute.struct.MemRange[cutlass.Int64, num_ab_stage]
        ab_empty_mbar_ptr: cute.struct.MemRange[cutlass.Int64, num_ab_stage]

        # ACC pipeline barriers (full/empty)
        acc_full_mbar_ptr: cute.struct.MemRange[cutlass.Int64, num_acc_stage]
        acc_empty_mbar_ptr: cute.struct.MemRange[cutlass.Int64, num_acc_stage]

        # 2CTA용 dealloc barrier (원본 호환)
        tmem_dealloc_mbar_ptr: cutlass.Int64

        # TMEM allocator 상태 공유 버퍼
        tmem_holding_buf: cutlass.Int32

    smem = utils.SmemAllocator()
    storage = smem.allocate(SharedStorage)

    # Shared memory buffers (staged layouts)
    #   sA: (tileM, tileK, stage)
    #   sB1/sB2: (tileN, tileK, stage)
    #   sSFA/sSFB: scale factor blocks
    #   sC: epilogue intermediate buffer (TMA store용)
    sC = smem.allocate_tensor(
        element_type=C_DTYPE,
        layout=c_smem_layout_staged.outer,
        byte_alignment=BUFFER_ALIGN_BYTES,
        swizzle=c_smem_layout_staged.inner,
    )
    sA = smem.allocate_tensor(
        element_type=AB_DTYPE,
        layout=a_smem_layout_staged.outer,
        byte_alignment=BUFFER_ALIGN_BYTES,
        swizzle=a_smem_layout_staged.inner,
    )
    sB1 = smem.allocate_tensor(
        element_type=AB_DTYPE,
        layout=b_smem_layout_staged.outer,
        byte_alignment=BUFFER_ALIGN_BYTES,
        swizzle=b_smem_layout_staged.inner,
    )
    sB2 = smem.allocate_tensor(
        element_type=AB_DTYPE,
        layout=b_smem_layout_staged.outer,
        byte_alignment=BUFFER_ALIGN_BYTES,
        swizzle=b_smem_layout_staged.inner,
    )
    sSFA = smem.allocate_tensor(
        element_type=SF_DTYPE,
        layout=sfa_smem_layout_staged,
        byte_alignment=BUFFER_ALIGN_BYTES,
    )
    sSFB1 = smem.allocate_tensor(
        element_type=SF_DTYPE,
        layout=sfb_smem_layout_staged,
        byte_alignment=BUFFER_ALIGN_BYTES,
    )
    sSFB2 = smem.allocate_tensor(
        element_type=SF_DTYPE,
        layout=sfb_smem_layout_staged,
        byte_alignment=BUFFER_ALIGN_BYTES,
    )

    # -------------------------------------------------------------------------
    # Pipeline 초기화
    # -------------------------------------------------------------------------
    # multicast CTA 개수:
    # - A/SFA는 cluster N방향으로 공유되므로, "mcast_ctas_a"가 중요
    num_mcast_ctas_a = cute.size(cluster_layout_vmnk.shape[2])
    num_mcast_ctas_b = cute.size(cluster_layout_vmnk.shape[1])
    num_mcast_ctas_sfb = cute.size(cluster_layout_sfb_vmnk.shape[1])

    is_a_mcast = num_mcast_ctas_a > 1
    is_b_mcast = num_mcast_ctas_b > 1
    is_sfb_mcast = num_mcast_ctas_sfb > 1

    # AB pipeline:
    #   Producer(TMA load)가 shared에 데이터를 넣고, Consumer(MMA)가 사용할 수 있게 barrier로 동기화.
    ab_pipeline_producer_group = pipeline.CooperativeGroup(pipeline.Agent.Thread)

    # TMA producer 수는 multicast 조합에 의해 결정(원본 공식)
    num_tma_producer = num_mcast_ctas_a + num_mcast_ctas_b - 1
    ab_pipeline_consumer_group = pipeline.CooperativeGroup(
        pipeline.Agent.Thread, num_tma_producer
    )

    ab_pipeline = pipeline.PipelineTmaUmma.create(
        barrier_storage=storage.ab_full_mbar_ptr.data_ptr(),
        num_stages=num_ab_stage,
        producer_group=ab_pipeline_producer_group,
        consumer_group=ab_pipeline_consumer_group,
        tx_count=num_tma_load_bytes,         # TMA 완료를 byte-count로 감지
        cta_layout_vmnk=cluster_layout_vmnk, # multicast mask 생성 등에 사용
    )

    # ACC pipeline:
    #   MMA warp가 accumulator 결과를 "완료"로 표시하면 epilogue warps가 이를 consume.
    acc_pipeline_producer_group = pipeline.CooperativeGroup(pipeline.Agent.Thread)
    num_acc_consumer_threads = len(EPILOG_WARPS) * (2 if use_2cta_instrs else 1)
    acc_pipeline_consumer_group = pipeline.CooperativeGroup(
        pipeline.Agent.Thread, num_acc_consumer_threads
    )

    acc_pipeline = pipeline.PipelineUmmaAsync.create(
        barrier_storage=storage.acc_full_mbar_ptr.data_ptr(),
        num_stages=num_acc_stage,
        producer_group=acc_pipeline_producer_group,
        consumer_group=acc_pipeline_consumer_group,
        cta_layout_vmnk=cluster_layout_vmnk,
    )

    # TMEM allocator:
    #   - epilogue warp0가 allocate 수행
    #   - MMA warp는 wait_for_alloc 후 retrieve_ptr로 사용
    tmem = utils.TmemAllocator(
        storage.tmem_holding_buf,
        barrier_for_retrieve=TMEM_ALLOC_BARRIER,
        allocator_warp_id=EPILOG_WARPS[0],
        is_two_cta=use_2cta_instrs,
        two_cta_tmem_dealloc_mbar_ptr=storage.tmem_dealloc_mbar_ptr,
    )

    # cluster arrive: 여러 CTA가 협업하는 경우 relaxed arrive 필요
    if cute.size(CLUSTER_SHAPE_MN) > 1:
        cute.arch.cluster_arrive_relaxed()

    # -------------------------------------------------------------------------
    # Multicast mask 생성 (A/B/SF를 cluster 내에 multicast 로드할 때 필요)
    # -------------------------------------------------------------------------
    a_full_mcast_mask = None
    b_full_mcast_mask = None
    sfa_full_mcast_mask = None
    sfb_full_mcast_mask = None

    # A/SFA는 mcast_mode=2, B는 mcast_mode=1 (원본 패턴)
    if cutlass.const_expr(is_a_mcast or 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,
        )

    # -------------------------------------------------------------------------
    # CTA 타일 단위로 global tensor를 local_tile로 자름
    # -------------------------------------------------------------------------
    # (A/B는 (m,k,l)/(n,k,l), C는 (m,n,l))
    gA_mkl = cute.local_tile(
        mA_mkl, cute.slice_(mma_tiler, (None, 0, None)), (None, None, None)
    )
    gB1_nkl = cute.local_tile(
        mB1_nkl, cute.slice_(mma_tiler, (0, None, None)), (None, None, None)
    )
    gB2_nkl = cute.local_tile(
        mB2_nkl, cute.slice_(mma_tiler, (0, None, None)), (None, None, None)
    )
    gSFA_mkl = cute.local_tile(
        mSFA_mkl, cute.slice_(mma_tiler, (None, 0, None)), (None, None, None)
    )
    gSFB1_nkl = cute.local_tile(
        mSFB1_nkl, cute.slice_(mma_tiler_sfb, (0, None, None)), (None, None, None)
    )
    gSFB2_nkl = cute.local_tile(
        mSFB2_nkl, cute.slice_(mma_tiler_sfb, (0, None, None)), (None, None, None)
    )
    gC_mnl = cute.local_tile(
        mC_mnl, cute.slice_(mma_tiler, (None, None, 0)), (None, None, None)
    )

    # K 타일 개수
    k_tile_cnt = cute.size(gA_mkl, mode=[3])

    # -------------------------------------------------------------------------
    # MMA partition: 각 CTA slice가 담당하는 부분을 얻음
    # -------------------------------------------------------------------------
    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)
    tCgB1 = thr_mma.partition_B(gB1_nkl)
    tCgB2 = thr_mma.partition_B(gB2_nkl)
    tCgSFA = thr_mma.partition_A(gSFA_mkl)
    tCgSFB1 = thr_mma_sfb.partition_B(gSFB1_nkl)
    tCgSFB2 = thr_mma_sfb.partition_B(gSFB2_nkl)
    tCgC = thr_mma.partition_C(gC_mnl)

    # -------------------------------------------------------------------------
    # TMA partition: (global tensor) <-> (shared buffer) 매칭되는 view 생성
    # -------------------------------------------------------------------------
    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 multicast dimension
        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
    )
    tBsB1, tBgB1 = cpasync.tma_partition(
        tma_atom_b1,
        block_in_cluster_coord_vmnk[1],      # B multicast dimension
        b_cta_layout,
        cute.group_modes(sB1, 0, 3),
        cute.group_modes(tCgB1, 0, 3),
    )
    tBsB2, tBgB2 = cpasync.tma_partition(
        tma_atom_b2,
        block_in_cluster_coord_vmnk[1],
        b_cta_layout,
        cute.group_modes(sB2, 0, 3),
        cute.group_modes(tCgB2, 0, 3),
    )

    # scale factors는 internal_type(Int16) 경로라 cute.nvgpu.cpasync.tma_partition 사용
    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
    )
    tBsSFB1, tBgSFB1 = cute.nvgpu.cpasync.tma_partition(
        tma_atom_sfb1,
        block_in_cluster_coord_sfb_vmnk[1],
        sfb_cta_layout,
        cute.group_modes(sSFB1, 0, 3),
        cute.group_modes(tCgSFB1, 0, 3),
    )
    tBsSFB1 = cute.filter_zeros(tBsSFB1)
    tBgSFB1 = cute.filter_zeros(tBgSFB1)

    tBsSFB2, tBgSFB2 = cute.nvgpu.cpasync.tma_partition(
        tma_atom_sfb2,
        block_in_cluster_coord_sfb_vmnk[1],
        sfb_cta_layout,
        cute.group_modes(sSFB2, 0, 3),
        cute.group_modes(tCgSFB2, 0, 3),
    )
    tBsSFB2 = cute.filter_zeros(tBsSFB2)
    tBgSFB2 = cute.filter_zeros(tBgSFB2)

    # -------------------------------------------------------------------------
    # shared buffer를 UMMA fragment view로 변환
    # -------------------------------------------------------------------------
    tCrA = tiled_mma.make_fragment_A(sA)
    tCrB1 = tiled_mma.make_fragment_B(sB1)
    tCrB2 = tiled_mma.make_fragment_B(sB2)

    # accumulator 텐서의 shape만 만든 fake fragment (실제 포인터는 TMEM allocate 후 얻음)
    acc_shape = tiled_mma.partition_shape_C(mma_tiler[:2])
    tCtAcc_fake = tiled_mma.make_fragment_C(
        cute.append(acc_shape, num_acc_stage)
    )

    # -------------------------------------------------------------------------
    # cluster sync: shared buffer/pipe 초기화 이후 동기화
    # -------------------------------------------------------------------------
    if cute.size(CLUSTER_SHAPE_MN) > 1:
        cute.arch.cluster_wait()
    else:
        CTA_SYNC_BARRIER.arrive_and_wait()

    # =========================================================================
    # (1) TMA warp: persistent scheduler로 GMEM -> SMEM 로드
    # =========================================================================
    if warp_idx == TMA_WARP_ID:
        tile_sched = utils.StaticPersistentTileScheduler.create(
            tile_sched_params, cute.arch.block_idx(), cute.arch.grid_dim()
        )
        work_tile = tile_sched.initial_work_tile_info()

        ab_producer_state = pipeline.make_pipeline_state(
            pipeline.PipelineUserType.Producer, num_ab_stage
        )

        while work_tile.is_valid_tile:
            # work tile index는 (cta_linear, n_tile, l) 같은 형태로 들어옴
            cur_tile_coord = work_tile.tile_idx

            # mma_tile_coord_mnl: (m_tile, n_tile, l)
            mma_tile_coord_mnl = (
                cur_tile_coord[0] // cute.size(tiled_mma.thr_id.shape),
                cur_tile_coord[1],
                cur_tile_coord[2],
            )

            # 현재 타일의 global slice를 선택
            tAgA_slice = tAgA[
                (None, mma_tile_coord_mnl[0], None, mma_tile_coord_mnl[2])
            ]
            tBgB1_slice = tBgB1[
                (None, mma_tile_coord_mnl[1], None, mma_tile_coord_mnl[2])
            ]
            tBgB2_slice = tBgB2[
                (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])
            ]

            # N=64 타일일 경우 SFB는 padded N=128이므로 slice_n 보정이 필요
            slice_n = mma_tile_coord_mnl[1]
            if cutlass.const_expr(cta_tile_shape_mnk[1] == 64):
                slice_n = mma_tile_coord_mnl[1] // 2

            tBgSFB1_slice = tBgSFB1[(None, slice_n, None, mma_tile_coord_mnl[2])]
            tBgSFB2_slice = tBgSFB2[(None, slice_n, None, mma_tile_coord_mnl[2])]

            # pipeline stage 획득 루프
            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
                )

            # k_tile_cnt 만큼 GMEM->SMEM 로드 (pipeline stage 회전)
            for k_tile_iter in cutlass.range(0, k_tile_cnt, 1, unroll=1):
                ab_pipeline.producer_acquire(
                    ab_producer_state, peek_ab_empty_status
                )
                bar = ab_pipeline.producer_get_barrier(ab_producer_state)

                # TMA load: A/B1/B2/SFA/SFB1/SFB2
                cute.copy(
                    tma_atom_a,
                    tAgA_slice[(None, ab_producer_state.count)],
                    tAsA[(None, ab_producer_state.index)],
                    tma_bar_ptr=bar,
                    mcast_mask=a_full_mcast_mask,
                )
                cute.copy(
                    tma_atom_b1,
                    tBgB1_slice[(None, ab_producer_state.count)],
                    tBsB1[(None, ab_producer_state.index)],
                    tma_bar_ptr=bar,
                    mcast_mask=b_full_mcast_mask,
                )
                cute.copy(
                    tma_atom_b2,
                    tBgB2_slice[(None, ab_producer_state.count)],
                    tBsB2[(None, ab_producer_state.index)],
                    tma_bar_ptr=bar,
                    mcast_mask=b_full_mcast_mask,
                )
                cute.copy(
                    tma_atom_sfa,
                    tAgSFA_slice[(None, ab_producer_state.count)],
                    tAsSFA[(None, ab_producer_state.index)],
                    tma_bar_ptr=bar,
                    mcast_mask=sfa_full_mcast_mask,
                )
                cute.copy(
                    tma_atom_sfb1,
                    tBgSFB1_slice[(None, ab_producer_state.count)],
                    tBsSFB1[(None, ab_producer_state.index)],
                    tma_bar_ptr=bar,
                    mcast_mask=sfb_full_mcast_mask,
                )
                cute.copy(
                    tma_atom_sfb2,
                    tBgSFB2_slice[(None, ab_producer_state.count)],
                    tBsSFB2[(None, ab_producer_state.index)],
                    tma_bar_ptr=bar,
                    mcast_mask=sfb_full_mcast_mask,
                )

                # 다음 stage로 advance
                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
                    )

            # 다음 work tile로 이동
            tile_sched.advance_to_next_work()
            work_tile = tile_sched.get_current_work()

        # producer tail (남은 stage 정리)
        ab_pipeline.producer_tail(ab_producer_state)

    # =========================================================================
    # (2) MMA warp: SMEM->TMEM(SF) copy + UMMA dual GEMM
    # =========================================================================
    if warp_idx == MMA_WARP_ID:
        # epilogue warp0가 TMEM allocate 끝낼 때까지 기다림
        tmem.wait_for_alloc()

        # accumulator TMEM base ptr
        acc_tmem_ptr = tmem.retrieve_ptr(ACC_DTYPE)
        tCtAcc1_base = cute.make_tensor(acc_tmem_ptr, tCtAcc_fake.layout)

        # acc2는 acc1 다음 column offset에 배치
        acc2_ptr = acc_tmem_ptr + tcgen05.find_tmem_tensor_col_offset(tCtAcc1_base)
        tCtAcc2_base = cute.make_tensor(acc2_ptr, tCtAcc_fake.layout)

        # TMEM layout for scale factors
        tCtSFA_layout = blockscaled_utils.make_tmem_layout_sfa(
            tiled_mma,
            mma_tiler,
            SF_VEC_SIZE,
            cute.slice_(sfa_smem_layout_staged, (None, None, None, 0)),
        )
        tCtSFB_layout = blockscaled_utils.make_tmem_layout_sfb(
            tiled_mma,
            mma_tiler,
            SF_VEC_SIZE,
            cute.slice_(sfb_smem_layout_staged, (None, None, None, 0)),
        )

        # TMEM column offsets 계산
        acc1_cols = tcgen05.find_tmem_tensor_col_offset(tCtAcc1_base)
        acc2_cols = tcgen05.find_tmem_tensor_col_offset(tCtAcc2_base)

        # SFA는 acc1+acc2 다음에 배치
        sfa_ptr_base = cute.recast_ptr(
            acc_tmem_ptr + acc1_cols + acc2_cols, dtype=SF_DTYPE
        )
        tCtSFA = cute.make_tensor(sfa_ptr_base, tCtSFA_layout)

        # SFB1은 SFA 다음에 배치
        sfb1_ptr_base = cute.recast_ptr(
            acc_tmem_ptr
            + acc1_cols
            + acc2_cols
            + tcgen05.find_tmem_tensor_col_offset(tCtSFA),
            dtype=SF_DTYPE,
        )
        tCtSFB1 = cute.make_tensor(sfb1_ptr_base, tCtSFB_layout)

        # SFB2는 SFB1 다음에 배치
        sfb2_ptr_base = cute.recast_ptr(
            acc_tmem_ptr
            + acc1_cols
            + acc2_cols
            + tcgen05.find_tmem_tensor_col_offset(tCtSFA)
            + tcgen05.find_tmem_tensor_col_offset(tCtSFB1),
            dtype=SF_DTYPE,
        )
        tCtSFB2 = cute.make_tensor(sfb2_ptr_base, tCtSFB_layout)

        # SMEM -> TMEM copy partitions for SFA/SFB1/SFB2
        tiled_copy_s2t_sfa, tCsSFA_compact_s2t, tCtSFA_compact_s2t = (
            _mainloop_s2t_copy_and_partition(sSFA, tCtSFA)
        )
        tiled_copy_s2t_sfb1, tCsSFB1_compact_s2t, tCtSFB1_compact_s2t = (
            _mainloop_s2t_copy_and_partition(sSFB1, tCtSFB1)
        )
        tiled_copy_s2t_sfb2, tCsSFB2_compact_s2t, tCtSFB2_compact_s2t = (
            _mainloop_s2t_copy_and_partition(sSFB2, tCtSFB2)
        )

        # persistent scheduler 생성
        tile_sched = utils.StaticPersistentTileScheduler.create(
            tile_sched_params, cute.arch.block_idx(), cute.arch.grid_dim()
        )
        work_tile = tile_sched.initial_work_tile_info()

        # AB pipeline consumer state (SMEM stage를 기다렸다가 사용)
        ab_consumer_state = pipeline.make_pipeline_state(
            pipeline.PipelineUserType.Consumer, num_ab_stage
        )
        # ACC pipeline producer state (MMA 결과 stage commit)
        acc_producer_state = pipeline.make_pipeline_state(
            pipeline.PipelineUserType.Producer, num_acc_stage
        )

        while work_tile.is_valid_tile:
            cur_tile_coord = work_tile.tile_idx
            mma_tile_coord_mnl = (
                cur_tile_coord[0] // cute.size(tiled_mma.thr_id.shape),
                cur_tile_coord[1],
                cur_tile_coord[2],
            )

            # 현재 acc stage 선택
            tCtAcc1 = tCtAcc1_base[(None, None, None, acc_producer_state.index)]
            tCtAcc2 = tCtAcc2_base[(None, None, None, acc_producer_state.index)]

            # AB consumer wait 준비
            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
                )

            # leader CTA만 acc_pipeline producer acquire (2CTA 대응)
            if is_leader_cta:
                acc_pipeline.producer_acquire(acc_producer_state)

            # N=64일 때 padded SFB를 TMEM에서 half-tile로 읽기 위한 pointer shift hack
            tCtSFB1_mma = tCtSFB1
            tCtSFB2_mma = tCtSFB2
            if cutlass.const_expr(cta_tile_shape_mnk[1] == 64):
                offset = cutlass.Int32((mma_tile_coord_mnl[1] % 2) * 2)

                sfb1_cols = (
                    acc1_cols
                    + acc2_cols
                    + tcgen05.find_tmem_tensor_col_offset(tCtSFA)
                )
                sfb2_cols = sfb1_cols + tcgen05.find_tmem_tensor_col_offset(
                    tCtSFB1
                )

                shifted_ptr1 = cute.recast_ptr(
                    acc_tmem_ptr + sfb1_cols + offset,
                    dtype=SF_DTYPE,
                )
                shifted_ptr2 = cute.recast_ptr(
                    acc_tmem_ptr + sfb2_cols + offset,
                    dtype=SF_DTYPE,
                )
                tCtSFB1_mma = cute.make_tensor(shifted_ptr1, tCtSFB_layout)
                tCtSFB2_mma = cute.make_tensor(shifted_ptr2, tCtSFB_layout)

            # tile 시작: accumulate false로 reset (첫 kblock은 overwrite)
            tiled_mma.set(tcgen05.Field.ACCUMULATE, False)

            # UMMA 내부 kblock 개수 (SMEM fragment에서 kblock 축)
            num_kblocks = cute.size(tCrA, mode=[2])

            # k_tile 루프: 각 타일(K chunk)에 대해
            for k_tile in range(k_tile_cnt):
                if is_leader_cta:
                    # shared stage full 기다림
                    ab_pipeline.consumer_wait(
                        ab_consumer_state, peek_ab_full_status
                    )

                    # scale factor를 shared -> tmem으로 copy
                    s2t_stage_coord = (
                        None,
                        None,
                        None,
                        None,
                        ab_consumer_state.index,
                    )
                    cute.copy(
                        tiled_copy_s2t_sfa,
                        tCsSFA_compact_s2t[s2t_stage_coord],
                        tCtSFA_compact_s2t,
                    )
                    cute.copy(
                        tiled_copy_s2t_sfb1,
                        tCsSFB1_compact_s2t[s2t_stage_coord],
                        tCtSFB1_compact_s2t,
                    )
                    cute.copy(
                        tiled_copy_s2t_sfb2,
                        tCsSFB2_compact_s2t[s2t_stage_coord],
                        tCtSFB2_compact_s2t,
                    )

                    # UMMA kblock 루프
                    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)

                        # SFA/SFB iterator 설정 (block-scaled UMMA)
                        tiled_mma.set(
                            tcgen05.Field.SFA,
                            tCtSFA[sf_kblock_coord].iterator,
                        )

                        # GEMM1: Acc1 += A@B1
                        tiled_mma.set(
                            tcgen05.Field.SFB,
                            tCtSFB1_mma[sf_kblock_coord].iterator,
                        )
                        cute.gemm(
                            tiled_mma,
                            tCtAcc1,
                            tCrA[kblock_coord],
                            tCrB1[kblock_coord],
                            tCtAcc1,
                        )

                        # GEMM2: Acc2 += A@B2
                        tiled_mma.set(
                            tcgen05.Field.SFB,
                            tCtSFB2_mma[sf_kblock_coord].iterator,
                        )
                        cute.gemm(
                            tiled_mma,
                            tCtAcc2,
                            tCrA[kblock_coord],
                            tCrB2[kblock_coord],
                            tCtAcc2,
                        )

                        # 첫 kblock 이후부터 accumulate true
                        if k_tile == 0 and kblock_idx == 0:
                            tiled_mma.set(tcgen05.Field.ACCUMULATE, True)

                    # AB stage release
                    ab_pipeline.consumer_release(ab_consumer_state)

                # 다음 AB stage로 advance
                ab_consumer_state.advance()
                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
                    )

            # MMA 결과가 준비됨을 ACC pipeline에 commit
            if is_leader_cta:
                acc_pipeline.producer_commit(acc_producer_state)
            acc_producer_state.advance()

            tile_sched.advance_to_next_work()
            work_tile = tile_sched.get_current_work()

        # producer tail
        acc_pipeline.producer_tail(acc_producer_state)

    # =========================================================================
    # (3) Epilogue warps: TMEM(acc) -> (silu+mul) -> SMEM -> TMA store -> GMEM
    # =========================================================================
    if warp_idx < MMA_WARP_ID:
        # TMEM allocate (warp0만 실제 allocate 수행, 나머지는 내부적으로 동기화 참여)
        tmem.allocate(NUM_TMEM_ALLOC_COLS)
        tmem.wait_for_alloc()

        # TMEM accumulator base ptr
        acc_tmem_ptr = tmem.retrieve_ptr(ACC_DTYPE)
        tCtAcc1_base = cute.make_tensor(acc_tmem_ptr, tCtAcc_fake.layout)
        acc2_ptr = acc_tmem_ptr + tcgen05.find_tmem_tensor_col_offset(tCtAcc1_base)
        tCtAcc2_base = cute.make_tensor(acc2_ptr, tCtAcc_fake.layout)

        epi_tidx = tidx

        # Acc1: TMEM->reg copy 파티셔닝
        tiled_copy_t2r, tTR_tAcc1_base, tTR_rAcc1 = _epilog_tmem_copy_and_partition(
            epi_tidx,
            tCtAcc1_base,
            tCgC,
            epi_tile,
            use_2cta_instrs,
            cta_tile_shape_mnk,
            C_LAYOUT,
            C_DTYPE,
            ACC_DTYPE,
        )

        # Acc2: 같은 copy op를 reuse해서 partition만 새로 만듦
        thr_copy_t2r = tiled_copy_t2r.get_slice(epi_tidx)
        tAcc2_epi = cute.flat_divide(
            tCtAcc2_base[((None, None), 0, 0, None)],
            epi_tile,
        )
        tTR_tAcc2_base = thr_copy_t2r.partition_S(tAcc2_epi)
        tTR_rAcc2 = cute.make_rmem_tensor(tTR_rAcc1.shape, ACC_DTYPE)

        # output register
        tTR_rC = cute.make_rmem_tensor(tTR_rAcc1.shape, C_DTYPE)

        # reg -> shared store
        tiled_copy_r2s, tRS_rC, tRS_sC = _epilog_smem_copy_and_partition(
            tiled_copy_t2r,
            tTR_rC,
            epi_tidx,
            sC,
            C_LAYOUT,
            C_DTYPE,
            ACC_DTYPE,
        )

        # shared -> global TMA store partition
        tma_atom_c, bSG_sC, bSG_gC_partitioned = _epilog_gmem_copy_and_partition(
            tma_atom_c,
            tCgC,
            epi_tile,
            sC,
        )

        # scheduler
        tile_sched = utils.StaticPersistentTileScheduler.create(
            tile_sched_params, cute.arch.block_idx(), cute.arch.grid_dim()
        )
        work_tile = tile_sched.initial_work_tile_info()

        # ACC pipeline consumer state (MMA 완료를 기다림)
        acc_consumer_state = pipeline.make_pipeline_state(
            pipeline.PipelineUserType.Consumer, num_acc_stage
        )

        # TMA store pipeline (SMEM -> GMEM)
        c_producer_group = pipeline.CooperativeGroup(
            pipeline.Agent.Thread, 32 * len(EPILOG_WARPS)
        )
        c_pipeline = pipeline.PipelineTmaStore.create(
            num_stages=num_c_stage,
            producer_group=c_producer_group,
        )

        while work_tile.is_valid_tile:
            cur_tile_coord = work_tile.tile_idx
            mma_tile_coord_mnl = (
                cur_tile_coord[0] // cute.size(tiled_mma.thr_id.shape),
                cur_tile_coord[1],
                cur_tile_coord[2],
            )

            # 현재 타일의 gC 파티션
            bSG_gC = bSG_gC_partitioned[(None, None, None, *mma_tile_coord_mnl)]

            # 현재 acc stage의 acc1/acc2 파티션
            tTR_tAcc1 = tTR_tAcc1_base[
                (None, None, None, None, None, acc_consumer_state.index)
            ]
            tTR_tAcc2 = tTR_tAcc2_base[
                (None, None, None, None, None, acc_consumer_state.index)
            ]

            # MMA 완료 대기
            acc_pipeline.consumer_wait(acc_consumer_state)

            # group_modes로 rank 정리
            tTR_tAcc1 = cute.group_modes(tTR_tAcc1, 3, cute.rank(tTR_tAcc1))
            tTR_tAcc2 = cute.group_modes(tTR_tAcc2, 3, cute.rank(tTR_tAcc2))
            bSG_gC = cute.group_modes(bSG_gC, 1, cute.rank(bSG_gC))

            # epilogue subtile 루프 (sC의 stage 버퍼를 회전)
            subtile_cnt = cute.size(tTR_tAcc1.shape, mode=[3])
            num_prev_subtiles = tile_sched.num_tiles_executed * subtile_cnt

            for subtile_idx in cutlass.range(subtile_cnt):
                tTR_tAcc1_mn = tTR_tAcc1[(None, None, None, subtile_idx)]
                tTR_tAcc2_mn = tTR_tAcc2[(None, None, None, subtile_idx)]

                # TMEM -> reg
                cute.copy(tiled_copy_t2r, tTR_tAcc1_mn, tTR_rAcc1)
                cute.copy(tiled_copy_t2r, tTR_tAcc2_mn, tTR_rAcc2)

                # vector load
                acc_vec1 = tiled_copy_r2s.retile(tTR_rAcc1).load()
                acc_vec2 = tiled_copy_r2s.retile(tTR_rAcc2).load()

                # fused: silu(acc1) * acc2
                acc_vec1 = epilogue_op(acc_vec1)
                acc_vec = acc_vec1 * acc_vec2

                # convert to fp16 and store into reg buffer
                tRS_rC.store(acc_vec.to(C_DTYPE))

                # stage buffer index
                c_buffer = (num_prev_subtiles + subtile_idx) % num_c_stage

                # reg -> shared
                cute.copy(
                    tiled_copy_r2s,
                    tRS_rC,
                    tRS_sC[(None, None, None, c_buffer)],
                )

                # shared fence + epilogue warp sync
                cute.arch.fence_proxy(
                    cute.arch.ProxyKind.async_shared,
                    space=cute.arch.SharedSpace.shared_cta,
                )
                EPILOG_SYNC_BARRIER.arrive_and_wait()

                # warp0가 TMA store 실행
                if warp_idx == EPILOG_WARPS[0]:
                    cute.copy(
                        tma_atom_c,
                        bSG_sC[(None, c_buffer)],
                        bSG_gC[(None, subtile_idx)],
                    )
                    c_pipeline.producer_commit()
                    c_pipeline.producer_acquire()

                EPILOG_SYNC_BARRIER.arrive_and_wait()

            # consumer release (한 thread만 수행)
            with cute.arch.elect_one():
                acc_pipeline.consumer_release(acc_consumer_state)
            acc_consumer_state.advance()

            tile_sched.advance_to_next_work()
            work_tile = tile_sched.get_current_work()

        # TMEM free
        tmem.relinquish_alloc_permit()
        EPILOG_SYNC_BARRIER.arrive_and_wait()
        tmem.free(acc_tmem_ptr)
        c_pipeline.producer_tail()


# =============================================================================
# 4) Host-side JIT wrapper (@cute.jit)
# =============================================================================
@cute.jit
def dual_gemm_silu_launch(
    a_ptr: cute.Pointer,
    b1_ptr: cute.Pointer,
    b2_ptr: cute.Pointer,
    sfa_ptr: cute.Pointer,
    sfb1_ptr: cute.Pointer,
    sfb2_ptr: cute.Pointer,
    c_ptr: cute.Pointer,
    problem_mnkl: Tuple[int, int, int, int],
    max_active_clusters: cutlass.Constexpr = MAX_ACTIVE_CLUSTERS,
    epilogue_op: cutlass.Constexpr = lambda x: x
    * (1.0 / (1.0 + cute.math.exp(-x, fastmath=True))),
):
    """
    host-side 준비 단계:
      1) global tensor 생성 (A/B/C)
      2) scale factor tensor 생성 (SFA/SFB)
      3) tiled_mma 생성 (blockscaled)
      4) shared layouts + TMA atoms 생성
      5) persistent scheduler grid 계산
      6) device kernel launch
    """
    m, n, k, l = problem_mnkl

    # -------------------------------------------------------------------------
    # Global tensor views
    # - A: (m, k, l) with K-major stride
    # - B: (n, k, l) with K-major stride
    # - C: (m, n, l) row-major
    # -------------------------------------------------------------------------
    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_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)),
        ),
    )
    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_tensor = cute.make_tensor(
        c_ptr,
        cute.make_layout((cute.assume(m, 32), n, l), stride=(n, 1, m * n)),
    )

    # -------------------------------------------------------------------------
    # Scale factor tensors
    # - 입력으로 주어지는 permuted scale factor 텐서는
    #   blockscaled_utils.tile_atom_to_shape_SF()로 만든 layout과 shape가 일치해야 함.
    # -------------------------------------------------------------------------
    sfa_layout = blockscaled_utils.tile_atom_to_shape_SF(a_tensor.shape, SF_VEC_SIZE)
    sfa_tensor = cute.make_tensor(sfa_ptr, sfa_layout)

    sfb_layout = blockscaled_utils.tile_atom_to_shape_SF(b_tensor1.shape, SF_VEC_SIZE)
    sfb_tensor1 = cute.make_tensor(sfb1_ptr, sfb_layout)
    sfb_tensor2 = cute.make_tensor(sfb2_ptr, sfb_layout)

    # -------------------------------------------------------------------------
    # tiled_mma 생성 (blockscaled trivial MMA)
    # -------------------------------------------------------------------------
    tiled_mma = sm100_utils.make_blockscaled_trivial_tiled_mma(
        AB_DTYPE,
        A_MAJOR_MODE,
        B_MAJOR_MODE,
        SF_DTYPE,
        SF_VEC_SIZE,
        CTA_GROUP,
        MMA_INST_SHAPE_MN,
    )

    # SFB는 N이 128 배수가 아닐 수 있어 padding용 MMA를 하나 더 만든다.
    mma_inst_shape_mn_sfb = (
        MMA_INST_SHAPE_MN[0] // (2 if USE_2CTA_INSTRS else 1),
        cute.round_up(MMA_INST_SHAPE_MN[1], 128),
    )
    tiled_mma_sfb = sm100_utils.make_blockscaled_trivial_tiled_mma(
        AB_DTYPE,
        A_MAJOR_MODE,
        B_MAJOR_MODE,
        SF_DTYPE,
        SF_VEC_SIZE,
        tcgen05.CtaGroup.ONE,
        mma_inst_shape_mn_sfb,
    )

    # UMMA instruction K shape (보통 64)
    mma_inst_shape_k = cute.size(tiled_mma.shape_mnk, mode=[2])

    # 최종 CTA tile shape
    mma_tiler = (
        MMA_INST_SHAPE_MN[0],
        MMA_INST_SHAPE_MN[1],
        mma_inst_shape_k * MMA_INST_TILE_K,   # 보통 64*4=256
    )
    mma_tiler_sfb = (
        mma_inst_shape_mn_sfb[0],
        mma_inst_shape_mn_sfb[1],
        mma_inst_shape_k * MMA_INST_TILE_K,
    )

    # CTA가 실제로 담당하는 M은 2CTA 여부에 따라 나뉠 수 있음
    cta_tile_shape_mnk = (
        mma_tiler[0] // cute.size(tiled_mma.thr_id.shape),
        mma_tiler[1],
        mma_tiler[2],
    )

    # -------------------------------------------------------------------------
    # Cluster layout (multicast/persistent scheduler용)
    # -------------------------------------------------------------------------
    cluster_layout_vmnk = cute.tiled_divide(
        cute.make_layout((*CLUSTER_SHAPE_MN, 1)),
        (tiled_mma.thr_id.shape,),
    )
    cluster_layout_sfb_vmnk = cute.tiled_divide(
        cute.make_layout((*CLUSTER_SHAPE_MN, 1)),
        (tiled_mma_sfb.thr_id.shape,),
    )

    # -------------------------------------------------------------------------
    # Epilogue tile shape (C layout/타일 크기에 맞게 계산)
    # -------------------------------------------------------------------------
    epi_tile = sm100_utils.compute_epilogue_tile_shape(
        cta_tile_shape_mnk,
        USE_2CTA_INSTRS,
        C_LAYOUT,
        C_DTYPE,
    )

    # -------------------------------------------------------------------------
    # Stage counts (shared capacity 기반)
    # -------------------------------------------------------------------------
    num_acc_stage, num_ab_stage, num_c_stage = _compute_stages(
        tiled_mma,
        mma_tiler,
        AB_DTYPE,
        AB_DTYPE,
        epi_tile,
        C_DTYPE,
        C_LAYOUT,
        SF_DTYPE,
        SF_VEC_SIZE,
        SMEM_CAPACITY_BYTES,
        occupancy=1,
    )

    # -------------------------------------------------------------------------
    # Shared memory staged layouts 생성
    # -------------------------------------------------------------------------
    a_smem_layout_staged = sm100_utils.make_smem_layout_a(
        tiled_mma, mma_tiler, AB_DTYPE, num_ab_stage
    )
    b_smem_layout_staged = sm100_utils.make_smem_layout_b(
        tiled_mma, mma_tiler, AB_DTYPE, num_ab_stage
    )
    sfa_smem_layout_staged = blockscaled_utils.make_smem_layout_sfa(
        tiled_mma, mma_tiler, SF_VEC_SIZE, num_ab_stage
    )
    sfb_smem_layout_staged = blockscaled_utils.make_smem_layout_sfb(
        tiled_mma, mma_tiler, SF_VEC_SIZE, num_ab_stage
    )
    c_smem_layout_staged = sm100_utils.make_smem_layout_epi(
        C_DTYPE, C_LAYOUT, epi_tile, num_c_stage
    )

    atom_thr_size = cute.size(tiled_mma.thr_id.shape)

    # -------------------------------------------------------------------------
    # ✅ 2CTA-safe minimal patch:
    # TMA descriptor 생성 시 "v-dim(2CTA)"를 포함하지 않도록 v=1로 고정.
    # (rank-4: v, cluster_m, cluster_n, k)
    # -------------------------------------------------------------------------
    cluster_shape_vmnk_for_tma = cute.make_layout((1, *CLUSTER_SHAPE_MN, 1)).shape

    # -------------------------------------------------------------------------
    # TMA atoms 생성 (A/B/SF)
    # -------------------------------------------------------------------------
    a_op = sm100_utils.cluster_shape_to_tma_atom_A(CLUSTER_SHAPE_MN, tiled_mma.thr_id)
    a_smem_layout = cute.slice_(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,
        mma_tiler,
        tiled_mma,
        cluster_shape_vmnk_for_tma,
    )

    b_op = sm100_utils.cluster_shape_to_tma_atom_B(CLUSTER_SHAPE_MN, tiled_mma.thr_id)
    b_smem_layout = cute.slice_(b_smem_layout_staged, (None, None, None, 0))
    tma_atom_b1, tma_tensor_b1 = cute.nvgpu.make_tiled_tma_atom_B(
        b_op,
        b_tensor1,
        b_smem_layout,
        mma_tiler,
        tiled_mma,
        cluster_shape_vmnk_for_tma,
    )
    tma_atom_b2, tma_tensor_b2 = cute.nvgpu.make_tiled_tma_atom_B(
        b_op,
        b_tensor2,
        b_smem_layout,
        mma_tiler,
        tiled_mma,
        cluster_shape_vmnk_for_tma,
    )

    sfa_op = sm100_utils.cluster_shape_to_tma_atom_A(CLUSTER_SHAPE_MN, tiled_mma.thr_id)
    sfa_smem_layout = cute.slice_(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,
        mma_tiler,
        tiled_mma,
        cluster_shape_vmnk_for_tma,
        internal_type=cutlass.Int16,
    )

    sfb_op = sm100_utils.cluster_shape_to_tma_atom_SFB(
        CLUSTER_SHAPE_MN, tiled_mma.thr_id
    )
    sfb_smem_layout = cute.slice_(sfb_smem_layout_staged, (None, None, None, 0))
    tma_atom_sfb1, tma_tensor_sfb1 = cute.nvgpu.make_tiled_tma_atom_B(
        sfb_op,
        sfb_tensor1,
        sfb_smem_layout,
        mma_tiler_sfb,
        tiled_mma_sfb,
        cluster_shape_vmnk_for_tma,
        internal_type=cutlass.Int16,
    )
    tma_atom_sfb2, tma_tensor_sfb2 = cute.nvgpu.make_tiled_tma_atom_B(
        sfb_op,
        sfb_tensor2,
        sfb_smem_layout,
        mma_tiler_sfb,
        tiled_mma_sfb,
        cluster_shape_vmnk_for_tma,
        internal_type=cutlass.Int16,
    )

    # -------------------------------------------------------------------------
    # AB pipeline의 tx_count용: "한 stage에서 로드하는 총 바이트"
    # (TMA 완료를 tx_count 바이트 기준으로 감지)
    # -------------------------------------------------------------------------
    a_copy_size = cute.size_in_bytes(AB_DTYPE, a_smem_layout)
    b_copy_size = cute.size_in_bytes(AB_DTYPE, b_smem_layout)
    sfa_copy_size = cute.size_in_bytes(SF_DTYPE, sfa_smem_layout)
    sfb_copy_size = cute.size_in_bytes(SF_DTYPE, sfb_smem_layout)

    num_tma_load_bytes = (
        a_copy_size + 2 * b_copy_size + sfa_copy_size + 2 * sfb_copy_size
    ) * atom_thr_size

    # -------------------------------------------------------------------------
    # TMA store atom for C (shared -> global)
    # -------------------------------------------------------------------------
    epi_smem_layout = cute.slice_(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,
        epi_tile,
    )

    # -------------------------------------------------------------------------
    # Persistent scheduler grid 계산
    # -------------------------------------------------------------------------
    tile_sched_params, grid = _compute_grid(
        c_tensor,
        cta_tile_shape_mnk,
        CLUSTER_SHAPE_MN,
        max_active_clusters,
    )

    # -------------------------------------------------------------------------
    # device kernel launch
    # -------------------------------------------------------------------------
    dual_gemm_silu_kernel(
        tiled_mma,
        tiled_mma_sfb,
        tma_atom_a,
        tma_tensor_a,
        tma_atom_b1,
        tma_tensor_b1,
        tma_atom_b2,
        tma_tensor_b2,
        tma_atom_sfa,
        tma_tensor_sfa,
        tma_atom_sfb1,
        tma_tensor_sfb1,
        tma_atom_sfb2,
        tma_tensor_sfb2,
        tma_atom_c,
        tma_tensor_c,
        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,
        tile_sched_params,
        mma_tiler,
        mma_tiler_sfb,
        cta_tile_shape_mnk,
        num_acc_stage,
        num_ab_stage,
        num_c_stage,
        num_tma_load_bytes,
        epilogue_op,
    ).launch(
        grid=grid,
        block=[THREADS_PER_CTA, 1, 1],
        cluster=(*CLUSTER_SHAPE_MN, 1),
        min_blocks_per_mp=1,
    )
    return


# =============================================================================
# 5) Compile cache + custom_kernel (평가 프레임워크 entry)
# =============================================================================
_compiled: Dict[Tuple[Tuple[int, int, int, int], int], object] = {}


def _choose_max_active_clusters(problem_size: Tuple[int, int, int, int]) -> int:
    m, n, k, l = problem_size
    # CTA tile sizes are fixed by MMA_INST_SHAPE_MN.
    tile_m, tile_n = MMA_INST_SHAPE_MN

    grid_m = m // tile_m
    grid_n = n // tile_n
    total_clusters = (grid_m * grid_n) // (CLUSTER_SHAPE_MN[0] * CLUSTER_SHAPE_MN[1])

    # Only force persistence when there's at least 2 waves worth of clusters.
    if total_clusters > _PERSISTENT_ACTIVE_CLUSTERS_CAP:
        return _PERSISTENT_ACTIVE_CLUSTERS_CAP
    return total_clusters


def compile_kernel(problem_size: Tuple[int, int, int, int]):
    """
    문제 크기별로 JIT compile 결과를 캐싱.
    (m,n,k,l에 따라 specialization)
    """
    ps = tuple(problem_size)
    max_active_clusters = _choose_max_active_clusters(ps)
    key = (ps, max_active_clusters)
    if key in _compiled:
        return _compiled[key]

    # compile용 dummy pointers
    a_ptr = make_ptr(AB_DTYPE, 0, cute.AddressSpace.gmem, assumed_align=16)
    b1_ptr = make_ptr(AB_DTYPE, 0, cute.AddressSpace.gmem, assumed_align=16)
    b2_ptr = make_ptr(AB_DTYPE, 0, cute.AddressSpace.gmem, assumed_align=16)
    sfa_ptr = make_ptr(SF_DTYPE, 0, cute.AddressSpace.gmem, assumed_align=32)
    sfb1_ptr = make_ptr(SF_DTYPE, 0, cute.AddressSpace.gmem, assumed_align=32)
    sfb2_ptr = make_ptr(SF_DTYPE, 0, cute.AddressSpace.gmem, assumed_align=32)
    c_ptr = make_ptr(C_DTYPE, 0, cute.AddressSpace.gmem, assumed_align=16)

    compiled = cute.compile(
        dual_gemm_silu_launch,
        a_ptr,
        b1_ptr,
        b2_ptr,
        sfa_ptr,
        sfb1_ptr,
        sfb2_ptr,
        c_ptr,
        ps,                  # runtime tuple
        max_active_clusters, # constexpr
        options=f"--opt-level {COMPILE_OPT_LEVEL}",
    )

    _compiled[key] = compiled
    return compiled


def custom_kernel(data: input_t) -> output_t:
    """
    평가 프레임워크가 호출하는 entrypoint.

    data 구성 (문제 템플릿):
      (a, b1, b2, sfa_cpu, sfb1_cpu, sfb2_cpu, sfa_perm, sfb1_perm, sfb2_perm, c)

    - a/b는 torch e2m1_x2 packing을 쓰므로
      a.shape = (m, k_packed, l) 이고 실제 logical K = 2 * k_packed
    - scale factor는 permuted 텐서(sfa_perm/sfb*_perm)를 그대로 사용
    """
    a, b1, b2, sfa_cpu_unused, sfb1_cpu_unused, sfb2_cpu_unused, sfa_perm, sfb1_perm, sfb2_perm, c = data

    m, k_packed, l = a.shape
    n, _, _ = b1.shape

    # Torch는 FP4 두 값을 1바이트에 pack하므로 logical K가 2배
    k = k_packed * 2

    problem_size = (m, n, k, l)
    compiled = compile_kernel(problem_size)

    # 실제 torch 텐서 포인터를 CuTe pointer로 래핑
    a_ptr = make_ptr(AB_DTYPE, a.data_ptr(), cute.AddressSpace.gmem, assumed_align=16)
    b1_ptr = make_ptr(AB_DTYPE, b1.data_ptr(), cute.AddressSpace.gmem, assumed_align=16)
    b2_ptr = make_ptr(AB_DTYPE, b2.data_ptr(), cute.AddressSpace.gmem, assumed_align=16)

    sfa_ptr = make_ptr(SF_DTYPE, sfa_perm.data_ptr(), cute.AddressSpace.gmem, assumed_align=32)
    sfb1_ptr = make_ptr(SF_DTYPE, sfb1_perm.data_ptr(), cute.AddressSpace.gmem, assumed_align=32)
    sfb2_ptr = make_ptr(SF_DTYPE, sfb2_perm.data_ptr(), cute.AddressSpace.gmem, assumed_align=32)

    c_ptr = make_ptr(C_DTYPE, c.data_ptr(), cute.AddressSpace.gmem, assumed_align=16)

    # 실행
    compiled(a_ptr, b1_ptr, b2_ptr, sfa_ptr, sfb1_ptr, sfb2_ptr, c_ptr, problem_size)
    return c
scrolls · 1704 lines total

Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0

Best evidence level for this revision: reported

JSON