AI-Hypercomputer / AI-Hypercomputer/maxtext

get_context_partition_and_sub_seq reads context_autoregressive, so MoE subgroup splitting never engages under ici_context_parallelism

Aperta
#4,931 0 commenti 0 reazioni 0 assegnatari Vedi su GitHub
Lingua principale
Python
Stelle
2.4k
Fork
607
Merge medio
2g 19h
PR unite (30g)
158

Descrizione

`moe.py` already has machinery to keep a sequence-sharded input through the MoE block. It reads a different mesh axis from the one `ici_context_parallelism` sets, so it never engages.

## The mismatch

`moe.py:707` is the only `mesh.shape.get("context...")` in the file:

```python
def get_context_autoregressive_parallelism_size(self):
return self.mesh.shape.get("context_autoregressive", 1)
```

`context`, `context_usp_ulysses` and `context_autoregressive` are three separate mesh axes (`base.yml:531`) with three separate knobs: `ici_context_parallelism`, `ici_context_usp_ulysses_parallelism`, `ici_context_autoregressive_parallelism`.

`get_context_partition_and_sub_seq` (`moe.py:2488`) calls it:

```python
def get_context_partition_and_sub_seq(self, seq_len):
cp = self.get_context_autoregressive_parallelism_size()
if seq_len % cp != 0:
cp = 1
sub_seq = seq_len // cp
return cp, sub_seq
```

Under `ici_context_parallelism` that returns `cp = 1` and `sub_seq = seq_len`, so the subgroup path collapses to the unsharded one.

## What is already built for the split

The dense and capacity path carries the subgroup shape end to end:

- `(batch, cp, sub_seq, emb)` reshape and sharding constraint at `:2892`
- `"BNSM,BNSEC -> EBNCM"` dispatch einsums at `:2887`
- `generate_masks_subgroup` threading `(batch, cp, sub_seq, ...)` through every mask reshape at `:2504-2583`

## Scope

Fixing the axis read would help configurations on that path. It wouldn't help megablox, which flattens at `:883-885` whatever the mesh says:

```python
bsz_times_seq_len = inputs_shape[0] * inputs_shape[1]
inputs_2d = jnp.reshape(inputs, (bsz_times_seq_len, inputs_shape[2]))
```

`fused_moe_matmul` at `:3065` flattens the same way.

## What I'm running instead

I substitute a plain `MlpBlock` for `RoutedMoE` when context parallelism is active. That keeps the sequence sharding and lets a long-context run proceed. It isn't the fix. On a v5p at sequence 8192 the grouped matmul reaches 29.01% MFU and `MlpBlock` 27.82%, so the swap costs throughput, and it only works for a dense configuration, which has nothing to route.

An xplane capture of that run shows `gmm` custom-calls doing the MLP work, so I'm on the megablox path and the axis fix alone wouldn't have helped me.

## Reproduce

Run a model that uses `RoutedMoE` with `ici_context_parallelism=4` and a sequence that doesn't fit on one device. The MoE block sees a full-length sequence whatever the mesh says.

One caveat on reading the memory, because it caught me elsewhere: the local batch is `per_device_batch_size × ici_context_parallelism`, so raising the context degree at a fixed `per_device_batch_size` grows the batch and hides the effect. Set `per_device_batch_size = 1/ici_context_parallelism` before attributing anything to this mismatch. See #4933.

## Not verified

I read the code, I didn't run it. Whether changing the axis read is enough depends on what else the subgroup path assumes about `context_autoregressive`. Treat this as a located mismatch rather than a proposed patch.

cc @mmcsa

Guida per i contributori

Apri la guida per i contributori

Valutazione

Questa issue non è ancora stata valutata.

Ricevi le nuove issue nella tua casella

Un breve riepilogo di issue GitHub adatte ai principianti.