From 1e97d2ed70b3cea57548273c4338033e77c5137d Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Fri, 25 Sep 2026 15:41:17 +0000 Subject: [PATCH 1/4] Release HybridEP dispatcher state after combine MCore's HybridEP manager keeps routing_map, token_probs and dispatched_probs after combine. The dispatched probabilities keep each MoE layer's checkpoint graph, its recomputed input and that input's gradient alive until the layer's next dispatch. Clear them after combine_postprocess for exact flex dispatchers with a HybridEP manager, as #858 and #861 do for the EP1 all-to-all. Co-Authored-By: Claude Opus 5.5 (1M context) --- src/art/trainer_rank/_impl.py | 30 +++ .../test_dispatcher_graph_retention.py | 188 ++++++++++++++++++ 2 files changed, 218 insertions(+) diff --git a/src/art/trainer_rank/_impl.py b/src/art/trainer_rank/_impl.py index 308738c28..572ad7b67 100644 --- a/src/art/trainer_rank/_impl.py +++ b/src/art/trainer_rank/_impl.py @@ -1250,6 +1250,19 @@ def _moe_combine_postprocess(dispatcher: Any, hidden_states: torch.Tensor): return result +def _release_hybridep_state(manager: Any) -> None: + # The dispatched probabilities hold this layer's checkpoint graph, with its + # recomputed input and that input's gradient, until the next dispatch. + manager.routing_map = manager.token_probs = manager.dispatched_probs = None + + +def _hybridep_combine_postprocess(dispatcher: Any, hidden_states: torch.Tensor): + result = type(dispatcher).combine_postprocess(dispatcher, hidden_states) + # Backward saves its own; the next setup and dispatch recreate these. + _release_hybridep_state(dispatcher._comm_manager) + return result + + def _configure_moe_dispatcher_caches(model: Sequence[torch.nn.Module]) -> None: for chunk in model: for module in chunk.modules(): @@ -1258,8 +1271,25 @@ def _configure_moe_dispatcher_caches(model: Sequence[torch.nn.Module]) -> None: continue from megatron.core.transformer.moe.token_dispatcher import ( MoEAlltoAllTokenDispatcher, + MoEFlexTokenDispatcher, + _HybridEPManager, ) + if ( + type(dispatcher) is MoEFlexTokenDispatcher + and type(getattr(dispatcher, "_comm_manager", None)) is _HybridEPManager + and "combine_postprocess" not in vars(dispatcher) + and getattr(dispatcher.config, "cuda_graph_impl", "none") == "none" + ): + # HybridEP keeps its routing inputs and dispatched probabilities + # after combine; CUDA graph capture reads them back instead. + setattr( + dispatcher, + "combine_postprocess", + partial(_hybridep_combine_postprocess, dispatcher), + ) + _release_hybridep_state(dispatcher._comm_manager) + continue if type(dispatcher) is not MoEAlltoAllTokenDispatcher or ( "dispatch_preprocess" in vars(dispatcher) ): diff --git a/tests/integration/megatron/model_support/test_dispatcher_graph_retention.py b/tests/integration/megatron/model_support/test_dispatcher_graph_retention.py index 3968367a0..427777bd5 100644 --- a/tests/integration/megatron/model_support/test_dispatcher_graph_retention.py +++ b/tests/integration/megatron/model_support/test_dispatcher_graph_retention.py @@ -625,3 +625,191 @@ def test_dispatcher_custom_combine_is_preserved(): assert dispatcher.probs.numel() == 0 assert dispatcher.routing_map is not None assert dispatcher.reversed_local_input_permutation_mapping is not None + + +def _hybridep_dispatch(*, x, routing_map, probs, num_local_experts, **_): + # CPU stand-in for HybridEP's fused dispatch: routed rows, their + # differentiable probabilities, per-expert counts and a combine handle. + rows, columns = routing_map.nonzero(as_tuple=True) + counts = routing_map.sum(0) + return x[rows], probs[rows, columns], None, counts, (rows, x.shape[0]) + + +def _hybridep_combine(*, x, handle, **_): + rows, tokens = handle + return x.new_zeros(tokens, x.shape[-1]).index_add(0, rows, x) + + +@pytest.fixture +def cpu_hybridep(cpu_checkpoint_rng, monkeypatch): + from megatron.core.transformer.moe import token_dispatcher + + monkeypatch.setattr(token_dispatcher, "hybrid_ep_dispatch", _hybridep_dispatch) + monkeypatch.setattr(token_dispatcher, "hybrid_ep_combine", _hybridep_combine) + + +def _flex_dispatcher(manager: str = "hybridep") -> Any: + # Upstream flex dispatcher and HybridEP manager methods, with CPU fused + # kernels and without distributed initialization. + from megatron.core.transformer.moe.token_dispatcher import ( + _DeepepManager, + _HybridEPManager, + ) + + config = SimpleNamespace( + cuda_graph_impl="none", + fp8=None, + fp4=None, + moe_hybridep_num_sms=1, + moe_router_topk=2, + ) + dispatcher: Any = object.__new__(MoEFlexTokenDispatcher) + dispatcher.config = config + dispatcher.tp_size = dispatcher.ep_size = 1 + dispatcher.num_local_experts = 4 + comm: Any = object.__new__( + _HybridEPManager if manager == "hybridep" else _DeepepManager + ) + comm.group = None + comm.num_local_experts = comm.num_experts = 4 + comm.config = config + comm.drop_and_pad = False + comm.num_permuted_tokens = comm.pad_multiple = comm.handle = None + comm.token_probs = None + dispatcher._comm_manager = comm + return dispatcher + + +class _FlexRouterLayer(torch.nn.Module): + def __init__(self, manager: str = "hybridep"): + super().__init__() + self.weight = torch.nn.Parameter(torch.randn(8, 4)) + self.token_dispatcher = _flex_dispatcher(manager) + + def forward(self, value): + probs = (value.reshape(-1, 8) @ self.weight).softmax(-1) + routing = torch.zeros_like(probs, dtype=torch.bool) + routing.scatter_(1, probs.topk(2, dim=-1).indices, True) + dispatcher = self.token_dispatcher + hidden, token_probs = dispatcher.dispatch_preprocess(value, routing, probs) + routed, routed_probs = dispatcher.token_dispatch(hidden, token_probs) + routed, _, routed_probs = dispatcher.dispatch_postprocess(routed, routed_probs) + transformed = routed.tanh() * routed_probs[:, None] + combined = dispatcher.token_combine(dispatcher.combine_preprocess(transformed)) + return value + 0.2 * dispatcher.combine_postprocess(combined) + + +def _run_checkpointed_flex_router(model): + model.zero_grad(set_to_none=True) + initial = torch.linspace(-1, 1, 88).reshape(1, 11, 8).requires_grad_() + inputs = [] + + def checkpointed(layer): + def compute(value): + if torch.is_grad_enabled(): + assert value.is_leaf + inputs.append(weakref.ref(value)) + return layer(value) + + return compute + + hidden = initial + for layer in model: + hidden = mcore_random.CheckpointFunction.apply( + checkpointed(layer), False, hidden + ) + loss = hidden.square().sum() + loss.backward() + gc.collect() + assert len(inputs) == len(model) + alive = [reference() is not None for reference in inputs] + gradients = [] + for value in [initial, *model.parameters()]: + assert value.grad is not None + gradients.append(value.grad.clone()) + for layer in model: + comm = cast(Any, layer).token_dispatcher._comm_manager + comm.routing_map = comm.token_probs = comm.dispatched_probs = None + gc.collect() + assert all(reference() is None for reference in inputs) + return loss.detach(), gradients, alive + + +def test_hybridep_state_releases_checkpoint_inputs(cpu_hybridep): + torch.manual_seed(954) + model = torch.nn.ModuleList([_FlexRouterLayer() for _ in range(4)]) + original = MoEFlexTokenDispatcher.combine_postprocess + reference_loss, reference_grads, retained = _run_checkpointed_flex_router(model) + # Upstream HybridEP keeps each layer's checkpoint graph after backward. + assert all(retained) + _configure_moe_dispatcher_caches([model]) + dispatchers = [cast(Any, layer).token_dispatcher for layer in model] + for dispatcher in dispatchers: + assert isinstance(dispatcher.combine_postprocess, partial) + assert dispatcher._comm_manager.token_probs is None + loss, gradients, retained = _run_checkpointed_flex_router(model) + assert not any(retained) + assert torch.equal(loss, reference_loss) + for actual, expected in zip(gradients, reference_grads, strict=True): + assert torch.equal(actual, expected) + assert MoEFlexTokenDispatcher.combine_postprocess is original + adapted = [dispatcher.combine_postprocess for dispatcher in dispatchers] + _configure_moe_dispatcher_caches([model]) + assert [dispatcher.combine_postprocess for dispatcher in dispatchers] == adapted + + +def test_hybridep_state_released_after_each_combine(cpu_hybridep): + layer = _FlexRouterLayer() + _configure_moe_dispatcher_caches([layer]) + comm = layer.token_dispatcher._comm_manager + value = torch.randn(1, 11, 8, requires_grad=True) + output = layer(value) + # Backward keeps what it needs through the graph, not the manager. + assert comm.routing_map is None + assert comm.token_probs is None + assert comm.dispatched_probs is None + assert comm.handle is None + output.square().sum().backward() + assert value.grad is not None and torch.isfinite(value.grad).all() + assert layer.weight.grad is not None and torch.isfinite(layer.weight.grad).all() + + +@pytest.mark.parametrize("case", ["deepep", "cuda_graph", "custom_combine"]) +def test_other_flex_dispatchers_keep_their_state(cpu_hybridep, case): + layer = _FlexRouterLayer("deepep" if case == "deepep" else "hybridep") + dispatcher = layer.token_dispatcher + if case == "cuda_graph": + # CUDA graph capture reads routing inputs back from the manager. + dispatcher.config.cuda_graph_impl = "transformer_engine" + combine = dispatcher.combine_postprocess + if case == "custom_combine": + dispatcher.combine_postprocess = combine + probs = dispatcher._comm_manager.token_probs = torch.ones(1) + _configure_moe_dispatcher_caches([layer]) + assert "combine_postprocess" not in vars(dispatcher) or ( + dispatcher.combine_postprocess is combine + ) + assert dispatcher._comm_manager.token_probs is probs + + +@pytest.mark.parametrize("round_trip", ["pickle", "deepcopy"]) +def test_hybridep_adaptation_survives_serialization(cpu_hybridep, round_trip): + torch.manual_seed(954) + model = torch.nn.ModuleList([_FlexRouterLayer() for _ in range(2)]) + expected_loss, expected_grads, retained = _run_checkpointed_flex_router(model) + assert all(retained) + _configure_moe_dispatcher_caches([model]) + clone = ( + pickle.loads(pickle.dumps(model)) if round_trip == "pickle" else deepcopy(model) + ) + for layer in clone: + dispatcher: Any = cast(Any, layer).token_dispatcher + combine = dispatcher.combine_postprocess + assert isinstance(combine, partial) and combine.args[0] is dispatcher + _configure_moe_dispatcher_caches([clone]) + assert dispatcher.combine_postprocess is combine + loss, gradients, retained = _run_checkpointed_flex_router(clone) + assert not any(retained) + assert torch.equal(loss, expected_loss) + for actual, expected in zip(gradients, expected_grads, strict=True): + assert torch.equal(actual, expected) From 0fe4e1553bacff4f99a186ad5f63800a76e1d3fb Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Fri, 25 Sep 2026 16:13:30 +0000 Subject: [PATCH 2/4] Cover compiled, pending-graph, repeated-backward and subclass HybridEP cases Co-Authored-By: Claude Opus 5.5 (1M context) --- .../test_dispatcher_graph_retention.py | 131 +++++++++++++++--- 1 file changed, 111 insertions(+), 20 deletions(-) diff --git a/tests/integration/megatron/model_support/test_dispatcher_graph_retention.py b/tests/integration/megatron/model_support/test_dispatcher_graph_retention.py index 427777bd5..08e6ff582 100644 --- a/tests/integration/megatron/model_support/test_dispatcher_graph_retention.py +++ b/tests/integration/megatron/model_support/test_dispatcher_graph_retention.py @@ -699,7 +699,7 @@ def forward(self, value): return value + 0.2 * dispatcher.combine_postprocess(combined) -def _run_checkpointed_flex_router(model): +def _run_checkpointed_flex_router(model, *, install_before_backward=False): model.zero_grad(set_to_none=True) initial = torch.linspace(-1, 1, 88).reshape(1, 11, 8).requires_grad_() inputs = [] @@ -719,6 +719,9 @@ def compute(value): checkpointed(layer), False, hidden ) loss = hidden.square().sum() + if install_before_backward: + # This forward ran unadapted; installation must release its state. + _configure_moe_dispatcher_caches([model]) loss.backward() gc.collect() assert len(inputs) == len(model) @@ -735,27 +738,101 @@ def compute(value): return loss.detach(), gradients, alive -def test_hybridep_state_releases_checkpoint_inputs(cpu_hybridep): +@pytest.mark.parametrize("pending_graph", [False, True]) +@pytest.mark.parametrize("compiled", [False, True]) +def test_hybridep_state_releases_checkpoint_inputs( + cpu_hybridep, compiled, pending_graph +): torch.manual_seed(954) model = torch.nn.ModuleList([_FlexRouterLayer() for _ in range(4)]) + backend = CompileCounterWithBackend("aot_eager") if compiled else None + if backend is not None: + model = torch.nn.ModuleList( + [ + cast(torch.nn.Module, torch.compile(layer, backend=backend)) + for layer in model + ] + ) original = MoEFlexTokenDispatcher.combine_postprocess - reference_loss, reference_grads, retained = _run_checkpointed_flex_router(model) - # Upstream HybridEP keeps each layer's checkpoint graph after backward. - assert all(retained) - _configure_moe_dispatcher_caches([model]) - dispatchers = [cast(Any, layer).token_dispatcher for layer in model] - for dispatcher in dispatchers: - assert isinstance(dispatcher.combine_postprocess, partial) - assert dispatcher._comm_manager.token_probs is None - loss, gradients, retained = _run_checkpointed_flex_router(model) - assert not any(retained) - assert torch.equal(loss, reference_loss) - for actual, expected in zip(gradients, reference_grads, strict=True): - assert torch.equal(actual, expected) - assert MoEFlexTokenDispatcher.combine_postprocess is original - adapted = [dispatcher.combine_postprocess for dispatcher in dispatchers] - _configure_moe_dispatcher_caches([model]) - assert [dispatcher.combine_postprocess for dispatcher in dispatchers] == adapted + try: + reference_loss, reference_grads, retained = _run_checkpointed_flex_router(model) + # Upstream HybridEP keeps each layer's checkpoint graph after backward. + assert all(retained) + # Install into the already-warmed (compiled) model without a reset. + if not pending_graph: + _configure_moe_dispatcher_caches([model]) + loss, gradients, retained = _run_checkpointed_flex_router( + model, install_before_backward=pending_graph + ) + assert not any(retained) + assert torch.equal(loss, reference_loss) + for actual, expected in zip(gradients, reference_grads, strict=True): + assert torch.equal(actual, expected) + dispatchers = [cast(Any, layer).token_dispatcher for layer in model] + for dispatcher in dispatchers: + assert isinstance(dispatcher.combine_postprocess, partial) + assert dispatcher._comm_manager.token_probs is None + assert MoEFlexTokenDispatcher.combine_postprocess is original + adapted = [dispatcher.combine_postprocess for dispatcher in dispatchers] + _configure_moe_dispatcher_caches([model]) + assert [d.combine_postprocess for d in dispatchers] == adapted + if backend is not None: + assert backend.frame_count > 0 + finally: + torch.compiler.reset() + + +@pytest.mark.parametrize("compiled", [False, True]) +@pytest.mark.parametrize("checkpointed", [False, True]) +def test_hybridep_state_allows_outstanding_forwards_and_repeated_backward( + cpu_hybridep, compiled, checkpointed +): + torch.manual_seed(848) + reference = torch.nn.Sequential(_FlexRouterLayer(), _FlexRouterLayer()) + adapted = deepcopy(reference) + _configure_moe_dispatcher_caches([adapted]) + if compiled: + reference = cast(torch.nn.Module, torch.compile(reference, backend="aot_eager")) + adapted = cast(torch.nn.Module, torch.compile(adapted, backend="aot_eager")) + + def run(model): + inputs = [ + torch.linspace(-1 + offset, 1 + offset, 88) + .reshape(1, 11, 8) + .requires_grad_() + for offset in (0, 0.3) + ] + outputs = [ + mcore_random.CheckpointFunction.apply(model, False, value) + if checkpointed + else model(value) + for value in inputs + ] + losses = [output.square().sum() for output in outputs] + losses[1].backward(retain_graph=True) + losses[0].backward() + losses[1].backward() + gradients = [] + for tensor in [*inputs, *model.parameters()]: + assert tensor.grad is not None + gradients.append(tensor.grad.clone()) + return [output.detach() for output in outputs], gradients + + try: + expected_outputs, expected_grads = run(reference) + outputs, grads = run(adapted) + for actual, expected in zip( + [*outputs, *grads], [*expected_outputs, *expected_grads], strict=True + ): + assert torch.equal(actual, expected) + for module in adapted.modules(): + if isinstance(module, _FlexRouterLayer): + comm = module.token_dispatcher._comm_manager + assert comm.routing_map is None + assert comm.token_probs is None + assert comm.dispatched_probs is None + finally: + torch.compiler.reset() def test_hybridep_state_released_after_each_combine(cpu_hybridep): @@ -774,10 +851,24 @@ def test_hybridep_state_released_after_each_combine(cpu_hybridep): assert layer.weight.grad is not None and torch.isfinite(layer.weight.grad).all() -@pytest.mark.parametrize("case", ["deepep", "cuda_graph", "custom_combine"]) +@pytest.mark.parametrize( + "case", + [ + "deepep", + "cuda_graph", + "custom_combine", + "dispatcher_subclass", + "manager_subclass", + ], +) def test_other_flex_dispatchers_keep_their_state(cpu_hybridep, case): layer = _FlexRouterLayer("deepep" if case == "deepep" else "hybridep") dispatcher = layer.token_dispatcher + if case == "dispatcher_subclass": + dispatcher.__class__ = type("CustomFlex", (MoEFlexTokenDispatcher,), {}) + if case == "manager_subclass": + manager_type = type(dispatcher._comm_manager) + dispatcher._comm_manager.__class__ = type("CustomManager", (manager_type,), {}) if case == "cuda_graph": # CUDA graph capture reads routing inputs back from the manager. dispatcher.config.cuda_graph_impl = "transformer_engine" From 9d6158f7abc2c7b496d4a056a4066ef068012d2f Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Fri, 25 Sep 2026 16:21:29 +0000 Subject: [PATCH 3/4] Test that installing the HybridEP adaptation drops state already held Co-Authored-By: Claude Opus 5.5 (1M context) --- .../test_dispatcher_graph_retention.py | 20 +++++++++++++++++++ 1 file changed, 20 insertions(+) diff --git a/tests/integration/megatron/model_support/test_dispatcher_graph_retention.py b/tests/integration/megatron/model_support/test_dispatcher_graph_retention.py index 08e6ff582..2d2c26e76 100644 --- a/tests/integration/megatron/model_support/test_dispatcher_graph_retention.py +++ b/tests/integration/megatron/model_support/test_dispatcher_graph_retention.py @@ -851,6 +851,26 @@ def test_hybridep_state_released_after_each_combine(cpu_hybridep): assert layer.weight.grad is not None and torch.isfinite(layer.weight.grad).all() +def test_hybridep_install_releases_existing_state(cpu_hybridep): + # A forward that ran before installation leaves its state on the manager; + # installation itself must drop it, not only later combines. + layer = _FlexRouterLayer() + value = torch.randn(1, 11, 8, requires_grad=True) + output = layer(value) + comm = layer.token_dispatcher._comm_manager + held = weakref.ref(comm.dispatched_probs) + assert comm.routing_map is not None and comm.token_probs is not None + _configure_moe_dispatcher_caches([layer]) + assert comm.routing_map is None + assert comm.token_probs is None + assert comm.dispatched_probs is None + output.square().sum().backward() + assert value.grad is not None and torch.isfinite(value.grad).all() + del output + gc.collect() + assert held() is None + + @pytest.mark.parametrize( "case", [ From 94d67933f474cb2e3375296fce9876abb5cade08 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Fri, 25 Sep 2026 16:35:12 +0000 Subject: [PATCH 4/4] Type the serialized HybridEP test clone as Any instead of casting each layer Co-Authored-By: Claude Opus 5.5 (1M context) --- .../megatron/model_support/test_dispatcher_graph_retention.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/tests/integration/megatron/model_support/test_dispatcher_graph_retention.py b/tests/integration/megatron/model_support/test_dispatcher_graph_retention.py index 2d2c26e76..8b81ef201 100644 --- a/tests/integration/megatron/model_support/test_dispatcher_graph_retention.py +++ b/tests/integration/megatron/model_support/test_dispatcher_graph_retention.py @@ -910,11 +910,11 @@ def test_hybridep_adaptation_survives_serialization(cpu_hybridep, round_trip): expected_loss, expected_grads, retained = _run_checkpointed_flex_router(model) assert all(retained) _configure_moe_dispatcher_caches([model]) - clone = ( + clone: Any = ( pickle.loads(pickle.dumps(model)) if round_trip == "pickle" else deepcopy(model) ) for layer in clone: - dispatcher: Any = cast(Any, layer).token_dispatcher + dispatcher = layer.token_dispatcher combine = dispatcher.combine_postprocess assert isinstance(combine, partial) and combine.args[0] is dispatcher _configure_moe_dispatcher_caches([clone])