Replay AutoEP's cached DeepEP dispatch in the layout it was built with (#8636)
AutoEP's combine backward replays the forward DeepEP dispatch against
the cached handle (`DeepEPExchange.dispatch_with_handle`).
In DeepEP v2, `ElasticBuffer.dispatch` takes `do_expand` from its
argument, which defaults to `False`, **even when a cached handle is
supplied**: the cached path reuses the handle's `topk_idx`, expert count
and alignment, but not its `do_expand`. `combine`, by contrast, passes
`handle.do_expand` (`deep_ep/buffers/elastic.py`). AutoEP's forward
dispatch uses `do_expand=True` so that arrivals are grouped per expert
for the grouped GEMM, which means the replay was asking for the
unexpanded layout and scattering gradient rows into a different layout
than the one the experts produced.
The replay now passes `do_expand=handle.do_expand`.
### Tests
- New CPU tests in `TestCachedDispatchLayout`
(`tests/unit/module_inject/test_auto_ep_comm.py`):
- `deepep_combine` forward, backward and two SGD steps run against a
buffer that follows DeepEP's behavior (a cached dispatch takes
`do_expand` from its argument, combine from the handle). The combined
rows, the row and weight gradients and the updated weights must match an
`index_add` oracle bit-exactly in both layouts. On master the expanded
case fails, with 11 of 12 row-gradient elements wrong; with this change
it passes.
- The forward dispatch asks for the expanded layout that the grouped
GEMM needs.
The whole file passes (45 tests).
- On H100,
`test_autoep_deepep_parity.py::test_cleanup_matches_legacy_preparation[False-True]`
(activation checkpointing off, skewed routing) fails on master and
passes with this change.
- A diagnostic on 16 H100s compared the replayed rows with the
forward's. The replay reproduced the forward row order on 0 of 16 ranks
before this change and on 16 of 16 after it, and the relative error of
the adjoint identity `<combine(rows), g> = <rows, replay(g)>` fell from
0.69 to 1.1e-5.
`test_cleanup_matches_legacy_preparation[True-False]` (non-reentrant
activation checkpointing) fails on master with or without this change,
for a separate reason. Non-reentrant checkpointing reruns the dispatch
during recompute, DeepEP's receive order depends on arrival timing
unless its deterministic receive order is enabled, and the backward
still replays against the original forward's handle. In a same-seed
comparison of checkpointing on against off, gradients were bit-exact
with reentrant checkpointing, with DeepEP's deterministic receive order
and with the collective backend, and far apart with DeepEP's default
order under non-reentrant checkpointing. That is not addressed here.
Signed-off-by: yh0903 <helloyu0903@gmail.com>
Co-authored-by: pengdurice <pengduhit@gmail.com>