submission 383336
currybab · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 1769 lines, June 9 Researcher Reciprocity License v1.0.
submission_4.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-modal-nvfp4-dual-gemm-383336?include=source"interfacepython
Compatibility
measured onNVIDIA B200
declared hardwareNVIDIA B200
architecturessm_100
dtypesfp8_e4m3, nvfp4
Benchmark evidence
1 measurement across 1 GPU, fastest first.
Operation / workload
Hardware
Latency
Rank
Observed
Reported · How evidence levels are derived →
Source and license
sourceavailable
revision digestsha256:19adb5b055ccca22c3bd1d46a1436e4432a7a1b4824646e90245887d16f04bbf
license declaredunknown
license concludedunknown
authorscurrybab
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
AB_DTYPE = cutlass.Float4E2M1FN # nvfp4fused-epilogue
Epilogue에서 accumulator(TMEM)를 register로 load하기 위한 copy/partition 생성.mbarrier
CTA_SYNC_BARRIER = pipeline.NamedBarrier(barrier_id=1, num_threads=THREADS_PER_CTA)persistent-kernel
_PERSISTENT_ACTIVE_CLUSTERS_CAP = 32shared-memory
SMEM_CAPACITY_BYTES = utils.get_smem_capacity_in_bytes("sm_100")tcgen05
A_MAJOR_MODE = tcgen05.OperandMajorMode.Ktile-k = 4
MMA_INST_TILE_K = 4warp-specialization
ab_pipeline_producer_group = pipeline.CooperativeGroup(pipeline.Agent.Thread)Kernel source
submission_4.py1769 lines
from typing import Dict, Tuple, Union
import cuda.bindings.driver as cuda
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
# =============================================================================
# 개요 (NO-PERSISTENT VARIANT)
# =============================================================================
# 이 커널은 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를 사용합니다.
# 이 파일(submission_4_no_persistent.py)은 A/B 실험을 위해 scheduler를 "비-persistent처럼"
# 동작하도록 강제합니다. (max_active_clusters = total_clusters, num_acc_stage=1)
# 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로 공유.
# - (4,1): M방향으로 4 CTA가 묶여서 B/SFB를 multicast로 공유.
#
# NOTE: cluster shape는 TMA atom 생성/cluster launch에 compile-time 상수로 필요하다.
# 일부 환경(Python/CUTLASS DSL 버전)에서 constexpr tuple 인자가 불안정할 수 있어,
# dual_gemm_silu_launch는 전역 CLUSTER_SHAPE_MN을 참조하도록 유지하고,
# compile_kernel()에서 compilation 직전에 값을 바꿔서 shape-specialize 한다.
DEFAULT_CLUSTER_SHAPE_MN = (1, 4)
CLUSTER_SHAPE_MN = DEFAULT_CLUSTER_SHAPE_MN
# 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,
force_num_acc_stage: int,
) -> Tuple[int, int, int]:
"""
shared memory capacity에서 pipeline stage 수(AB stage, C stage) 등을 계산.
dual-GEMM이라 B/SFB가 2번 로드되는 것까지 고려해서 stage 수를 산정.
"""
# ACC stage:
# - num_acc_stage=2: overlap epilogue of tile i with mainloop of tile i+1 (only matters if persistent is active)
# - num_acc_stage=1: less TMEM/pipeline overhead for "single-tile-per-cluster" cases
num_acc_stage = int(force_num_acc_stage)
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 필요
cluster_size_mn = num_mcast_ctas_a * num_mcast_ctas_b
if cluster_size_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 cluster_size_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,
num_acc_stage_override: cutlass.Constexpr = 2,
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 (BlockScaled SF layouts assume inst_k_tiles=4)
)
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,
force_num_acc_stage=num_acc_stage_override,
)
# -------------------------------------------------------------------------
# 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], Tuple[int, int], int, int], object
] = {}
def _ceil_div(a: int, b: int) -> int:
return (a + b - 1) // b
def _count_total_clusters(
problem_size: Tuple[int, int, int, int], cluster_shape_mn: Tuple[int, int]
) -> int:
m, n, _k, l = problem_size
tile_m, tile_n = MMA_INST_SHAPE_MN
cluster_m, cluster_n = cluster_shape_mn
grid_m = m // tile_m
grid_n = n // tile_n
# Cluster mode requires full clusters. Model the scheduler as rounding up to full clusters.
clusters_m = _ceil_div(grid_m, cluster_m)
clusters_n = _ceil_div(grid_n, cluster_n)
return clusters_m * clusters_n * l
def _device_wave_cluster_cap(cluster_shape_mn: Tuple[int, int]) -> int:
"""
Approximate "one-wave" active clusters for the current device.
With occupancy=1, each CTA typically occupies one SM; a (cluster_m, cluster_n) cluster
uses cluster_m*cluster_n SMs concurrently.
"""
try:
dev = torch.cuda.current_device()
sm_count = torch.cuda.get_device_properties(dev).multi_processor_count
cluster_size = cluster_shape_mn[0] * cluster_shape_mn[1]
return max(1, sm_count // cluster_size)
except Exception:
return _PERSISTENT_ACTIVE_CLUSTERS_CAP
def _choose_max_active_clusters(
problem_size: Tuple[int, int, int, int], cluster_shape_mn: Tuple[int, int]
) -> int:
# Force "no persistent": launch enough active clusters to cover all work.
return _count_total_clusters(problem_size, cluster_shape_mn)
def _choose_num_acc_stage(
problem_size: Tuple[int, int, int, int],
cluster_shape_mn: Tuple[int, int],
max_active_clusters: int,
) -> int:
# No persistent -> no need for double-buffered accumulators.
return 1
def _choose_cluster_shape(problem_size: Tuple[int, int, int, int]) -> Tuple[int, int]:
# Baseline: A/SFA multicast across N (cluster_n=4).
# Other shapes (e.g., (4,1)) did not show wins for the current test set.
return DEFAULT_CLUSTER_SHAPE_MN
def compile_kernel(problem_size: Tuple[int, int, int, int]):
"""
문제 크기별로 JIT compile 결과를 캐싱.
(m,n,k,l에 따라 specialization)
"""
ps = tuple(problem_size)
cluster_shape_mn = _choose_cluster_shape(ps)
max_active_clusters = _choose_max_active_clusters(ps, cluster_shape_mn)
num_acc_stage_override = _choose_num_acc_stage(
ps, cluster_shape_mn, max_active_clusters
)
key = (ps, cluster_shape_mn, max_active_clusters, num_acc_stage_override)
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)
global CLUSTER_SHAPE_MN
prev_cluster_shape = CLUSTER_SHAPE_MN
CLUSTER_SHAPE_MN = cluster_shape_mn
try:
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
num_acc_stage_override, # constexpr
options=f"--opt-level {COMPILE_OPT_LEVEL}",
)
finally:
CLUSTER_SHAPE_MN = prev_cluster_shape
_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 · 1769 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 283829.
- from __future__ import annotations-from typing import Dict, Tuple, Union+ import cuda.bindings.driver as cuda+import torchimport cutlass⋯ 8 unchanged linesfrom task import input_t, output_t# =============================================================================- # 개요+ # 개요 (NO-PERSISTENT VARIANT)# =============================================================================# 이 커널은 Blackwell(SM100, B200)에서 "block-scaled NVFP4 GEMM"을 2번 수행하고,# epilogue에서 아래 연산을 fused로 수행합니다.⋯ 5 unchanged lines# - Accumulator : Float32 (TMEM)# - Output C : Float16#- # 또한 persistent tile scheduler를 사용하여, 하나의 CTA가 여러 타일을 순차적으로 처리합니다.+ # 또한 persistent tile scheduler를 사용합니다.+ # 이 파일(submission_4_no_persistent.py)은 A/B 실험을 위해 scheduler를 "비-persistent처럼"+ # 동작하도록 강제합니다. (max_active_clusters = total_clusters, num_acc_stage=1)# Warp specialization:# - warp 5: TMA load 전담 (global -> shared)# - warp 4: MMA 전담 (shared -> TMEM(SF), UMMA 2x GEMM)⋯ 28 unchanged linesMMA_INST_TILE_K = 4# 클러스터(CTA 그룹) 형태: (cluster_m, cluster_n)- # (1,4)이면 N방향으로 4 CTA가 묶여서 A/SFA를 multicast로 공유.- CLUSTER_SHAPE_MN = (1, 4)+ # - (1,4): N방향으로 4 CTA가 묶여서 A/SFA를 multicast로 공유.+ # - (4,1): M방향으로 4 CTA가 묶여서 B/SFB를 multicast로 공유.+ #+ # NOTE: cluster shape는 TMA atom 생성/cluster launch에 compile-time 상수로 필요하다.+ # 일부 환경(Python/CUTLASS DSL 버전)에서 constexpr tuple 인자가 불안정할 수 있어,+ # dual_gemm_silu_launch는 전역 CLUSTER_SHAPE_MN을 참조하도록 유지하고,+ # compile_kernel()에서 compilation 직전에 값을 바꿔서 shape-specialize 한다.+ DEFAULT_CLUSTER_SHAPE_MN = (1, 4)+ CLUSTER_SHAPE_MN = DEFAULT_CLUSTER_SHAPE_MN# persistent scheduler에서 활성 클러스터 상한# NOTE: If this is >= total clusters, the kernel degenerates to non-persistent (1 tile per cluster),⋯ 50 unchanged linessf_vec_size: int,smem_capacity: int,occupancy: int,+ force_num_acc_stage: 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+ # ACC stage:+ # - num_acc_stage=2: overlap epilogue of tile i with mainloop of tile i+1 (only matters if persistent is active)+ # - num_acc_stage=1: less TMEM/pipeline overhead for "single-tile-per-cluster" cases+ num_acc_stage = int(force_num_acc_stage)num_c_stage = 2# 1-stage 기준으로 각 buffer가 얼마나 shared memory를 먹는지 계산⋯ 440 unchanged lines)# cluster arrive: 여러 CTA가 협업하는 경우 relaxed arrive 필요- if cute.size(CLUSTER_SHAPE_MN) > 1:+ cluster_size_mn = num_mcast_ctas_a * num_mcast_ctas_b+ if cluster_size_mn > 1:cute.arch.cluster_arrive_relaxed()# -------------------------------------------------------------------------⋯ 147 unchanged lines# -------------------------------------------------------------------------# cluster sync: shared buffer/pipe 초기화 이후 동기화# -------------------------------------------------------------------------- if cute.size(CLUSTER_SHAPE_MN) > 1:+ if cluster_size_mn > 1:cute.arch.cluster_wait()else:CTA_SYNC_BARRIER.arrive_and_wait()⋯ 548 unchanged linesc_ptr: cute.Pointer,problem_mnkl: Tuple[int, int, int, int],max_active_clusters: cutlass.Constexpr = MAX_ACTIVE_CLUSTERS,+ num_acc_stage_override: cutlass.Constexpr = 2,epilogue_op: cutlass.Constexpr = lambda x: x* (1.0 / (1.0 + cute.math.exp(-x, fastmath=True))),):⋯ 87 unchanged linesmma_tiler = (MMA_INST_SHAPE_MN[0],MMA_INST_SHAPE_MN[1],- mma_inst_shape_k * MMA_INST_TILE_K, # 보통 64*4=256+ mma_inst_shape_k * MMA_INST_TILE_K, # 64*4=256 (BlockScaled SF layouts assume inst_k_tiles=4))mma_tiler_sfb = (mma_inst_shape_mn_sfb[0],⋯ 45 unchanged linesSF_VEC_SIZE,SMEM_CAPACITY_BYTES,occupancy=1,+ force_num_acc_stage=num_acc_stage_override,)# -------------------------------------------------------------------------⋯ 69 unchanged linesinternal_type=cutlass.Int16,)- sfb_op = sm100_utils.cluster_shape_to_tma_atom_SFB(- CLUSTER_SHAPE_MN, tiled_mma.thr_id- )+ 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,⋯ 42 unchanged lines# Persistent scheduler grid 계산# -------------------------------------------------------------------------tile_sched_params, grid = _compute_grid(- c_tensor,- cta_tile_shape_mnk,- CLUSTER_SHAPE_MN,- max_active_clusters,+ c_tensor, cta_tile_shape_mnk, CLUSTER_SHAPE_MN, max_active_clusters)# -------------------------------------------------------------------------⋯ 45 unchanged lines# =============================================================================# 5) Compile cache + custom_kernel (평가 프레임워크 entry)# =============================================================================- _compiled: Dict[Tuple[Tuple[int, int, int, int], int], object] = {}+ _compiled: Dict[+ Tuple[Tuple[int, int, int, int], Tuple[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.+ def _ceil_div(a: int, b: int) -> int:+ return (a + b - 1) // b+++ def _count_total_clusters(+ problem_size: Tuple[int, int, int, int], cluster_shape_mn: Tuple[int, int]+ ) -> int:+ m, n, _k, l = problem_sizetile_m, tile_n = MMA_INST_SHAPE_MN+ cluster_m, cluster_n = cluster_shape_mngrid_m = m // tile_mgrid_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:+ # Cluster mode requires full clusters. Model the scheduler as rounding up to full clusters.+ clusters_m = _ceil_div(grid_m, cluster_m)+ clusters_n = _ceil_div(grid_n, cluster_n)+ return clusters_m * clusters_n * l+++ def _device_wave_cluster_cap(cluster_shape_mn: Tuple[int, int]) -> int:+ """+ Approximate "one-wave" active clusters for the current device.+ With occupancy=1, each CTA typically occupies one SM; a (cluster_m, cluster_n) cluster+ uses cluster_m*cluster_n SMs concurrently.+ """+ try:+ dev = torch.cuda.current_device()+ sm_count = torch.cuda.get_device_properties(dev).multi_processor_count+ cluster_size = cluster_shape_mn[0] * cluster_shape_mn[1]+ return max(1, sm_count // cluster_size)+ except Exception:return _PERSISTENT_ACTIVE_CLUSTERS_CAP- return total_clusters+ def _choose_max_active_clusters(+ problem_size: Tuple[int, int, int, int], cluster_shape_mn: Tuple[int, int]+ ) -> int:+ # Force "no persistent": launch enough active clusters to cover all work.+ return _count_total_clusters(problem_size, cluster_shape_mn)+++ def _choose_num_acc_stage(+ problem_size: Tuple[int, int, int, int],+ cluster_shape_mn: Tuple[int, int],+ max_active_clusters: int,+ ) -> int:+ # No persistent -> no need for double-buffered accumulators.+ return 1+++ def _choose_cluster_shape(problem_size: Tuple[int, int, int, int]) -> Tuple[int, int]:+ # Baseline: A/SFA multicast across N (cluster_n=4).+ # Other shapes (e.g., (4,1)) did not show wins for the current test set.+ return DEFAULT_CLUSTER_SHAPE_MN++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)+ cluster_shape_mn = _choose_cluster_shape(ps)+ max_active_clusters = _choose_max_active_clusters(ps, cluster_shape_mn)+ num_acc_stage_override = _choose_num_acc_stage(+ ps, cluster_shape_mn, max_active_clusters+ )+ key = (ps, cluster_shape_mn, max_active_clusters, num_acc_stage_override)if key in _compiled:return _compiled[key]⋯ 6 unchanged linessfb2_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}",- )+ global CLUSTER_SHAPE_MN+ prev_cluster_shape = CLUSTER_SHAPE_MN+ CLUSTER_SHAPE_MN = cluster_shape_mn+ try:+ 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+ num_acc_stage_override, # constexpr+ options=f"--opt-level {COMPILE_OPT_LEVEL}",+ )+ finally:+ CLUSTER_SHAPE_MN = prev_cluster_shape_compiled[key] = compiledreturn compiled
scrolls · 264 diff lines total
Best evidence level for this revision: reported
JSON