Skip to content

Commit b0f74ac

Browse files
authored
fix(megatron): isolate Mamba state across packed sequences (#2056)
## Summary Megatron Core derives PackedSeqParams.seq_idx only when total_tokens is set. Mamba uses those labels to reset recurrent state between packed documents; without them, state can flow across document boundaries in a packed THD row. Pass the global padded token count when constructing PackedSeqParams. Using the global count keeps the labels aligned with cu_seqlens_q_padded under context parallelism. ## Testing - uv run --no-sync pytest -q tests/backends/skyrl_train/distributed/test_preprocess_packed_seqs_multiseq.py::TestSubSeqLengths::test_multiseq_row_emits_padded_cu_seqlens_entries tests/backends/skyrl_train/distributed/test_preprocess_packed_seqs_multiseq.py::TestMultiSeqCPLayout::test_roundtrip_recovers_full_layout_cp2 <!-- CURSOR_SUMMARY --> --- > [!NOTE] > **Medium Risk** > Fixes correctness of Mamba state in packed training (wrong labels can silently corrupt gradients); change is small and localized to packed-seq metadata construction. > > **Overview** > **Sets `total_tokens` on `PackedSeqParams`** when building packed THD inputs in `preprocess_packed_seqs`, using the global padded token count (`cu_seqlens_padded_cpu[-1]`). Megatron Core only derives per-token `seq_idx` when `total_tokens` is present; Mamba layers use those labels to reset recurrent state at document boundaries. > > Without this, **recurrent state can carry across packed sub-sequences** in the same THD row. Using the **global** padded total keeps `seq_idx` aligned with `cu_seqlens_q_padded` under context parallelism. > > Tests extend megatron stubs and assert `total_tokens` matches the padded cumulative length in multi-subseq and CP round-trip cases. > > <sup>Reviewed by [Cursor Bugbot](https://cursor.com/bugbot) for commit 6558f94. Bugbot is set up for automated code reviews on this repo. Configure [here](https://www.cursor.com/dashboard/bugbot).</sup> <!-- /CURSOR_SUMMARY -->
1 parent 6ecd004 commit b0f74ac

4 files changed

Lines changed: 8 additions & 0 deletions

File tree

skyrl/backends/skyrl_train/distributed/megatron/megatron_utils.py

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -525,6 +525,7 @@ def preprocess_packed_seqs(
525525
remain_start:remain_end
526526
]
527527

528+
# Mamba derives per-token document labels from the global padded token count.
528529
packed_seq_params = PackedSeqParams(
529530
qkv_format="thd",
530531
cu_seqlens_q=cu_seqlens_padded,
@@ -533,6 +534,7 @@ def preprocess_packed_seqs(
533534
max_seqlen_kv=max_seqlen_in_batch,
534535
cu_seqlens_q_padded=cu_seqlens_padded,
535536
cu_seqlens_kv_padded=cu_seqlens_padded,
537+
total_tokens=cu_seqlens_padded_cpu[-1],
536538
)
537539
if pre_process:
538540
return input_ids_rmpad.unsqueeze(0), packed_seq_params

tests/backends/skyrl_train/distributed/test_preprocess_packed_seqs_cp.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -34,6 +34,7 @@ class _PackedSeqParams:
3434
max_seqlen_kv: Any = None
3535
cu_seqlens_q_padded: Any = None
3636
cu_seqlens_kv_padded: Any = None
37+
total_tokens: Any = None
3738

3839

3940
_MEGATRON_MODULES = [

tests/backends/skyrl_train/distributed/test_preprocess_packed_seqs_multiseq.py

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -28,6 +28,7 @@ class _PackedSeqParams:
2828
max_seqlen_kv: Any = None
2929
cu_seqlens_q_padded: Any = None
3030
cu_seqlens_kv_padded: Any = None
31+
total_tokens: Any = None
3132

3233

3334
_MEGATRON_MODULES = [
@@ -175,6 +176,7 @@ def test_multiseq_row_emits_padded_cu_seqlens_entries(self):
175176
)
176177

177178
assert params.cu_seqlens_q.tolist() == [0, 16, 32]
179+
assert params.total_tokens == 32
178180
assert packed.shape == (1, 32)
179181
assert packed[0, :3].tolist() == [11, 12, 13]
180182
assert packed[0, 16:20].tolist() == [21, 22, 23, 24]
@@ -345,6 +347,8 @@ def _run_roundtrip(self, tp_size, cp_size, sub_seq_lengths, fp8_enabled=True):
345347
per_rank_out.append(packed_r.to(torch.float32))
346348
params_cp = params_r # cu_seqlens are global, identical across ranks
347349

350+
assert params_cp.total_tokens == params_cp.cu_seqlens_q_padded[-1]
351+
348352
# Every rank's local buffer must be the same length (== total/cp).
349353
for r in range(1, cp_size):
350354
assert per_rank_out[r].shape == per_rank_out[0].shape

tests/train/test_packing_round_trip.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -48,6 +48,7 @@ class _PackedSeqParams:
4848
max_seqlen_kv: Any = None
4949
cu_seqlens_q_padded: Any = None
5050
cu_seqlens_kv_padded: Any = None
51+
total_tokens: Any = None
5152

5253

5354
_MEGATRON_MODULES = [

0 commit comments

Comments
 (0)