Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
233 changes: 113 additions & 120 deletions examples/blockscaled_gemm_sm100/gemm_mxfp8_blockscaled_1d1d.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,10 @@
# MXFP8 Block-Scaled GEMM on SM100
# Blockscale size: (M, N, K) = (1, 1, 128)
#
# Persistent 2-CTA variant uses PersistentTileScheduler: each warp role owns a
# scheduler instance and drives ``while sched.valid()``. ``sched.current_iter[0]``
# is the wave index for pipeline/double-buffering; ``sched.m_idx[0]`` /
# ``sched.n_idx[0]`` decode the tile (with ``bx = m_idx * cluster_size + cta_id``).

import argparse
import torch
Expand Down Expand Up @@ -347,16 +352,14 @@ def mxfp8_blockscaled_gemm_2cta_persistent(
C = T.empty((M, N), out_dtype)

sm_num = driver.get_num_sms()
num_clusters = sm_num // 2
m_blocks = T.ceildiv(M, block_M)
m_clusters = m_blocks // 2
n_blocks = T.ceildiv(N, block_N)
assert K % (2 * block_K) == 0 # for simplicity
waves = T.ceildiv(m_blocks * n_blocks, sm_num)
group_size = 16 # in cluster
cluster_size = 2
group_size = 16
assert n_blocks % (2 * group_size) == 0 # Please adjust group_size if not satisfied

with T.ClusterKernel(sm_num, threads=256, cluster_dims=2) as (block_id):
with T.ClusterKernel(sm_num, threads=256, cluster_dims=cluster_size) as (block_id):
cta_id = T.block_rank_in_cluster()
T.assume(cta_id < 2)

Expand Down Expand Up @@ -386,130 +389,120 @@ def mxfp8_blockscaled_gemm_2cta_persistent(
warp_idx = tx // 32

if warp_idx == 0:
for w in T.unroll(waves):
cluster_id = block_id // 2
tile_id = num_clusters * w + cluster_id
bx_cluster = (tile_id // group_size) % m_clusters
bx = bx_cluster * 2 + cta_id
by = (tile_id % group_size) + (tile_id // group_size) // m_clusters * group_size

if bx * block_M < M and by * block_N < N:
for k in T.serial(k_iters):
phase = w * k_iters + k
stage = phase % num_stages
parity = (phase // num_stages) & 1
T.mbarrier_wait_parity(consumed[stage], parity ^ 1)
sched = T.PersistentTileScheduler("sched_tma", m_blocks, n_blocks, swizzle_size=group_size, cluster_size=cluster_size)
sched.init(block_id // cluster_size)
while sched.valid():
bx = sched.m_idx[0] * cluster_size + cta_id
by = sched.n_idx[0]
Comment on lines +392 to +396

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

🩺 Stability & Availability | 🟠 Major

🧩 Analysis chain

🏁 Script executed:

#!/bin/bash
rg -nP -C3 'mxfp8_blockscaled_gemm_2cta_persistent|block_M\s*=' examples/blockscaled_gemm_sm100/

Repository: tile-ai/tilelang

Length of output: 7196


🏁 Script executed:

# Get the outline of the mxfp8_blockscaled_gemm_2cta_persistent function
ast-grep outline examples/blockscaled_gemm_sm100/gemm_mxfp8_blockscaled_1d1d.py --view expanded | grep -A 50 "mxfp8_blockscaled_gemm_2cta_persistent"

Repository: tile-ai/tilelang

Length of output: 528


🏁 Script executed:

# Read the function definition and the scheduler/bx calculation area
sed -n '324,410p' examples/blockscaled_gemm_sm100/gemm_mxfp8_blockscaled_1d1d.py

Repository: tile-ai/tilelang

Length of output: 3433


🏁 Script executed:

# Check the cluster_size and m_blocks calculation
sed -n '358,370p' examples/blockscaled_gemm_sm100/gemm_mxfp8_blockscaled_1d1d.py

Repository: tile-ai/tilelang

Length of output: 682


🏁 Script executed:

# Look for any M divisibility checks in the entry point function or callers
rg -n 'assert.*M.*%|M.*%.*block_M|assert.*m_blocks' examples/blockscaled_gemm_sm100/gemm_mxfp8_blockscaled_1d1d.py

Repository: tile-ai/tilelang

Length of output: 154


🏁 Script executed:

# Check line 501 (C store) and surrounding context for guards
sed -n '495,510p' examples/blockscaled_gemm_sm100/gemm_mxfp8_blockscaled_1d1d.py

Repository: tile-ai/tilelang

Length of output: 826


🏁 Script executed:

# Find the wrapper/entry point that calls mxfp8_blockscaled_gemm_2cta_persistent
sed -n '630,680p' examples/blockscaled_gemm_sm100/gemm_mxfp8_blockscaled_1d1d.py

Repository: tile-ai/tilelang

Length of output: 2559


🏁 Script executed:

# Check for M divisibility assertions in the main entry point
rg -B5 -A5 'def blockscaled_gemm\(' examples/blockscaled_gemm_sm100/gemm_mxfp8_blockscaled_1d1d.py | head -40

Repository: tile-ai/tilelang

Length of output: 154


🏁 Script executed:

# Look at the full blockscaled_gemm function to see if it enforces M divisibility
sed -n '595,675p' examples/blockscaled_gemm_sm100/gemm_mxfp8_blockscaled_1d1d.py

Repository: tile-ai/tilelang

Length of output: 3529


🏁 Script executed:

# Check grouped_gemm caller as well
sed -n '540,580p' examples/blockscaled_gemm_sm100/grouped_gemm_mxfp8_blockscaled_1d1d.py

Repository: tile-ai/tilelang

Length of output: 1104


🏁 Script executed:

# Search for any tests or examples that show how these functions are called
rg -n 'blockscaled_gemm\(' examples/blockscaled_gemm_sm100/ --type py -A3 | head -60

Repository: tile-ai/tilelang

Length of output: 1650


🏁 Script executed:

# Check if there are any guards around bx accesses (lines 400-415 for TMA load and 495-510 for store)
sed -n '398,415p' examples/blockscaled_gemm_sm100/gemm_mxfp8_blockscaled_1d1d.py

Repository: tile-ai/tilelang

Length of output: 1083


🏁 Script executed:

# Look for any validation in the kernel entry or wrapper functions
sed -n '185,200p' examples/blockscaled_gemm_sm100/gemm_mxfp8_blockscaled_1d1d.py

Repository: tile-ai/tilelang

Length of output: 565


🏁 Script executed:

# Check if there are any wrapper or high-level APIs that validate M
rg -n 'def.*blockscaled_gemm' examples/blockscaled_gemm_sm100/gemm_mxfp8_blockscaled_1d1d.py | head -10

Repository: tile-ai/tilelang

Length of output: 371


🏁 Script executed:

# Search for any comments or documentation about M constraints
rg -B2 'M.*divisible|M.*multiple|cluster_size.*block_M' examples/blockscaled_gemm_sm100/gemm_mxfp8_blockscaled_1d1d.py

Repository: tile-ai/tilelang

Length of output: 154


🏁 Script executed:

# Check if there are any conditional guards around the bx TMA access or store
sed -n '390,425p' examples/blockscaled_gemm_sm100/gemm_mxfp8_blockscaled_1d1d.py

Repository: tile-ai/tilelang

Length of output: 1866


🏁 Script executed:

# Look at the C store operation more carefully for any guards
sed -n '490,510p' examples/blockscaled_gemm_sm100/gemm_mxfp8_blockscaled_1d1d.py

Repository: tile-ai/tilelang

Length of output: 1033


🏁 Script executed:

# Check git history to see if bx < m_blocks was recently removed
git log -p --all -S "bx < m_blocks" -- examples/blockscaled_gemm_sm100/gemm_mxfp8_blockscaled_1d1d.py 2>/dev/null | head -80

Repository: tile-ai/tilelang

Length of output: 154


🏁 Script executed:

# Check all usages of bx in the file to see if there are any guards
rg -n 'bx' examples/blockscaled_gemm_sm100/gemm_mxfp8_blockscaled_1d1d.py | grep -v 'by\|swap'

Repository: tile-ai/tilelang

Length of output: 978


🏁 Script executed:

# Let's check what the clustered scheduler actually does - look for its implementation
rg -n 'class PersistentTileScheduler|def PersistentTileScheduler' --type py 2>/dev/null | head -5

Repository: tile-ai/tilelang

Length of output: 154


🏁 Script executed:

# Check if M divisibility is validated at the wrapper/API level
sed -n '18,75p' examples/blockscaled_gemm_sm100/gemm_mxfp8_blockscaled_1d1d.py

Repository: tile-ai/tilelang

Length of output: 2225


🏁 Script executed:

# Search for any M divisibility checks anywhere in the repo
rg -n 'M.*%.*block_M|assert.*M' examples/blockscaled_gemm_sm100/ --type py

Repository: tile-ai/tilelang

Length of output: 576


Out-of-bounds bx index when M is not divisible by cluster_size * block_M.

The calculation bx = sched.m_idx[0] * cluster_size + cta_id can exceed m_blocks when M is not a multiple of cluster_size * block_M (i.e., 256). For example, if M = 257, the scheduler creates m_blocks = 3, but the tail cluster allows bx = 3 when cta_id = 1, causing out-of-bounds reads in the TMA load (line 404) and stores in the C epilogue (line 501).

The test suite only exercises M = 8192, which is perfectly divisible by 256, masking this issue. Add an assertion that M % (cluster_size * block_M) == 0, or restore a bx < m_blocks guard around the tensor accesses.

🤖 Prompt for AI Agents
Verify each finding against current code. Fix only still-valid issues, skip the
rest with a brief reason, keep changes minimal, and validate.

In `@examples/blockscaled_gemm_sm100/gemm_mxfp8_blockscaled_1d1d.py` around lines
392 - 396, The calculation of bx using sched.m_idx[0] * cluster_size + cta_id
can exceed m_blocks when M is not divisible by cluster_size * block_M, causing
out-of-bounds accesses during TMA loads and C epilogue stores. Fix this by
adding an assertion at the point where cluster_size and block_M are defined to
ensure M % (cluster_size * block_M) == 0, or alternatively add a bounds check
guard condition (bx < m_blocks) around the tensor access operations that use bx,
such as the TMA load operation and the C epilogue store operations, to prevent
accessing invalid indices.


for k in T.serial(k_iters):
phase = sched.current_iter[0] * k_iters + k
stage = phase % num_stages
parity = (phase // num_stages) & 1
T.mbarrier_wait_parity(consumed[stage], parity ^ 1)
T.tma_copy(
A[bx * block_M : (bx + 1) * block_M, k * block_K : (k + 1) * block_K],
A_shared[stage, :, :],
barrier=loaded[stage],
)
if transpose_B:
T.tma_copy(
A[bx * block_M : (bx + 1) * block_M, k * block_K : (k + 1) * block_K],
A_shared[stage, :, :],
B[
by * block_N + cta_id * half_N : by * block_N + (cta_id + 1) * half_N,
k * block_K : (k + 1) * block_K,
],
B_shared[stage, :, :],
barrier=loaded[stage],
)
if transpose_B:
T.tma_copy(
B[
by * block_N + cta_id * half_N : by * block_N + (cta_id + 1) * half_N,
k * block_K : (k + 1) * block_K,
],
B_shared[stage, :, :],
barrier=loaded[stage],
)
else:
T.tma_copy(
B[
k * block_K : (k + 1) * block_K,
by * block_N + cta_id * half_N : by * block_N + (cta_id + 1) * half_N,
],
B_shared[stage, :, :],
barrier=loaded[stage],
)
if k % sf_load_period == 0:
sf_group_idx = k // sf_load_period
T.tma_copy(
SFA[sf_group_idx * M + bx * block_M : sf_group_idx * M + (bx + 1) * block_M],
SFA_shared[stage, :],
barrier=loaded[stage],
)
T.tma_copy(
SFB[sf_group_idx * N + by * block_N : sf_group_idx * N + (by + 1) * block_N],
SFB_shared[stage, :],
barrier=loaded[stage],
)
T.mbarrier_arrive(loaded[stage])

elif warp_idx == 1 and cta_id == 0:
for w in T.unroll(waves):
cluster_id = block_id // 2
tile_id = num_clusters * w + cluster_id
bx_cluster = (tile_id // group_size) % m_clusters
bx = bx_cluster * 2 + cta_id
by = (tile_id % group_size) + (tile_id // group_size) // m_clusters * group_size

if bx * block_M < M and by * block_N < N:
T.mbarrier_wait_parity(tmem_empty, (w & 1) ^ 1)
for k in T.serial(k_iters):
phase = w * k_iters + k
stage = phase % num_stages
parity = (phase // num_stages) & 1
T.mbarrier_wait_parity(with_sf_full[stage], parity)
if k % sf_load_period == 0:
T.tcgen05_cp_warpx4(SFA_shared[stage, :], SFA_tmem, use_2cta=True)
T.tcgen05_cp_warpx4(SFB_shared[stage, :], SFB_tmem, use_2cta=True)
T.tcgen05_gemm_blockscaled(
A_shared[stage, :, :],
else:
T.tma_copy(
B[
k * block_K : (k + 1) * block_K,
by * block_N + cta_id * half_N : by * block_N + (cta_id + 1) * half_N,
],
B_shared[stage, :, :],
C_tmem,
SFA_tmem,
SFB_tmem,
transpose_B=transpose_B,
mbar=consumed[stage],
clear_accum=k == 0,
k_start=k * block_K,
sf_a_granularity_k=sf_granularity_k,
sf_b_granularity_k=sf_granularity_k,
use_2cta=True,
barrier=loaded[stage],
)
T.tcgen05_mma_arrive(tmem_full, arrive_2cta=True)
if k % sf_load_period == 0:
sf_group_idx = k // sf_load_period
T.tma_copy(
SFA[sf_group_idx * M + bx * block_M : sf_group_idx * M + (bx + 1) * block_M],
SFA_shared[stage, :],
barrier=loaded[stage],
)
T.tma_copy(
SFB[sf_group_idx * N + by * block_N : sf_group_idx * N + (by + 1) * block_N],
SFB_shared[stage, :],
barrier=loaded[stage],
)
T.mbarrier_arrive(loaded[stage])
sched.next_tile()

elif warp_idx == 1 and cta_id == 0:
sched = T.PersistentTileScheduler("sched_mma", m_blocks, n_blocks, swizzle_size=group_size, cluster_size=cluster_size)
sched.init(block_id // cluster_size)
while sched.valid():
T.mbarrier_wait_parity(tmem_empty, (sched.current_iter[0] & 1) ^ 1)
for k in T.serial(k_iters):
phase = sched.current_iter[0] * k_iters + k
stage = phase % num_stages
parity = (phase // num_stages) & 1
T.mbarrier_wait_parity(with_sf_full[stage], parity)
if k % sf_load_period == 0:
T.tcgen05_cp_warpx4(SFA_shared[stage, :], SFA_tmem, use_2cta=True)
T.tcgen05_cp_warpx4(SFB_shared[stage, :], SFB_tmem, use_2cta=True)
T.tcgen05_gemm_blockscaled(
A_shared[stage, :, :],
B_shared[stage, :, :],
C_tmem,
SFA_tmem,
SFB_tmem,
transpose_B=transpose_B,
mbar=consumed[stage],
clear_accum=k == 0,
k_start=k * block_K,
sf_a_granularity_k=sf_granularity_k,
sf_b_granularity_k=sf_granularity_k,
use_2cta=True,
)
T.tcgen05_mma_arrive(tmem_full, arrive_2cta=True)
sched.next_tile()

elif warp_idx == 2:
for w in T.unroll(waves):
cluster_id = block_id // 2
tile_id = num_clusters * w + cluster_id
bx_cluster = (tile_id // group_size) % m_clusters
bx = bx_cluster * 2 + cta_id
by = (tile_id % group_size) + (tile_id // group_size) // m_clusters * group_size

if bx * block_M < M and by * block_N < N:
for k in T.serial(k_iters):
phase = w * k_iters + k
stage = phase % num_stages
parity = (phase // num_stages) & 1
T.mbarrier_wait_parity(loaded[stage], parity)
if k % sf_load_period == 0:
T.tcgen05_sf_warp_transpose(SFA_shared[stage, :])
T.tcgen05_sf_warp_transpose(SFB_shared[stage, :])
T.fence_proxy_async()
T.mbarrier_arrive(with_sf_full[stage], 0)
sched = T.PersistentTileScheduler("sched_sf", m_blocks, n_blocks, swizzle_size=group_size, cluster_size=cluster_size)
sched.init(block_id // cluster_size)
while sched.valid():
for k in T.serial(k_iters):
phase = sched.current_iter[0] * k_iters + k
stage = phase % num_stages
parity = (phase // num_stages) & 1
T.mbarrier_wait_parity(loaded[stage], parity)
if k % sf_load_period == 0:
T.tcgen05_sf_warp_transpose(SFA_shared[stage, :])
T.tcgen05_sf_warp_transpose(SFB_shared[stage, :])
T.fence_proxy_async()
T.mbarrier_arrive(with_sf_full[stage], 0)
sched.next_tile()

elif 128 <= tx < 256:
for w in T.unroll(waves):
cluster_id = block_id // 2
tile_id = num_clusters * w + cluster_id
bx_cluster = (tile_id // group_size) % m_clusters
bx = bx_cluster * 2 + cta_id
by = (tile_id % group_size) + (tile_id // group_size) // m_clusters * group_size

if bx * block_M < M and by * block_N < N:
T.mbarrier_wait_parity(tmem_full, w & 1)
T.copy(C_tmem, C_local)
T.mbarrier_arrive(tmem_empty, 0)

if use_tma_store:
for i in T.unroll(T.ceildiv(block_N, store_block_N)):
T.copy(C_local[:, i * store_block_N : (i + 1) * store_block_N], C_shared)
T.copy(C_shared, C[bx * block_M, by * block_N + i * store_block_N])
else:
T.copy(C_local, C_local_cast)
T.copy(C_local_cast, C[bx * block_M, by * block_N])
sched = T.PersistentTileScheduler("sched_epi", m_blocks, n_blocks, swizzle_size=group_size, cluster_size=cluster_size)
sched.init(block_id // cluster_size)
while sched.valid():
bx = sched.m_idx[0] * cluster_size + cta_id
by = sched.n_idx[0]

T.mbarrier_wait_parity(tmem_full, sched.current_iter[0] & 1)
T.copy(C_tmem, C_local)
T.mbarrier_arrive(tmem_empty, 0)

if use_tma_store:
for i in T.unroll(T.ceildiv(block_N, store_block_N)):
T.copy(C_local[:, i * store_block_N : (i + 1) * store_block_N], C_shared)
T.copy(C_shared, C[bx * block_M, by * block_N + i * store_block_N])
else:
T.copy(C_local, C_local_cast)
T.copy(C_local_cast, C[bx * block_M, by * block_N])
sched.next_tile()
return C


Expand Down
Loading
Loading