Skip to content
Open
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
Original file line number Diff line number Diff line change
Expand Up @@ -50,6 +50,7 @@
MoEWeightMode,
WGradInputOrder,
WgradSfTensormapConstructor,
gmem_ptr_to_generic,
)
from ..moe_sched_extension import (
WgradScaledGemmSchedExtension,
Expand Down Expand Up @@ -1023,10 +1024,23 @@ class SharedStorage:
sched_pipeline.consumer_release(sched_consumer_state)
sched_consumer_state.advance()

last_acquired_expert = cutlass.Int32(-1)
while work_tile_info.is_valid_tile:
k_tile_cnt = work_tile_info.k_tile_cnt
ext.update_expert_info(offs, work_tile_info.expert_idx)

# The preceding device kernel writes descriptors through the generic
# proxy. Acquire them before TMA reads, including on graph replay.
# Descriptors remain immutable within this kernel, so consecutive
# tiles of the same expert can reuse the acquire.
if work_tile_info.expert_idx != last_acquired_expert:
if cutlass.const_expr(self.input_order == WGradInputOrder.TensorRagged):
cpasync.fence_tma_desc_acquire(gmem_ptr_to_generic(desc_workspace.get_desc_ptr("a", work_tile_info.expert_idx)))
cpasync.fence_tma_desc_acquire(gmem_ptr_to_generic(desc_workspace.get_desc_ptr("b", work_tile_info.expert_idx)))
cpasync.fence_tma_desc_acquire(gmem_ptr_to_generic(desc_workspace.get_desc_ptr("sfa", work_tile_info.expert_idx)))
cpasync.fence_tma_desc_acquire(gmem_ptr_to_generic(desc_workspace.get_desc_ptr("sfb", work_tile_info.expert_idx)))
last_acquired_expert = work_tile_info.expert_idx

real_a, desc_ptr_a = ext.get_gmem_tensor(
"a",
mA_mkl,
Expand Down Expand Up @@ -1458,10 +1472,19 @@ class SharedStorage:
sched_pipeline.consumer_release(sched_consumer_state)
sched_consumer_state.advance()

last_acquired_expert = cutlass.Int32(-1)
while work_tile_info.is_valid_tile:
k_tile_cnt = work_tile_info.k_tile_cnt
ext.update_expert_info(offs, work_tile_info.expert_idx)

# Only the warp issuing TMA stores needs the output descriptor.
# Empty experts still store zeros and must acquire it as well.
if cutlass.const_expr(self.weight_mode == MoEWeightMode.DISCRETE):
if warp_idx == self.epilogue_warp_id[0]:
if work_tile_info.expert_idx != last_acquired_expert:
cpasync.fence_tma_desc_acquire(gmem_ptr_to_generic(desc_workspace.get_desc_ptr("c", work_tile_info.expert_idx)))
last_acquired_expert = work_tile_info.expert_idx

real_c, desc_ptr_c = ext.get_gmem_tensor(
"c",
mC_mnl,
Expand Down
79 changes: 79 additions & 0 deletions test/python/fe_api/grouped_gemm/test_wgrad_tma_descriptor_reuse.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,79 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0

"""A cached wgrad plan must acquire descriptors rewritten between graph nodes."""

import pytest
import torch
import cudnn


@pytest.mark.L0
@pytest.mark.parametrize("discrete", [False, True])
@pytest.mark.parametrize("ragged", [False, True])
@pytest.mark.parametrize("accumulate", [False, True])
def test_wgrad_tma_descriptor_reuse(discrete, ragged, accumulate):
if torch.cuda.get_device_capability() not in ((10, 0), (10, 3)):
pytest.skip("Exercises the SM100 blockscaled wgrad kernel.")

stream = torch.cuda.Stream()
calls, outputs = [], []
# Identical signatures reuse one plan/workspace, but the device helper must
# rebind its descriptors to distinct input, scale and output addresses.
for value, scale_byte, tokens in ((1, 127, (128, 128, 0, 256)), (2, 128, (256, 0, 128, 128))):
a = torch.full((512, 256), value, device="cuda").to(torch.float8_e4m3fn).T
b = torch.ones((512, 256), device="cuda").to(torch.float8_e4m3fn)
# Inputs are constant, so their physical layout is also valid in ragged mode.
sa = torch.full((256, 16), scale_byte, device="cuda", dtype=torch.uint8).view(torch.float8_e8m0fnu)
sb = sa.clone()
offsets = torch.tensor(tokens, device="cuda", dtype=torch.int32).cumsum(0, dtype=torch.int32)
if discrete:
dest = [torch.empty((256, 256), device="cuda", dtype=torch.bfloat16) for _ in tokens]
pointers = torch.tensor([out.data_ptr() for out in dest], device="cuda", dtype=torch.int64)
output_args = dict(output_mode="discrete", wgrad_ptrs=pointers)
else:
dense = torch.empty((4, 256, 256), device="cuda", dtype=torch.bfloat16)
dest = list(dense.unbind())
output_args = dict(output_mode="dense", wgrad_tensor=dense)
calls.append(
dict(
a_tensor=a,
b_tensor=b,
sfa_tensor=sa,
sfb_tensor=sb,
offsets_tensor=offsets,
**output_args,
wgrad_dtype=torch.bfloat16,
acc_dtype=torch.float32,
sf_vec_size=32,
input_order="tensor_ragged" if ragged else "tensor2d",
accumulate_on_output=accumulate,
current_stream=stream.cuda_stream,
)
)
scale = 2 ** (scale_byte - 127)
outputs.extend((out, n * value * scale * scale + int(accumulate)) for out, n in zip(dest, tokens))

def run():
for kwargs in calls:
cudnn.grouped_gemm_wgrad_wrapper_sm100(**kwargs)

def check(fn):
# Reset outside capture on every replay: stale warmup outputs must not
# hide a missing store, and accumulation always starts from a known value.
for out, _ in outputs:
out.fill_(1 if accumulate else -1)
fn()
for out, expected in outputs:
torch.testing.assert_close(out, torch.full_like(out, expected), rtol=0, atol=0)

stream.wait_stream(torch.cuda.current_stream())
with torch.cuda.stream(stream):
for _ in range(3):
check(run)
graph = torch.cuda.CUDAGraph()
with torch.cuda.graph(graph, stream=stream):
run()
for _ in range(8):
check(graph.replay)
torch.cuda.current_stream().wait_stream(stream)
Loading