Skip to content

Commit a954ff9

Browse files
RissyRanGoogle-ML-Automation
authored andcommitted
Fix ring ragged unsort gradient and clamp local group size.
This change fixes issues in the ragged sort implementation when using ring of experts: - In `moe.py`, clamp `local_group_size` to `buffer_size` to prevent overflow. - In `ragged_sort.py`, apply mask to `grad_sorted_tokens` in `_ring_ragged_unsort` backward pass to ensure gradients outside shard boundaries are zeroed out. - In `ops.py` (megablox), add `group_offset` to `_gmm_fwd` and `_gmm_bwd`. PiperOrigin-RevId: 949759227
1 parent e10af2f commit a954ff9

3 files changed

Lines changed: 13 additions & 0 deletions

File tree

src/maxtext/kernels/megablox/ops.py

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -204,6 +204,7 @@ def _gmm_fwd(
204204
),
205205
preferred_element_type=preferred_element_type,
206206
partial_sum=partial_sum,
207+
group_offset=group_offset,
207208
)
208209
elif use_tokamax_backend:
209210
# manual_axis_type is for gmm with shard_map check_vma=True, needs tokamax > 0.0.12
@@ -331,6 +332,7 @@ def _gmm_bwd(
331332
tile_n=tiling[5],
332333
),
333334
preferred_element_type=lhs_dtype,
335+
group_offset=group_offset,
334336
)
335337

336338
# tgmm_v2_op requires lhs and rhs to have the same dtype.

src/maxtext/kernels/ragged/ragged_sort.py

Lines changed: 8 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -386,6 +386,10 @@ def _ring_ragged_unsort_bwd(res, g_out):
386386
flops_override=gather_flops_override,
387387
bytes_accessed_override=gather_bytes_accessed_override,
388388
)
389+
# Mask out gradients that correspond to elements outside the valid shard
390+
# output range.
391+
mask = (jnp.arange(n) >= shard_output_start) & (jnp.arange(n) < shard_output_end)
392+
grad_sorted_tokens = jnp.where(mask[:, None], grad_sorted_tokens, 0.0)
389393
else:
390394
# Slice the inverse permutation to match the packed local buffer.
391395
padded_idx_inv = jnp.pad(idx_inv, (0, buffer_size))
@@ -405,6 +409,10 @@ def _ring_ragged_unsort_bwd(res, g_out):
405409
flops_override=gather_flops_override,
406410
bytes_accessed_override=gather_bytes_accessed_override,
407411
)
412+
# Mask out gradients for elements beyond the valid limit of the local buffer.
413+
limit = jnp.minimum(shard_output_end - shard_output_start, buffer_size)
414+
mask = jnp.arange(buffer_size) < limit
415+
grad_sorted_tokens = jnp.where(mask[:, None], grad_sorted_tokens, 0.0)
408416
return grad_sorted_tokens, None, None, None
409417

410418
_ring_ragged_unsort.defvjp(_ring_ragged_unsort_fwd, _ring_ragged_unsort_bwd)

src/maxtext/layers/moe.py

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -919,6 +919,9 @@ def permute(
919919
local_num_experts,
920920
axis=0,
921921
)
922+
# Clamp local_group_size to buffer_size to ensure we don't exceed buffer
923+
# capacity by leveraging the helper _truncate_matrix.
924+
local_group_size = _truncate_matrix(local_group_size[:, None], buffer_size)[:, 0]
922925
expert_indices = jnp.arange(local_num_experts)
923926
sorted_experts = jnp.repeat(
924927
expert_indices,

0 commit comments

Comments
 (0)