Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
30 changes: 30 additions & 0 deletions src/art/trainer_rank/_impl.py
Original file line number Diff line number Diff line change
Expand Up @@ -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():
Expand All @@ -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)
):
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -625,3 +625,302 @@ 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, *, install_before_backward=False):
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()
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)
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


@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
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):
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()


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",
[
"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"
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: Any = (
pickle.loads(pickle.dumps(model)) if round_trip == "pickle" else deepcopy(model)
)
for layer in clone:
dispatcher = 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)
Loading