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
54 changes: 54 additions & 0 deletions tests/pytorch/test_fused_router.py
Original file line number Diff line number Diff line change
Expand Up @@ -1119,6 +1119,60 @@ def test_fused_moe_aux_loss(dtype, num_tokens, num_experts, topk, expert_multipl
torch.testing.assert_close(probs.grad, probs_clone.grad, atol=atol, rtol=rtol)


@pytest.mark.parametrize("path", ["forward", "graph_safe_forward", "backward"])
def test_fused_moe_aux_loss_int64_offsets(path):
num_rows, num_cols = 8_388_609, 256 # Above INT_MAX elements.
bytes_needed = num_rows * num_cols * 2
torch.cuda.empty_cache()
free_bytes, _ = torch.cuda.mem_get_info()
if free_bytes < bytes_needed + (1 << 30):
pytest.skip("Needs 4.0 GiB for BF16 input/output plus 1 GiB headroom")

tokens_per_expert = torch.ones(num_cols, device="cuda", dtype=torch.int32)
try:
if path == "backward":
tokens_per_expert = torch.arange(1, num_cols + 1, device="cuda", dtype=torch.int32)
grad_probs = tex.fused_moe_aux_loss_bwd(
Const_buf=torch.tensor([0.25, 0.0], device="cuda", dtype=torch.float32),
tokens_per_expert=tokens_per_expert,
num_rows=num_rows,
num_cols=num_cols,
grad_aux_loss=torch.tensor(2.0, device="cuda", dtype=torch.bfloat16),
)
expected = (0.5 * tokens_per_expert).to(torch.bfloat16)
for start in range(0, num_rows, 65536):
chunk = grad_probs[start : start + 65536]
torch.testing.assert_close(chunk, expected.expand_as(chunk), atol=0, rtol=0)
else:
probs = torch.zeros((num_rows, num_cols), device="cuda", dtype=torch.bfloat16)
# Nonzero values on both sides of the int32 offset boundary.
probs[-2:] = 1
arguments = dict(
probs=probs,
tokens_per_expert=tokens_per_expert,
num_experts=num_cols,
num_rows=num_rows,
num_cols=num_cols,
topk=1,
coeff=1.0,
)
if path == "forward":
aux_loss, const_buf = tex.fused_moe_aux_loss_fwd(
total_num_tokens=num_cols, **arguments
)
else:
total = torch.tensor(num_cols, device="cuda", dtype=torch.int64)
aux_loss, const_buf = tex.fused_moe_aux_loss_fwd_graph_safe(
total_num_tokens=total, **arguments
)
torch.testing.assert_close(aux_loss, probs.new_tensor(2.0), atol=0, rtol=0)
torch.testing.assert_close(
const_buf, const_buf.new_tensor([1.0 / num_cols, 2.0]), atol=0, rtol=0
)
except torch.cuda.OutOfMemoryError:
pytest.skip("Could not allocate the BF16 input/output tensor")


def test_fused_moe_aux_loss_cuda_graph_capture():
"""CUDA-graph-safe path: total_num_tokens is a device tensor whose value
changes between replays. Forward and backward must both observe the new
Expand Down
21 changes: 12 additions & 9 deletions transformer_engine/common/fused_router/fused_moe_aux_loss.cu
Original file line number Diff line number Diff line change
Expand Up @@ -37,11 +37,12 @@ __global__ void fused_moe_aux_loss_forward_kernel(const DataType* probs,

// Grid-stride over rows so that every row is processed exactly once.
// Each thread processes a subset of columns.
for (int col = threadIdx.x; col < num_cols; col += blockDim.x) {
for (int64_t col = threadIdx.x; col < num_cols; col += blockDim.x) {
CompType col_sum = CompType(0);

// Accumulate probs over the rows assigned to this CTA (grid-stride).
for (int row = blockIdx.x; row < num_rows; row += gridDim.x) {
#pragma unroll 4
for (int64_t row = blockIdx.x; row < num_rows; row += gridDim.x) {
col_sum += CompType(probs[row * num_cols + col]);
}

Expand Down Expand Up @@ -144,9 +145,10 @@ __global__ void fused_moe_aux_loss_forward_kernel_graph_safe(
int num_experts, int num_rows, int num_cols, int topk, float coeff, float* Coeff_buf) {
// Reduction body matches the scalar-input kernel above.
CompType thread_sum = CompType(0);
for (int col = threadIdx.x; col < num_cols; col += blockDim.x) {
for (int64_t col = threadIdx.x; col < num_cols; col += blockDim.x) {
CompType col_sum = CompType(0);
for (int row = blockIdx.x; row < num_rows; row += gridDim.x) {
#pragma unroll 4
for (int64_t row = blockIdx.x; row < num_rows; row += gridDim.x) {
col_sum += CompType(probs[row * num_cols + col]);
}
col_sum *= CompType(tokens_per_expert[col]);
Expand Down Expand Up @@ -231,17 +233,18 @@ __global__ void fused_moe_aux_loss_backward_kernel(const float* Const_buf,
const IndexType* tokens_per_expert, int num_rows,
int num_cols, DataType* grad_aux_loss,
DataType* grad_probs) {
int global_warp_num = gridDim.x * blockDim.x / kThreadsPerWarp;
int global_warp_id = (blockIdx.x * blockDim.x + threadIdx.x) / kThreadsPerWarp;
int64_t global_warp_num = static_cast<int64_t>(gridDim.x) * blockDim.x / kThreadsPerWarp;
int64_t global_warp_id =
(static_cast<int64_t>(blockIdx.x) * blockDim.x + threadIdx.x) / kThreadsPerWarp;
int lane_id = threadIdx.x % kThreadsPerWarp;

// Loop: for all positions in each row
for (int i = lane_id; i < num_cols; i += kThreadsPerWarp) {
for (int64_t i = lane_id; i < num_cols; i += kThreadsPerWarp) {
float C_coeff = Const_buf[0];
CompType tokens_per_expert_i = static_cast<CompType>(tokens_per_expert[i]);
CompType grad_aux_loss_value = static_cast<CompType>(grad_aux_loss[0]);
// Loop: for all rows
for (int j = global_warp_id; j < num_rows; j += global_warp_num) {
for (int64_t j = global_warp_id; j < num_rows; j += global_warp_num) {
grad_probs[j * num_cols + i] = C_coeff * tokens_per_expert_i * grad_aux_loss_value;
}
}
Expand All @@ -254,7 +257,7 @@ void fused_moe_aux_loss_backward_kernel_launcher(const float* Const_buf,
DataType* grad_probs, cudaStream_t stream) {
// Meta data for the kernel
int block_size = 256;
int grid_size = (num_rows + block_size - 1) / block_size;
int grid_size = (static_cast<int64_t>(num_rows) + block_size - 1) / block_size;
fused_moe_aux_loss_backward_kernel<DataType, IndexType><<<grid_size, block_size, 0, stream>>>(
Const_buf, tokens_per_expert, num_rows, num_cols, grad_aux_loss, grad_probs);
NVTE_CHECK_CUDA(cudaGetLastError());
Expand Down
Loading