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
128 changes: 99 additions & 29 deletions flash_qla/ops/gated_delta_rule/chunk/blackwell/fused_fwd.py
Original file line number Diff line number Diff line change
Expand Up @@ -40,6 +40,7 @@ def tilelang_fused_chunk_gdr_fwd(
store_h,
store_o,
is_varlen,
has_incomplete_tile,
is_cp,
state_v_first,
block_DV=128,
Expand All @@ -50,6 +51,16 @@ def tilelang_fused_chunk_gdr_fwd(
raw_batch_size = T.dynamic("raw_batch_size")
block_S = chunk_size

# Producer warp layout for the 512-thread CTA. Threads 0..383 are
# consumers; the remaining four warps run tcgen, data staging, and store.
WARP_SIZE = 32
CONSUMER_WARP_COUNT = 12
TCGEN_PRODUCER_END = (CONSUMER_WARP_COUNT + 1) * WARP_SIZE
QK_PRODUCER_BEGIN = TCGEN_PRODUCER_END
QK_PRODUCER_END = QK_PRODUCER_BEGIN + WARP_SIZE
VA_PRODUCER_BEGIN = QK_PRODUCER_END
VA_PRODUCER_END = VA_PRODUCER_BEGIN + WARP_SIZE

if is_varlen:
q_shape = (1, num_tokens, Hg, DK)
k_shape = (1, num_tokens, Hg, DK)
Expand Down Expand Up @@ -472,7 +483,7 @@ def tilelang_fused_chunk_gdr_fwd_kernel(
else:
T.set_max_nreg(PRODUCER_NREG, 0)

if tx < 384 + 32:
if tx < TCGEN_PRODUCER_END:
for i_s in T.serial(num_iters):
T.barrier_arrive(bar_0)

Expand Down Expand Up @@ -595,7 +606,7 @@ def tilelang_fused_chunk_gdr_fwd_kernel(
use_2cta=False,
)

elif tx < 384 + 64:
elif tx < QK_PRODUCER_END:
for i_s in T.serial(num_unmasked_iters):
T.barrier_wait(data_is_free[i_s % 2], (i_s // 2 + 1) % 2)
left = seq_start_idx + i_s * block_S
Expand Down Expand Up @@ -628,41 +639,50 @@ def tilelang_fused_chunk_gdr_fwd_kernel(

T.barrier_arrive(data_is_ready[i_s % 2])

if num_unmasked_iters < num_iters:
# Specialize pure full-tile batches without the larger
# async-tail body. The caller derives this from the chunk
# count that it already synchronizes to the host.
if has_incomplete_tile and num_unmasked_iters < num_iters:
T.barrier_wait(data_is_free[num_unmasked_iters % 2], (num_unmasked_iters // 2 + 1) % 2)
left = seq_start_idx + num_unmasked_iters * block_S
right = left + block_S

# Load Q
for j_s, j_k in T.Parallel(block_S, DK):
if left + j_s < seq_end_idx:
q_shared[num_unmasked_iters % 2, j_s, j_k] = q[batch_idx, left + j_s, bhg, j_k]
else:
q_shared[num_unmasked_iters % 2, j_s, j_k] = 0
# Load K
for j_s, j_k in T.Parallel(block_S, DK):
if left + j_s < seq_end_idx:
k_shared[num_unmasked_iters % 2, j_s, j_k] = k[batch_idx, left + j_s, bhg, j_k]
else:
k_shared[num_unmasked_iters % 2, j_s, j_k] = 0
# Load V
for j_s, j_v in T.Parallel(block_S, block_DV):
if left + j_s < seq_end_idx:
v_shared[num_unmasked_iters % 2, j_s, j_v] = v[batch_idx, left + j_s, bh, DV_start + j_v]
else:
v_shared[num_unmasked_iters % 2, j_s, j_v] = 0
# Load A
for j_s, j_t in T.Parallel(block_S, block_S):
if left + j_s < seq_end_idx:
a_shared[num_unmasked_iters % 2, j_s, j_t] = a[batch_idx, left + j_s, bh, j_t]
else:
a_shared[num_unmasked_iters % 2, j_s, j_t] = 0
# TMA remains the fast path for complete 64-row tiles.
# An incomplete packed-varlen tile must predicate each
# row so it does not read into the next sequence. This
# warp stages Q/K with zero-fill cp.async; the following
# producer warp stages V/A in parallel.
tail_lane = tx - QK_PRODUCER_BEGIN
for i_copy in T.serial(block_S * DK // (WARP_SIZE * 8)):
copy_offset = (i_copy * WARP_SIZE + tail_lane) * 8
j_s = copy_offset // DK
j_k = copy_offset % DK
T.ptx_cp_async(
T.access_ptr(q_shared[num_unmasked_iters % 2, j_s, j_k], "w", 8),
T.access_ptr(q[batch_idx, left + j_s, bhg, j_k], "r", 8),
8,
left + j_s < seq_end_idx,
)
T.ptx_commit_group()

for i_copy in T.serial(block_S * DK // (WARP_SIZE * 8)):
copy_offset = (i_copy * WARP_SIZE + tail_lane) * 8
j_s = copy_offset // DK
j_k = copy_offset % DK
T.ptx_cp_async(
T.access_ptr(k_shared[num_unmasked_iters % 2, j_s, j_k], "w", 8),
T.access_ptr(k[batch_idx, left + j_s, bhg, j_k], "r", 8),
8,
left + j_s < seq_end_idx,
)
T.ptx_commit_group()
T.ptx_wait_group(0)
T.fence_proxy_async()

T.barrier_arrive(data_is_ready[num_unmasked_iters % 2])

elif tx < 384 + 96:
for i_s in T.serial(num_iters):
elif tx < VA_PRODUCER_END:
for i_s in T.serial(num_unmasked_iters):
T.barrier_wait(data_is_free[i_s % 2], (i_s // 2 + 1) % 2)
left = seq_start_idx + i_s * block_S
right = left + block_S
Expand Down Expand Up @@ -690,6 +710,53 @@ def tilelang_fused_chunk_gdr_fwd_kernel(

T.barrier_arrive(data_is_ready[i_s % 2])

if has_incomplete_tile and num_unmasked_iters < num_iters:
T.barrier_wait(data_is_free[num_unmasked_iters % 2], (num_unmasked_iters // 2 + 1) % 2)
left = seq_start_idx + num_unmasked_iters * block_S
tail_lane = tx - VA_PRODUCER_BEGIN

for i_copy in T.serial(block_S * block_DV // (WARP_SIZE * 8)):
copy_offset = (i_copy * WARP_SIZE + tail_lane) * 8
j_s = copy_offset // block_DV
j_v = copy_offset % block_DV
T.ptx_cp_async(
T.access_ptr(v_shared[num_unmasked_iters % 2, j_s, j_v], "w", 8),
T.access_ptr(v[batch_idx, left + j_s, bh, DV_start + j_v], "r", 8),
8,
left + j_s < seq_end_idx,
)
T.ptx_commit_group()

for i_copy in T.serial(block_S * block_S // (WARP_SIZE * 8)):
copy_offset = (i_copy * WARP_SIZE + tail_lane) * 8
j_s = copy_offset // block_S
j_t = copy_offset % block_S
T.ptx_cp_async(
T.access_ptr(a_shared[num_unmasked_iters % 2, j_s, j_t], "w", 8),
T.access_ptr(a[batch_idx, left + j_s, bh, j_t], "r", 8),
8,
left + j_s < seq_end_idx,
)
T.ptx_commit_group()

# Scalar loads execute while the V/A copies are in flight.
for j_s in T.Parallel(block_S):
if left + j_s < seq_end_idx:
b_shared[num_unmasked_iters % 2, j_s] = b[batch_idx, left + j_s, bh]
else:
b_shared[num_unmasked_iters % 2, j_s] = 0
# Preserve the original tail semantics by extending
# gamma with the last valid value.
for j_s in T.Parallel(block_S):
if left + j_s < seq_end_idx:
g_shared[num_unmasked_iters % 2, j_s] = g[batch_idx, left + j_s, bh]
else:
g_shared[num_unmasked_iters % 2, j_s] = g[batch_idx, seq_end_idx - 1, bh]

T.ptx_wait_group(0)
T.fence_proxy_async()
T.barrier_arrive(data_is_ready[num_unmasked_iters % 2])

else:
for i_s in T.serial(num_unmasked_iters):
right = seq_start_idx + i_s * block_S
Expand Down Expand Up @@ -808,9 +875,11 @@ def fused_gdr_fwd(
)
seqlen_dtype = torch.int32
is_varlen = False
has_incomplete_tile = num_tokens % chunk_size != 0
else:
real_batch_size = len(cu_seqlens) - 1
chunk_offsets, num_chunks = prepare_chunk_offsets(cu_seqlens, chunk_size)
has_incomplete_tile = num_chunks * chunk_size != num_tokens
chunk_offsets = chunk_offsets.to(cu_seqlens.dtype)
num_chunks = num_chunks if output_h else 0
seqlen_dtype = cu_seqlens.dtype
Expand Down Expand Up @@ -884,6 +953,7 @@ def fused_gdr_fwd(
store_h=output_h,
store_o=output_o,
is_varlen=is_varlen,
has_incomplete_tile=has_incomplete_tile,
is_cp=is_cp,
state_v_first=state_v_first,
block_DV=block_DV,
Expand Down
92 changes: 68 additions & 24 deletions flash_qla/ops/gated_delta_rule/chunk/blackwell/prepare_h.py
Original file line number Diff line number Diff line change
Expand Up @@ -32,6 +32,7 @@ def tilelang_prepare_h(
store_final_state,
store_h,
is_varlen,
has_incomplete_tile,
is_cp,
state_v_first,
num_stages=2,
Expand All @@ -41,6 +42,16 @@ def tilelang_prepare_h(
num_chunks = T.dynamic("num_chunks")
block_S = chunk_size

# Producer warp layout for the 512-thread CTA. Threads 0..383 are
# consumers; the remaining four warps run tcgen, data staging, and store.
WARP_SIZE = 32
CONSUMER_WARP_COUNT = 12
TCGEN_PRODUCER_END = (CONSUMER_WARP_COUNT + 1) * WARP_SIZE
K_PRODUCER_BEGIN = TCGEN_PRODUCER_END
K_PRODUCER_END = K_PRODUCER_BEGIN + WARP_SIZE
VA_PRODUCER_BEGIN = K_PRODUCER_END
VA_PRODUCER_END = VA_PRODUCER_BEGIN + WARP_SIZE

if is_varlen:
k_shape = (1, num_tokens, Hg, DK)
v_shape = (1, num_tokens, H, DV)
Expand Down Expand Up @@ -462,7 +473,7 @@ def tilelang_prepare_h_kernel(
else:
T.set_max_nreg(PRODUCER_NREG, 0)

if tx < 384 + 32:
if tx < TCGEN_PRODUCER_END:
for i_s in T.serial(num_iters):
T.barrier_arrive(bar_0)

Expand Down Expand Up @@ -552,7 +563,7 @@ def tilelang_prepare_h_kernel(
# use_2cta=False,
# )

elif tx < 384 + 64:
elif tx < K_PRODUCER_END:
for i_s in T.serial(num_iters):
T.barrier_wait(
data_is_free[i_s % num_stages], (i_s // num_stages + 1) % 2
Expand All @@ -567,50 +578,75 @@ def tilelang_prepare_h_kernel(
k_shared[i_s % num_stages, :, :],
barrier=data_is_ready[i_s % num_stages],
)
else:
for j_s, j_k in T.Parallel(block_S, DK):
if left + j_s < seq_end_idx:
k_shared[i_s % num_stages, j_s, j_k] = k[batch_idx, left + j_s, bhg, j_k]
else:
k_shared[i_s % num_stages, j_s, j_k] = 0
# Load V
elif has_incomplete_tile:
# Packed-varlen tails cannot use an unpredicated
# TMA tile without reading the next sequence.
tail_lane = tx - K_PRODUCER_BEGIN
for i_copy in T.serial(block_S * DK // (WARP_SIZE * 8)):
copy_offset = (i_copy * WARP_SIZE + tail_lane) * 8
j_s = copy_offset // DK
j_k = copy_offset % DK
T.ptx_cp_async(
T.access_ptr(k_shared[i_s % num_stages, j_s, j_k], "w", 8),
T.access_ptr(k[batch_idx, left + j_s, bhg, j_k], "r", 8),
8,
left + j_s < seq_end_idx,
)
T.ptx_commit_group()

# Complete tiles keep the faster TMA path for V/A.
# The following producer warp owns incomplete tails.
if right <= seq_end_idx:
T.tma_copy(
v[batch_idx, left:right, bh, 0:DV],
v_shared[i_s % num_stages, :, :],
barrier=data_is_ready[i_s % num_stages],
)
else:
for j_s, j_v in T.Parallel(block_S, DV):
if left + j_s < seq_end_idx:
v_shared[i_s % num_stages, j_s, j_v] = v[batch_idx, left + j_s, bh, j_v]
else:
v_shared[i_s % num_stages, j_s, j_v] = 0
# Load A
if right <= seq_end_idx:
T.tma_copy(
a[batch_idx, left:right, bh, 0:block_S],
a_shared[i_s % num_stages, :, :],
barrier=data_is_ready[i_s % num_stages],
)
else:
for j_s, j_t in T.Parallel(block_S, block_S):
if left + j_s < seq_end_idx:
a_shared[i_s % num_stages, j_s, j_t] = a[batch_idx, left + j_s, bh, j_t]
else:
a_shared[i_s % num_stages, j_s, j_t] = 0
elif has_incomplete_tile:
T.ptx_wait_group(0)
T.fence_proxy_async()

T.barrier_arrive(data_is_ready[i_s % num_stages])

elif tx < 384 + 96:
elif tx < VA_PRODUCER_END:
for i_s in T.serial(num_iters):
T.barrier_wait(
data_is_free[i_s % num_stages], (i_s // num_stages + 1) % 2
)
left = seq_start_idx + i_s * block_S
right = left + block_S

if has_incomplete_tile and right > seq_end_idx:
tail_lane = tx - VA_PRODUCER_BEGIN
for i_copy in T.serial(block_S * DV // (WARP_SIZE * 8)):
copy_offset = (i_copy * WARP_SIZE + tail_lane) * 8
j_s = copy_offset // DV
j_v = copy_offset % DV
T.ptx_cp_async(
T.access_ptr(v_shared[i_s % num_stages, j_s, j_v], "w", 8),
T.access_ptr(v[batch_idx, left + j_s, bh, j_v], "r", 8),
8,
left + j_s < seq_end_idx,
)
T.ptx_commit_group()

for i_copy in T.serial(block_S * block_S // (WARP_SIZE * 8)):
copy_offset = (i_copy * WARP_SIZE + tail_lane) * 8
j_s = copy_offset // block_S
j_t = copy_offset % block_S
T.ptx_cp_async(
T.access_ptr(a_shared[i_s % num_stages, j_s, j_t], "w", 8),
T.access_ptr(a[batch_idx, left + j_s, bh, j_t], "r", 8),
8,
left + j_s < seq_end_idx,
)
T.ptx_commit_group()

# Load gamma
if right <= seq_end_idx:
for j_s in T.Parallel(block_S):
Expand Down Expand Up @@ -642,6 +678,11 @@ def tilelang_prepare_h_kernel(
else:
b_shared[i_s % num_stages, j_s] = 0

# Scalar loads execute while tail V/A copies are in flight.
if has_incomplete_tile and right > seq_end_idx:
T.ptx_wait_group(0)
T.fence_proxy_async()

T.barrier_arrive(data_is_ready[i_s % num_stages])

else:
Expand Down Expand Up @@ -696,10 +737,12 @@ def fused_gdr_h(
(batch_size + 1), dtype=torch.int32, device=k.device
)
is_varlen = False
has_incomplete_tile = num_tokens % chunk_size != 0
is_cp = False
else:
real_batch_size = len(cu_seqlens) - 1
chunk_offsets, num_chunks = prepare_chunk_offsets(cu_seqlens, chunk_size)
has_incomplete_tile = num_chunks * chunk_size != num_tokens
chunk_offsets = chunk_offsets.to(cu_seqlens.dtype)
num_chunks = num_chunks if output_h else 0
is_varlen = True
Expand Down Expand Up @@ -757,6 +800,7 @@ def fused_gdr_h(
store_final_state=output_final_state,
store_h=output_h,
is_varlen=is_varlen,
has_incomplete_tile=has_incomplete_tile,
is_cp=is_cp,
state_v_first=state_v_first,
)
Expand Down