DeepSpeed
bc778b8c - Replay AutoEP's cached DeepEP dispatch in the layout it was built with (#8636)

Commit
10 days ago
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>
Author
Parents
Loading