From eecb6a8a1c74b22ffe2f161a6a8cac9ff590969d Mon Sep 17 00:00:00 2001 From: 3outeille Date: Wed, 16 Sep 2026 11:05:41 +0000 Subject: [PATCH 01/86] Add dense and expert device mesh views with a MeshManager `initialize_distributed_mesh` now builds two named views of the same ranks: `(pp, fsdp, tp)` for dense layers and `(pp, efsdp, ep)` for experts, both keeping size-one axes so callers select dimensions by name. `MeshManager` routes `ep`/`efsdp` lookups to the expert view and everything else to the dense view. `DistributedConfig` gains `ep_size` (defaults to `tp_size` when `enable_expert_parallel=True`) and `efsdp_size`, with size validation. Model execution is unchanged: expert sharding and FSDP still use the `tp` and `fsdp` axes, and loading rejects `ep_size != tp_size` until the all-to-all dispatcher lands. --- docs/source/en/expert_parallelism.md | 16 +- .../distributed/configuration_utils.py | 27 +- src/transformers/distributed/mixin.py | 26 +- src/transformers/distributed/utils.py | 56 +++-- src/transformers/modeling_utils.py | 5 +- tests/test_distributed_config.py | 232 ++++++++++++++++++ tests/test_fsdp_mixin.py | 2 +- 7 files changed, 325 insertions(+), 39 deletions(-) create mode 100644 tests/test_distributed_config.py diff --git a/docs/source/en/expert_parallelism.md b/docs/source/en/expert_parallelism.md index 59daa55f1bff..d01c1ad94576 100644 --- a/docs/source/en/expert_parallelism.md +++ b/docs/source/en/expert_parallelism.md @@ -66,7 +66,7 @@ distributed_config = DistributedConfig( model = AutoModelForCausalLM.from_pretrained("Qwen/Qwen3-30B-A3B", distributed_config=distributed_config) ``` -The model is loaded on a 2D `(fsdp, tp)` device mesh, and `tp_size * fsdp_size` must equal the number of processes. The expert parallel plan shards the experts across `tp`, then FSDP2 shards every parameter, experts included, across `fsdp` and owns their gradient reduction. Each `fsdp` rank trains on its own part of the batch. +The model is loaded on a `(pp, fsdp, tp)` device mesh with `pp_size=1`, and `tp_size * fsdp_size` must equal the number of processes. The expert parallel plan shards the experts across `tp`, then FSDP2 shards every parameter, experts included, across `fsdp` and owns their gradient reduction. Each `fsdp` rank trains on its own part of the batch. Load the model as usual, then train with [`Trainer`]. It takes the gradient norm across both meshes and gives each mesh its own optimizer param group. [`~Trainer.save_model`] gathers sharded weights into a regular checkpoint. This requires `accelerate>=1.12` so the `Trainer` can mirror `tp_size` and `fsdp_size` into [`~Accelerate.ParallelismConfig`]. @@ -81,6 +81,20 @@ The table below compares EP-only training with 2D EP+FSDP2 on 8xH100 GPUs. The w > [!WARNING] > Resuming from a checkpoint is not supported yet for models sharded at load time, so the [`Trainer`] only accepts `save_only_model=True` or `save_strategy="no"` for them. +## Mesh views + +`DistributedConfig` also accepts an explicit `ep_size`. For the current all-reduce implementation, +set `ep_size=tp_size`; `DistributedConfig(tp_size=4, ep_size=4)` is equivalent to +`DistributedConfig(tp_size=4, enable_expert_parallel=True)`. An explicit `ep_size=1` disables EP. + +Internally, a mesh manager provides two views of the same ranks: `(pp, fsdp, tp)` for dense layers +and `(pp, efsdp, ep)` for experts, where `efsdp_size = fsdp_size * tp_size // ep_size`. +Both retain size-one axes, so callers can select dimensions by name. The mesh builder supports +`ep_size` values that are multiples of `tp_size` and divide `fsdp_size * tp_size`. +Model loading currently rejects enabled EP layouts with `ep_size != tp_size` because all-reduce +requires identical tokens within each expert group. Expert sharding and FSDP continue to use the +`tp` and `fsdp` axes of the dense view. + ## API reference [[autodoc]] DistributedConfig diff --git a/src/transformers/distributed/configuration_utils.py b/src/transformers/distributed/configuration_utils.py index 55e7c85013dc..814c8335b2fd 100644 --- a/src/transformers/distributed/configuration_utils.py +++ b/src/transformers/distributed/configuration_utils.py @@ -33,7 +33,8 @@ class DistributedConfig: enable_sequence_parallel (`bool`, *optional*, defaults to `False`): Reserved for sequence parallelism. Not wired up yet. enable_expert_parallel (`bool`, *optional*, defaults to `False`): - Route MoE models through the expert-parallel path (``base_model_ep_plan``). + Route MoE models through the expert-parallel path (``base_model_ep_plan``). When `ep_size` is + omitted, sets it to `tp_size`. An explicit `ep_size` takes precedence. fsdp_size (`int`, *optional*): Number of devices for FSDP (data parallelism). If `None` and `tp_size` is set, defaults to 1. fsdp_cpu_offload (`bool`, *optional*, defaults to `False`): @@ -42,6 +43,9 @@ class DistributedConfig: Whether to enable mixed precision for FSDP2. pp_size (`int`, *optional*): Number of devices for pipeline parallelism. If `None` and another parallel mode is set, defaults to 1. + ep_size (`int`, *optional*): + Number of devices owning distinct expert shards. Defaults to 1, or to `tp_size` when + `enable_expert_parallel=True`. Model execution currently requires `ep_size=tp_size` when EP is enabled. """ tp_size: int | None = None @@ -52,10 +56,17 @@ class DistributedConfig: fsdp_cpu_offload: bool = False fsdp_mixed_precision: bool = False pp_size: int | None = None + ep_size: int | None = None + + @property + def efsdp_size(self) -> int: + """Size of the expert FSDP axis in the expert mesh view.""" + return self.fsdp_size * self.tp_size // self.ep_size def __post_init__(self): - if self.tp_plan is None and self.tp_size is None and self.fsdp_size is None and self.pp_size is None: - return + for value in (self.tp_size, self.fsdp_size, self.pp_size, self.ep_size): + if value is not None and value < 1: + raise ValueError(f"Parallelism sizes must be >= 1, got {value}.") if self.fsdp_size is None: self.fsdp_size = 1 @@ -73,6 +84,16 @@ def __post_init__(self): elif self.tp_size is None: self.tp_size = 1 + if self.ep_size is None: + self.ep_size = self.tp_size if self.enable_expert_parallel else 1 + self.enable_expert_parallel = self.ep_size > 1 + + if self.ep_size > 1: + if self.ep_size % self.tp_size: + raise ValueError("`ep_size` must be a multiple of `tp_size`.") + if (self.fsdp_size * self.tp_size) % self.ep_size: + raise ValueError("`ep_size` must divide `fsdp_size * tp_size`.") + if self.fsdp_size > 1 and self.pp_size > 1: raise ValueError( "Combining FSDP with pipeline parallelism is not supported yet. " diff --git a/src/transformers/distributed/mixin.py b/src/transformers/distributed/mixin.py index 646b8ac102cc..6f4e96b04d28 100644 --- a/src/transformers/distributed/mixin.py +++ b/src/transformers/distributed/mixin.py @@ -29,6 +29,7 @@ gather_state_dict_for_save, ) from .utils import ( + MeshManager, _distributed_barrier, _get_torch_distributed_rank, _is_torch_distributed_initialized, @@ -49,6 +50,7 @@ class DistributedMixin: """Distributed orchestration and save/load hooks for [`PreTrainedModel`].""" _device_mesh = None + _mesh_manager: MeshManager | None = None _tp_plan: dict[str, str] | None = None _ep_plan: dict[str, str] | None = None _tp_size = None @@ -142,13 +144,18 @@ def prepare_distribute_model( cls, distributed_config: DistributedConfig | dict | None, device_map=None, - ) -> tuple[DistributedConfig | None, object, object]: + ) -> tuple[DistributedConfig | None, object, MeshManager | None]: if distributed_config is None: return None, device_map, None if isinstance(distributed_config, dict): distributed_config = DistributedConfig.from_dict(distributed_config) + if distributed_config.ep_size > 1 and distributed_config.ep_size != distributed_config.tp_size: + raise ValueError( + "All-reduce expert parallelism requires `ep_size=tp_size` and identical tokens per EP group." + ) + if distributed_config.tp_size == 1 and distributed_config.fsdp_size == 1 and distributed_config.pp_size == 1: return distributed_config, device_map, None @@ -157,38 +164,39 @@ def prepare_distribute_model( if distributed_config.fsdp_size > 1 and not is_torch_greater_or_equal("2.7"): raise OSError("FSDP2 requires `torch>=2.7` (distributed checkpoint save/load).") - device_map, device_mesh = initialize_distributed_mesh(distributed_config) + device_map, mesh_manager = initialize_distributed_mesh(distributed_config) - return distributed_config, device_map, device_mesh + return distributed_config, device_map, mesh_manager @classmethod def maybe_distribute_model( cls, model: nn.Module, distributed_config: DistributedConfig | None, - device_mesh, + mesh_manager: MeshManager | None, ): """Apply TP or FSDP2 after model init, before weight loading.""" - if device_mesh is not None: + if mesh_manager is not None: model.config.distributed_config = distributed_config - model._device_mesh = device_mesh + model._mesh_manager = mesh_manager + model._device_mesh = mesh_manager.get_mesh(("pp", "fsdp", "tp")) model._tp_size = distributed_config.tp_size model._fsdp_size = distributed_config.fsdp_size if distributed_config.pp_size > 1: - pp_mesh = device_mesh["pp"] if device_mesh.ndim > 1 else device_mesh + pp_mesh = mesh_manager.get_mesh("pp") model = apply_pipeline_parallelism(model, pp_mesh) # Both may apply: the tensor/expert parallel plan shards across `tp` first, then FSDP2 # shards every parameter (the `tp`-sharded ones included) across `fsdp`. if distributed_config.tp_size > 1: - tp_mesh = device_mesh["tp"] if device_mesh.ndim > 1 else device_mesh + tp_mesh = mesh_manager.get_mesh("tp") if isinstance(distributed_config.tp_plan, dict): model.tp_plan = distributed_config.tp_plan model = apply_tensor_parallelism(model, tp_mesh) if distributed_config.fsdp_size > 1: - fsdp_mesh = device_mesh["fsdp"] if device_mesh.ndim > 1 else device_mesh + fsdp_mesh = mesh_manager.get_mesh("fsdp") model = apply_fully_sharded_data_parallelism(model, fsdp_mesh) return model diff --git a/src/transformers/distributed/utils.py b/src/transformers/distributed/utils.py index 571e4239f087..c843a9139f24 100644 --- a/src/transformers/distributed/utils.py +++ b/src/transformers/distributed/utils.py @@ -25,6 +25,7 @@ if TYPE_CHECKING: + from torch.distributed.device_mesh import DeviceMesh from torch.distributed.tensor import DTensor from .configuration_utils import DistributedConfig @@ -130,6 +131,20 @@ def _distributed_barrier(): torch.distributed.barrier() +class MeshManager: + """Named access to dense and expert parallel axes without exposing their view selection.""" + + def __init__(self, dense_mesh: DeviceMesh, expert_mesh: DeviceMesh): + self._dense_mesh = dense_mesh + self._expert_mesh = expert_mesh + + def get_mesh(self, dims: str | tuple[str, ...]) -> DeviceMesh: + """Select expert axes for `ep`/`efsdp`, otherwise dense axes; DeviceMesh handles slicing.""" + dims = (dims,) if isinstance(dims, str) else dims + mesh = self._expert_mesh if "ep" in dims or "efsdp" in dims else self._dense_mesh + return mesh[dims] + + # Retained for the legacy transformers.integrations.tensor_parallel API. def initialize_tensor_parallelism( tp_plan: str | dict[str, str] | None, tp_size: int | None = None, device_mesh=None, device_map=None @@ -223,22 +238,15 @@ def initialize_fully_sharded_data_parallelism(distributed_config: DistributedCon def initialize_distributed_mesh( distributed_config: DistributedConfig, -): - """Create a device mesh containing every configured parallel dimension.""" - mesh_shape = [] - mesh_dim_names = [] - - if distributed_config.pp_size > 1: - mesh_shape.append(distributed_config.pp_size) - mesh_dim_names.append("pp") - if distributed_config.fsdp_size > 1: - mesh_shape.append(distributed_config.fsdp_size) - mesh_dim_names.append("fsdp") - if distributed_config.tp_size > 1: - mesh_shape.append(distributed_config.tp_size) - mesh_dim_names.append("tp") - - if not mesh_shape: +) -> tuple[torch.device | None, MeshManager | None]: + """Build named dense and expert views, independently of the expert dispatcher. + + Both views include singleton dimensions so callers can always select their axes by name. + Each parameter's FSDP and TP/EP axes come from the same view. Separate roots avoid requiring + the newer `DeviceMesh._unflatten` API; the expert view is unused when EP is disabled. + """ + mesh_shape = (distributed_config.pp_size, distributed_config.fsdp_size, distributed_config.tp_size) + if mesh_shape == (1, 1, 1): return None, None device_type = torch._C._get_accelerator().type @@ -260,15 +268,17 @@ def initialize_distributed_mesh( else: device_map = torch.device(device_type) - device_mesh = torch.distributed.init_device_mesh( + dense_mesh = torch.distributed.init_device_mesh( device_type, - tuple(mesh_shape), - mesh_dim_names=tuple(mesh_dim_names), + mesh_shape, + mesh_dim_names=("pp", "fsdp", "tp"), ) - # A flattened sub-mesh, so an all-reduce over every rank is one collective instead of one per dimension. - if len(mesh_dim_names) > 1: - device_mesh._flatten("_".join(mesh_dim_names)) - return device_map, device_mesh + expert_mesh = torch.distributed.init_device_mesh( + device_type, + (distributed_config.pp_size, distributed_config.efsdp_size, distributed_config.ep_size), + mesh_dim_names=("pp", "efsdp", "ep"), + ) + return device_map, MeshManager(dense_mesh, expert_mesh) def gather_full_state_dict(model) -> dict[str, torch.Tensor]: diff --git a/src/transformers/modeling_utils.py b/src/transformers/modeling_utils.py index d8dd2981e274..07291b2870f6 100644 --- a/src/transformers/modeling_utils.py +++ b/src/transformers/modeling_utils.py @@ -4160,9 +4160,10 @@ def from_pretrained( distributed_config = DistributedConfig(tp_plan=tp_plan, tp_size=tp_size) if distributed_config is not None: - distributed_config, device_map, device_mesh = cls.prepare_distribute_model( + distributed_config, device_map, mesh_manager = cls.prepare_distribute_model( distributed_config, device_map=device_map ) + device_mesh = mesh_manager.get_mesh(("pp", "fsdp", "tp")) if mesh_manager is not None else None if gguf_file is not None and not is_accelerate_available(): raise ValueError("accelerate is required when loading a GGUF file `pip install accelerate`.") @@ -4309,7 +4310,7 @@ def from_pretrained( weight_conversions = get_model_conversion_mapping(model, key_mapping, hf_quantizer) if distributed_config is not None: - model = cls.maybe_distribute_model(model, distributed_config, device_mesh) + model = cls.maybe_distribute_model(model, distributed_config, mesh_manager) # Prepare the full device map if device_map is not None: diff --git a/tests/test_distributed_config.py b/tests/test_distributed_config.py new file mode 100644 index 000000000000..ba9f9c507f66 --- /dev/null +++ b/tests/test_distributed_config.py @@ -0,0 +1,232 @@ +# Copyright 2026 The HuggingFace Team. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +import os +import tempfile +import unittest +from datetime import timedelta +from unittest.mock import patch + +from transformers.distributed import DistributedConfig +from transformers.testing_utils import require_torch, require_torch_greater_or_equal +from transformers.utils import is_torch_available + + +if is_torch_available(): + import torch + import torch.distributed as dist + import torch.multiprocessing as mp + + from transformers.distributed.mixin import DistributedMixin + from transformers.distributed.utils import initialize_distributed_mesh + + +class DistributedConfigTest(unittest.TestCase): + def test_defaults_and_round_trip(self): + for kwargs in ({}, {"tp_size": 4}, {"fsdp_size": 4}, {"pp_size": 4}, {"tp_size": 2, "fsdp_size": 2}): + with self.subTest(kwargs=kwargs): + config = DistributedConfig(**kwargs) + self.assertEqual(config.ep_size, 1) + self.assertFalse(config.enable_expert_parallel) + self.assertEqual(config.efsdp_size, config.fsdp_size * config.tp_size) + self.assertEqual(DistributedConfig.from_dict(config.to_dict()), config) + + def test_legacy_and_explicit_ep_sizes(self): + legacy = DistributedConfig(tp_size=4, fsdp_size=2, enable_expert_parallel=True) + explicit = DistributedConfig(tp_size=4, fsdp_size=2, ep_size=4) + self.assertEqual(legacy, explicit) + self.assertEqual(DistributedConfig.from_dict(explicit.to_dict()), explicit) + for ep_size in (1, 4, 8): + with self.subTest(ep_size=ep_size): + config = DistributedConfig(tp_size=4, fsdp_size=2, ep_size=ep_size, enable_expert_parallel=True) + self.assertEqual(config.ep_size, ep_size) + self.assertEqual(config.enable_expert_parallel, ep_size > 1) + self.assertEqual((config.tp_size, config.fsdp_size), (4, 2)) + + def test_inferred_tp_size(self): + with patch.dict(os.environ, {"WORLD_SIZE": "8"}): + config = DistributedConfig(tp_plan="auto", fsdp_size=2, enable_expert_parallel=True) + self.assertEqual((config.tp_size, config.ep_size, config.efsdp_size), (4, 4, 2)) + + def test_expert_mesh_sizes(self): + for fsdp, tp, ep, efsdp in ((8, 1, 4, 2), (2, 2, 4, 1), (4, 2, 4, 2), (2, 2, 2, 2), (1, 4, 4, 1)): + with self.subTest(fsdp=fsdp, tp=tp, ep=ep): + config = DistributedConfig(fsdp_size=fsdp, tp_size=tp, ep_size=ep) + self.assertEqual(config.efsdp_size, efsdp) + self.assertEqual(DistributedConfig.from_dict(config.to_dict()), config) + + def test_invalid_sizes(self): + for name in ("tp_size", "fsdp_size", "pp_size", "ep_size"): + for value in (0, -1): + with self.subTest(name=name, value=value), self.assertRaisesRegex(ValueError, "must be >= 1"): + DistributedConfig(**{name: value}) + for kwargs, message in ( + ({"tp_size": 4, "ep_size": 2}, "multiple"), + ({"fsdp_size": 4, "ep_size": 3}, "must divide"), + ({"fsdp_size": 2, "pp_size": 2}, "pipeline parallelism"), + ({"ep_size": 2}, "must divide"), + ): + with self.subTest(kwargs=kwargs), self.assertRaisesRegex(ValueError, message): + DistributedConfig(**kwargs) + + +@require_torch +class DistributedMeshValidationTest(unittest.TestCase): + def test_disabled_mesh_does_not_initialize_distributed(self): + with patch("transformers.distributed.utils._ensure_torch_distributed") as initialize: + self.assertEqual(initialize_distributed_mesh(DistributedConfig()), (None, None)) + config, device_map, meshes = DistributedMixin.prepare_distribute_model({}, device_map="cpu") + self.assertEqual(config, DistributedConfig()) + self.assertEqual(device_map, "cpu") + self.assertIsNone(meshes) + initialize.assert_not_called() + + def test_model_loading_rejects_unsupported_ep_layout_before_initialization(self): + with patch("transformers.distributed.mixin.initialize_distributed_mesh") as initialize: + with self.assertRaisesRegex(ValueError, "ep_size=tp_size"): + DistributedMixin.prepare_distribute_model(DistributedConfig(fsdp_size=4, ep_size=2)) + initialize.assert_not_called() + + def test_world_size_mismatch(self): + with ( + patch("transformers.distributed.utils._ensure_torch_distributed"), + patch("torch._C._get_accelerator", return_value=torch.device("cpu")), + patch("torch.distributed.get_world_size", return_value=2), + self.assertRaisesRegex(RuntimeError, "requires 4 processes"), + ): + initialize_distributed_mesh(DistributedConfig(tp_size=4)) + + +def _mesh_worker(rank, rendezvous): + world_size = 4 + dist.init_process_group( + "gloo", + init_method=f"file://{rendezvous}", + rank=rank, + world_size=world_size, + timeout=timedelta(seconds=120), + ) + os.environ["LOCAL_RANK"] = str(rank) + try: + configs = [ + DistributedConfig(fsdp_size=4), + DistributedConfig(tp_size=4), + DistributedConfig(pp_size=4), + DistributedConfig(tp_size=2, fsdp_size=2), + DistributedConfig(tp_size=2, pp_size=2), + ] + configs += [DistributedConfig(fsdp_size=4, ep_size=ep) for ep in (2, 4)] + configs += [DistributedConfig(fsdp_size=2, tp_size=2, ep_size=ep) for ep in (2, 4)] + configs += [DistributedConfig(fsdp_size=1, tp_size=4, ep_size=4)] + for config in configs: + with patch("torch._C._get_accelerator", return_value=torch.device("cpu")): + _, meshes = initialize_distributed_mesh(config) + assert meshes.get_mesh(("pp", "fsdp", "tp")).mesh_dim_names == ("pp", "fsdp", "tp") + for axes in (("pp", "fsdp", "tp"), ("pp", "efsdp", "ep")): + assert meshes.get_mesh(axes).size() == world_size + assert meshes.get_mesh(axes).mesh_dim_names == axes + for name in axes: + assert meshes.get_mesh(name).size() == getattr(config, name + "_size") + assert meshes.get_mesh(("fsdp", "tp")).mesh_dim_names == ("fsdp", "tp") + assert meshes.get_mesh(("efsdp", "ep")).mesh_dim_names == ("efsdp", "ep") + for invalid in ("missing", ("tp", "ep"), ("fsdp", "efsdp")): + try: + meshes.get_mesh(invalid) + except KeyError: + pass + else: + raise AssertionError(f"Accepted invalid mesh dimensions: {invalid}") + stage_size = config.fsdp_size * config.tp_size + stage_start = rank // stage_size * stage_size + expert_rank = (rank - stage_start) % config.ep_size + ep_start = rank // config.ep_size * config.ep_size + assert dist.get_process_group_ranks(meshes.get_mesh("ep").get_group()) == list( + range(ep_start, ep_start + config.ep_size) + ) + assert dist.get_process_group_ranks(meshes.get_mesh("efsdp").get_group()) == list( + range(stage_start + expert_rank, stage_start + stage_size, config.ep_size) + ) + tp_start = rank // config.tp_size * config.tp_size + assert dist.get_process_group_ranks(meshes.get_mesh("tp").get_group()) == list( + range(tp_start, tp_start + config.tp_size) + ) + assert dist.get_process_group_ranks(meshes.get_mesh("fsdp").get_group()) == list( + range(stage_start + rank % config.tp_size, stage_start + stage_size, config.tp_size) + ) + assert dist.get_process_group_ranks(meshes.get_mesh("pp").get_group()) == list( + range(rank % stage_size, world_size, stage_size) + ) + assert meshes.get_mesh("ep").get_group() is meshes.get_mesh("ep").get_group() + finally: + dist.destroy_process_group() + + +def _dense_load_worker(rank, rendezvous): + from transformers import Qwen2Config, Qwen2ForCausalLM + + os.environ.update(RANK=str(rank), LOCAL_RANK=str(rank), WORLD_SIZE="2", LOCAL_WORLD_SIZE="2") + dist.init_process_group( + "gloo", init_method=f"file://{rendezvous}", rank=rank, world_size=2, timeout=timedelta(seconds=120) + ) + try: + torch.manual_seed(42) + config = Qwen2Config( + vocab_size=32, + hidden_size=8, + intermediate_size=8, + num_hidden_layers=2, + num_attention_heads=2, + num_key_value_heads=2, + ) + reference = Qwen2ForCausalLM(config).eval() + source = rendezvous + "_model" + if rank == 0: + reference.save_pretrained(source) + dist.barrier() + inputs = torch.tensor([[1, 2, 3]]) + generation_kwargs = { + "max_new_tokens": 2, + "do_sample": False, + "output_logits": True, + "return_dict_in_generate": True, + } + expected = reference.generate(inputs, **generation_kwargs) + for distributed_config in (DistributedConfig(tp_size=2), DistributedConfig(pp_size=2)): + with patch("torch._C._get_accelerator", return_value=torch.device("cpu")): + model = Qwen2ForCausalLM.from_pretrained(source, distributed_config=distributed_config).eval() + assert model._device_mesh is model._mesh_manager.get_mesh(("pp", "fsdp", "tp")) + actual = model.generate(inputs, **generation_kwargs) + torch.testing.assert_close(actual.sequences, expected.sequences) + torch.testing.assert_close(torch.stack(actual.logits), torch.stack(expected.logits)) + if distributed_config.tp_size > 1: + destination = rendezvous + "_saved" + model.save_pretrained(destination) + dist.barrier() + restored = Qwen2ForCausalLM.from_pretrained(destination).eval() + for name, param in restored.named_parameters(): + torch.testing.assert_close(param, dict(reference.named_parameters())[name], atol=0, rtol=0) + finally: + dist.destroy_process_group() + + +@require_torch +@require_torch_greater_or_equal("2.5") +class DistributedMeshTest(unittest.TestCase): + def test_mesh_groups(self): + with tempfile.TemporaryDirectory() as directory: + mp.spawn(_mesh_worker, args=(os.path.join(directory, "init"),), nprocs=4, join=True) + + def test_dense_load_generate_and_save(self): + with tempfile.TemporaryDirectory() as directory: + mp.spawn(_dense_load_worker, args=(os.path.join(directory, "init"),), nprocs=2, join=True) diff --git a/tests/test_fsdp_mixin.py b/tests/test_fsdp_mixin.py index 18fbcac99ee8..b59108cee436 100644 --- a/tests/test_fsdp_mixin.py +++ b/tests/test_fsdp_mixin.py @@ -539,7 +539,7 @@ def _test_fsdp2_expert_parallel_2d_vs_ddp_impl(rank, config_class, config_dict, distributed_config=DistributedConfig(tp_size=2, fsdp_size=dp, enable_expert_parallel=True), ) assert model.tp_size == 2 and model.fsdp_size == dp - assert model._device_mesh.mesh_dim_names == ("fsdp", "tp") + assert model._device_mesh.mesh_dim_names == ("pp", "fsdp", "tp") model.train() optimizer = torch.optim.Adam(model.parameters(), lr=LR, foreach=False) dp_rank = model._device_mesh["fsdp"].get_local_rank() From 29f512678e01adf44fc06fc44c1057ad94022b26 Mon Sep 17 00:00:00 2001 From: 3outeille Date: Wed, 16 Sep 2026 11:05:42 +0000 Subject: [PATCH 02/86] cleaning --- tests/test_distributed_config.py | 232 ------------------------------- 1 file changed, 232 deletions(-) delete mode 100644 tests/test_distributed_config.py diff --git a/tests/test_distributed_config.py b/tests/test_distributed_config.py deleted file mode 100644 index ba9f9c507f66..000000000000 --- a/tests/test_distributed_config.py +++ /dev/null @@ -1,232 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import os -import tempfile -import unittest -from datetime import timedelta -from unittest.mock import patch - -from transformers.distributed import DistributedConfig -from transformers.testing_utils import require_torch, require_torch_greater_or_equal -from transformers.utils import is_torch_available - - -if is_torch_available(): - import torch - import torch.distributed as dist - import torch.multiprocessing as mp - - from transformers.distributed.mixin import DistributedMixin - from transformers.distributed.utils import initialize_distributed_mesh - - -class DistributedConfigTest(unittest.TestCase): - def test_defaults_and_round_trip(self): - for kwargs in ({}, {"tp_size": 4}, {"fsdp_size": 4}, {"pp_size": 4}, {"tp_size": 2, "fsdp_size": 2}): - with self.subTest(kwargs=kwargs): - config = DistributedConfig(**kwargs) - self.assertEqual(config.ep_size, 1) - self.assertFalse(config.enable_expert_parallel) - self.assertEqual(config.efsdp_size, config.fsdp_size * config.tp_size) - self.assertEqual(DistributedConfig.from_dict(config.to_dict()), config) - - def test_legacy_and_explicit_ep_sizes(self): - legacy = DistributedConfig(tp_size=4, fsdp_size=2, enable_expert_parallel=True) - explicit = DistributedConfig(tp_size=4, fsdp_size=2, ep_size=4) - self.assertEqual(legacy, explicit) - self.assertEqual(DistributedConfig.from_dict(explicit.to_dict()), explicit) - for ep_size in (1, 4, 8): - with self.subTest(ep_size=ep_size): - config = DistributedConfig(tp_size=4, fsdp_size=2, ep_size=ep_size, enable_expert_parallel=True) - self.assertEqual(config.ep_size, ep_size) - self.assertEqual(config.enable_expert_parallel, ep_size > 1) - self.assertEqual((config.tp_size, config.fsdp_size), (4, 2)) - - def test_inferred_tp_size(self): - with patch.dict(os.environ, {"WORLD_SIZE": "8"}): - config = DistributedConfig(tp_plan="auto", fsdp_size=2, enable_expert_parallel=True) - self.assertEqual((config.tp_size, config.ep_size, config.efsdp_size), (4, 4, 2)) - - def test_expert_mesh_sizes(self): - for fsdp, tp, ep, efsdp in ((8, 1, 4, 2), (2, 2, 4, 1), (4, 2, 4, 2), (2, 2, 2, 2), (1, 4, 4, 1)): - with self.subTest(fsdp=fsdp, tp=tp, ep=ep): - config = DistributedConfig(fsdp_size=fsdp, tp_size=tp, ep_size=ep) - self.assertEqual(config.efsdp_size, efsdp) - self.assertEqual(DistributedConfig.from_dict(config.to_dict()), config) - - def test_invalid_sizes(self): - for name in ("tp_size", "fsdp_size", "pp_size", "ep_size"): - for value in (0, -1): - with self.subTest(name=name, value=value), self.assertRaisesRegex(ValueError, "must be >= 1"): - DistributedConfig(**{name: value}) - for kwargs, message in ( - ({"tp_size": 4, "ep_size": 2}, "multiple"), - ({"fsdp_size": 4, "ep_size": 3}, "must divide"), - ({"fsdp_size": 2, "pp_size": 2}, "pipeline parallelism"), - ({"ep_size": 2}, "must divide"), - ): - with self.subTest(kwargs=kwargs), self.assertRaisesRegex(ValueError, message): - DistributedConfig(**kwargs) - - -@require_torch -class DistributedMeshValidationTest(unittest.TestCase): - def test_disabled_mesh_does_not_initialize_distributed(self): - with patch("transformers.distributed.utils._ensure_torch_distributed") as initialize: - self.assertEqual(initialize_distributed_mesh(DistributedConfig()), (None, None)) - config, device_map, meshes = DistributedMixin.prepare_distribute_model({}, device_map="cpu") - self.assertEqual(config, DistributedConfig()) - self.assertEqual(device_map, "cpu") - self.assertIsNone(meshes) - initialize.assert_not_called() - - def test_model_loading_rejects_unsupported_ep_layout_before_initialization(self): - with patch("transformers.distributed.mixin.initialize_distributed_mesh") as initialize: - with self.assertRaisesRegex(ValueError, "ep_size=tp_size"): - DistributedMixin.prepare_distribute_model(DistributedConfig(fsdp_size=4, ep_size=2)) - initialize.assert_not_called() - - def test_world_size_mismatch(self): - with ( - patch("transformers.distributed.utils._ensure_torch_distributed"), - patch("torch._C._get_accelerator", return_value=torch.device("cpu")), - patch("torch.distributed.get_world_size", return_value=2), - self.assertRaisesRegex(RuntimeError, "requires 4 processes"), - ): - initialize_distributed_mesh(DistributedConfig(tp_size=4)) - - -def _mesh_worker(rank, rendezvous): - world_size = 4 - dist.init_process_group( - "gloo", - init_method=f"file://{rendezvous}", - rank=rank, - world_size=world_size, - timeout=timedelta(seconds=120), - ) - os.environ["LOCAL_RANK"] = str(rank) - try: - configs = [ - DistributedConfig(fsdp_size=4), - DistributedConfig(tp_size=4), - DistributedConfig(pp_size=4), - DistributedConfig(tp_size=2, fsdp_size=2), - DistributedConfig(tp_size=2, pp_size=2), - ] - configs += [DistributedConfig(fsdp_size=4, ep_size=ep) for ep in (2, 4)] - configs += [DistributedConfig(fsdp_size=2, tp_size=2, ep_size=ep) for ep in (2, 4)] - configs += [DistributedConfig(fsdp_size=1, tp_size=4, ep_size=4)] - for config in configs: - with patch("torch._C._get_accelerator", return_value=torch.device("cpu")): - _, meshes = initialize_distributed_mesh(config) - assert meshes.get_mesh(("pp", "fsdp", "tp")).mesh_dim_names == ("pp", "fsdp", "tp") - for axes in (("pp", "fsdp", "tp"), ("pp", "efsdp", "ep")): - assert meshes.get_mesh(axes).size() == world_size - assert meshes.get_mesh(axes).mesh_dim_names == axes - for name in axes: - assert meshes.get_mesh(name).size() == getattr(config, name + "_size") - assert meshes.get_mesh(("fsdp", "tp")).mesh_dim_names == ("fsdp", "tp") - assert meshes.get_mesh(("efsdp", "ep")).mesh_dim_names == ("efsdp", "ep") - for invalid in ("missing", ("tp", "ep"), ("fsdp", "efsdp")): - try: - meshes.get_mesh(invalid) - except KeyError: - pass - else: - raise AssertionError(f"Accepted invalid mesh dimensions: {invalid}") - stage_size = config.fsdp_size * config.tp_size - stage_start = rank // stage_size * stage_size - expert_rank = (rank - stage_start) % config.ep_size - ep_start = rank // config.ep_size * config.ep_size - assert dist.get_process_group_ranks(meshes.get_mesh("ep").get_group()) == list( - range(ep_start, ep_start + config.ep_size) - ) - assert dist.get_process_group_ranks(meshes.get_mesh("efsdp").get_group()) == list( - range(stage_start + expert_rank, stage_start + stage_size, config.ep_size) - ) - tp_start = rank // config.tp_size * config.tp_size - assert dist.get_process_group_ranks(meshes.get_mesh("tp").get_group()) == list( - range(tp_start, tp_start + config.tp_size) - ) - assert dist.get_process_group_ranks(meshes.get_mesh("fsdp").get_group()) == list( - range(stage_start + rank % config.tp_size, stage_start + stage_size, config.tp_size) - ) - assert dist.get_process_group_ranks(meshes.get_mesh("pp").get_group()) == list( - range(rank % stage_size, world_size, stage_size) - ) - assert meshes.get_mesh("ep").get_group() is meshes.get_mesh("ep").get_group() - finally: - dist.destroy_process_group() - - -def _dense_load_worker(rank, rendezvous): - from transformers import Qwen2Config, Qwen2ForCausalLM - - os.environ.update(RANK=str(rank), LOCAL_RANK=str(rank), WORLD_SIZE="2", LOCAL_WORLD_SIZE="2") - dist.init_process_group( - "gloo", init_method=f"file://{rendezvous}", rank=rank, world_size=2, timeout=timedelta(seconds=120) - ) - try: - torch.manual_seed(42) - config = Qwen2Config( - vocab_size=32, - hidden_size=8, - intermediate_size=8, - num_hidden_layers=2, - num_attention_heads=2, - num_key_value_heads=2, - ) - reference = Qwen2ForCausalLM(config).eval() - source = rendezvous + "_model" - if rank == 0: - reference.save_pretrained(source) - dist.barrier() - inputs = torch.tensor([[1, 2, 3]]) - generation_kwargs = { - "max_new_tokens": 2, - "do_sample": False, - "output_logits": True, - "return_dict_in_generate": True, - } - expected = reference.generate(inputs, **generation_kwargs) - for distributed_config in (DistributedConfig(tp_size=2), DistributedConfig(pp_size=2)): - with patch("torch._C._get_accelerator", return_value=torch.device("cpu")): - model = Qwen2ForCausalLM.from_pretrained(source, distributed_config=distributed_config).eval() - assert model._device_mesh is model._mesh_manager.get_mesh(("pp", "fsdp", "tp")) - actual = model.generate(inputs, **generation_kwargs) - torch.testing.assert_close(actual.sequences, expected.sequences) - torch.testing.assert_close(torch.stack(actual.logits), torch.stack(expected.logits)) - if distributed_config.tp_size > 1: - destination = rendezvous + "_saved" - model.save_pretrained(destination) - dist.barrier() - restored = Qwen2ForCausalLM.from_pretrained(destination).eval() - for name, param in restored.named_parameters(): - torch.testing.assert_close(param, dict(reference.named_parameters())[name], atol=0, rtol=0) - finally: - dist.destroy_process_group() - - -@require_torch -@require_torch_greater_or_equal("2.5") -class DistributedMeshTest(unittest.TestCase): - def test_mesh_groups(self): - with tempfile.TemporaryDirectory() as directory: - mp.spawn(_mesh_worker, args=(os.path.join(directory, "init"),), nprocs=4, join=True) - - def test_dense_load_generate_and_save(self): - with tempfile.TemporaryDirectory() as directory: - mp.spawn(_dense_load_worker, args=(os.path.join(directory, "init"),), nprocs=2, join=True) From 407452665db48fd131c21c8516b78cb36b906f8e Mon Sep 17 00:00:00 2001 From: 3outeille Date: Wed, 16 Sep 2026 11:05:42 +0000 Subject: [PATCH 03/86] clean --- src/transformers/distributed/utils.py | 7 +------ 1 file changed, 1 insertion(+), 6 deletions(-) diff --git a/src/transformers/distributed/utils.py b/src/transformers/distributed/utils.py index c843a9139f24..c06182d86b2e 100644 --- a/src/transformers/distributed/utils.py +++ b/src/transformers/distributed/utils.py @@ -239,12 +239,7 @@ def initialize_fully_sharded_data_parallelism(distributed_config: DistributedCon def initialize_distributed_mesh( distributed_config: DistributedConfig, ) -> tuple[torch.device | None, MeshManager | None]: - """Build named dense and expert views, independently of the expert dispatcher. - - Both views include singleton dimensions so callers can always select their axes by name. - Each parameter's FSDP and TP/EP axes come from the same view. Separate roots avoid requiring - the newer `DeviceMesh._unflatten` API; the expert view is unused when EP is disabled. - """ + """Create a device mesh containing every configured parallel dimension.""" mesh_shape = (distributed_config.pp_size, distributed_config.fsdp_size, distributed_config.tp_size) if mesh_shape == (1, 1, 1): return None, None From 578af0e7d5401037f89f081435413baa0e345981 Mon Sep 17 00:00:00 2001 From: 3outeille Date: Wed, 16 Sep 2026 12:15:08 +0000 Subject: [PATCH 04/86] update --- docs/source/en/expert_parallelism.md | 18 ++++++------ docs/source/en/fsdp.md | 2 +- docs/source/en/model_doc/hy_v4.md | 2 +- docs/source/en/model_doc/minimax_m3_vl.md | 2 +- .../distributed/configuration_utils.py | 28 +++++++++++++++---- src/transformers/distributed/mixin.py | 2 +- src/transformers/distributed/utils.py | 4 +++ .../deepseek_v4/test_modeling_deepseek_v4.py | 4 +-- tests/test_fsdp_mixin.py | 2 +- tests/test_tensor_parallel_mixin.py | 2 +- 10 files changed, 44 insertions(+), 22 deletions(-) diff --git a/docs/source/en/expert_parallelism.md b/docs/source/en/expert_parallelism.md index d01c1ad94576..8d23ca6c9feb 100644 --- a/docs/source/en/expert_parallelism.md +++ b/docs/source/en/expert_parallelism.md @@ -19,7 +19,7 @@ rendered properly in your Markdown viewer. ## DistributedConfig -Enable expert parallelism with the [`DistributedConfig`] class and the `enable_expert_parallel` argument. +Enable expert parallelism with the [`DistributedConfig`] class and the `ep_size` argument. The current all-reduce implementation requires `ep_size=tp_size`, so every rank in an expert group receives the same tokens. ```py import os @@ -30,7 +30,7 @@ from transformers.distributed.configuration_utils import DistributedConfig distributed_config = DistributedConfig( tp_size=int(os.environ["WORLD_SIZE"]), - enable_expert_parallel=True, + ep_size=int(os.environ["WORLD_SIZE"]), ) model = AutoModelForCausalLM.from_pretrained( @@ -42,7 +42,7 @@ model = AutoModelForCausalLM.from_pretrained( > [!TIP] > Expert parallelism automatically enables [tensor parallelism](./perf_infer_gpu_multi) for attention layers. -This argument switches to the `ep_plan` (expert parallel plan) defined in each MoE model's config file. The [`GroupedGemmParallel`] class splits expert weights so each device loads only its local experts. The `ep_router` routes tokens to experts and an all-reduce operation combines their outputs. +Setting `ep_size > 1` switches to the `ep_plan` (expert parallel plan) defined in each MoE model's config file. The [`GroupedGemmParallel`] class splits expert weights so each device loads only its local experts. The `ep_router` routes tokens to experts and an all-reduce operation combines their outputs. Launch your inference script with [torchrun](https://pytorch.org/docs/stable/elastic/run.html) and specify how many devices to use. The number of devices must evenly divide the total number of experts. @@ -52,16 +52,16 @@ torchrun --nproc-per-node 8 your_script.py ## Combining with FSDP2 -Expert parallelism only shards the experts. Everything else (attention, embeddings, norms) and its optimizer state is replicated on every expert-parallel rank, which limits how large a model you can train. Add [FSDP2](./fsdp) on a second mesh dimension with `fsdp_size`, and keep using `tp_size` for the expert parallel width (`tp_size` is the EP size). +Expert parallelism only shards the experts. Everything else (attention, embeddings, norms) and its optimizer state is replicated on every expert-parallel rank, which limits how large a model you can train. Add [FSDP2](./fsdp) on a second mesh dimension with `fsdp_size`, and keep `ep_size=tp_size` for the expert parallel width. ```py from transformers import AutoModelForCausalLM from transformers.distributed import DistributedConfig distributed_config = DistributedConfig( - tp_size=4, # expert parallel size + tp_size=4, + ep_size=4, # expert parallel size, must match tp_size fsdp_size=2, # data parallel shards - enable_expert_parallel=True, ) model = AutoModelForCausalLM.from_pretrained("Qwen/Qwen3-30B-A3B", distributed_config=distributed_config) ``` @@ -83,9 +83,9 @@ The table below compares EP-only training with 2D EP+FSDP2 on 8xH100 GPUs. The w ## Mesh views -`DistributedConfig` also accepts an explicit `ep_size`. For the current all-reduce implementation, -set `ep_size=tp_size`; `DistributedConfig(tp_size=4, ep_size=4)` is equivalent to -`DistributedConfig(tp_size=4, enable_expert_parallel=True)`. An explicit `ep_size=1` disables EP. +The legacy `enable_expert_parallel=True` flag is a deprecated alias for `ep_size=tp_size` when `ep_size` +is omitted, and will be removed in v5.20. It emits a `FutureWarning` and leaves `tp_size` and `fsdp_size` +unchanged. An explicit `ep_size` takes precedence over the flag, and `ep_size=1` disables EP. Internally, a mesh manager provides two views of the same ranks: `(pp, fsdp, tp)` for dense layers and `(pp, efsdp, ep)` for experts, where `efsdp_size = fsdp_size * tp_size // ep_size`. diff --git a/docs/source/en/fsdp.md b/docs/source/en/fsdp.md index dbbd629f22d1..5c12d7b55df3 100644 --- a/docs/source/en/fsdp.md +++ b/docs/source/en/fsdp.md @@ -122,7 +122,7 @@ TrainingArguments( > [!TIP] -> For mixture-of-experts models, `fsdp_size` can be combined with `tp_size` and `enable_expert_parallel=True` to shard the experts across one mesh dimension and everything else across the other. See [expert parallelism](./expert_parallelism#combining-with-fsdp2). +> For mixture-of-experts models, `fsdp_size` can be combined with `tp_size` and `ep_size` to shard the experts across one mesh dimension and everything else across the other. See [expert parallelism](./expert_parallelism#combining-with-fsdp2). ## Next steps diff --git a/docs/source/en/model_doc/hy_v4.md b/docs/source/en/model_doc/hy_v4.md index c6b3b31cf969..dd8cd789c622 100644 --- a/docs/source/en/model_doc/hy_v4.md +++ b/docs/source/en/model_doc/hy_v4.md @@ -83,7 +83,7 @@ model = AutoModelForCausalLM.from_pretrained( model = AutoModelForCausalLM.from_pretrained( model_id, dtype=torch.bfloat16, - distributed_config=DistributedConfig(tp_size=16, enable_expert_parallel=True), + distributed_config=DistributedConfig(tp_size=16, ep_size=16), ) ``` diff --git a/docs/source/en/model_doc/minimax_m3_vl.md b/docs/source/en/model_doc/minimax_m3_vl.md index 671d59922a44..2688f8c31ea2 100644 --- a/docs/source/en/model_doc/minimax_m3_vl.md +++ b/docs/source/en/model_doc/minimax_m3_vl.md @@ -167,7 +167,7 @@ model = AutoModelForImageTextToText.from_pretrained( quantization_config=FineGrainedFP8Config(dequantize=True), distributed_config=DistributedConfig( tp_size=int(os.environ["WORLD_SIZE"]), - enable_expert_parallel=True, + ep_size=int(os.environ["WORLD_SIZE"]), ), attn_implementation="kernels-staging/msa@v0", # MSA block-sparse attention kernel ) diff --git a/src/transformers/distributed/configuration_utils.py b/src/transformers/distributed/configuration_utils.py index 814c8335b2fd..ad33f96d39b7 100644 --- a/src/transformers/distributed/configuration_utils.py +++ b/src/transformers/distributed/configuration_utils.py @@ -14,6 +14,7 @@ import json import os +import warnings from dataclasses import asdict, dataclass from typing import Literal @@ -33,8 +34,8 @@ class DistributedConfig: enable_sequence_parallel (`bool`, *optional*, defaults to `False`): Reserved for sequence parallelism. Not wired up yet. enable_expert_parallel (`bool`, *optional*, defaults to `False`): - Route MoE models through the expert-parallel path (``base_model_ep_plan``). When `ep_size` is - omitted, sets it to `tp_size`. An explicit `ep_size` takes precedence. + Deprecated alias for `ep_size=tp_size` when `ep_size` is omitted, removed in v5.20. An explicit + `ep_size` takes precedence. This flag does not change `tp_size` or `fsdp_size`. fsdp_size (`int`, *optional*): Number of devices for FSDP (data parallelism). If `None` and `tp_size` is set, defaults to 1. fsdp_cpu_offload (`bool`, *optional*, defaults to `False`): @@ -44,8 +45,8 @@ class DistributedConfig: pp_size (`int`, *optional*): Number of devices for pipeline parallelism. If `None` and another parallel mode is set, defaults to 1. ep_size (`int`, *optional*): - Number of devices owning distinct expert shards. Defaults to 1, or to `tp_size` when - `enable_expert_parallel=True`. Model execution currently requires `ep_size=tp_size` when EP is enabled. + Number of devices owning distinct expert shards. Defaults to 1. Set it explicitly to enable EP. + Model execution currently requires `ep_size=tp_size` when EP is enabled. """ tp_size: int | None = None @@ -64,6 +65,11 @@ def efsdp_size(self) -> int: return self.fsdp_size * self.tp_size // self.ep_size def __post_init__(self): + self._resolve_parallelism() + self._validate_mesh_config() + + def _resolve_parallelism(self): + """Resolve parallel sizes and legacy EP settings.""" for value in (self.tp_size, self.fsdp_size, self.pp_size, self.ep_size): if value is not None and value < 1: raise ValueError(f"Parallelism sizes must be >= 1, got {value}.") @@ -84,10 +90,22 @@ def __post_init__(self): elif self.tp_size is None: self.tp_size = 1 + if self.enable_expert_parallel and self.ep_size is None: + self.ep_size = self.tp_size + warnings.warn( + f"`enable_expert_parallel` without `ep_size` is deprecated and will be removed in v5.20. " + f"Use ep_size={self.ep_size} instead.", + FutureWarning, + stacklevel=4, + ) + if self.ep_size is None: - self.ep_size = self.tp_size if self.enable_expert_parallel else 1 + self.ep_size = 1 + # Retain the legacy attribute for callers; internal EP decisions use ep_size. self.enable_expert_parallel = self.ep_size > 1 + def _validate_mesh_config(self): + """Validate mesh sizes before the model's expert plan is available.""" if self.ep_size > 1: if self.ep_size % self.tp_size: raise ValueError("`ep_size` must be a multiple of `tp_size`.") diff --git a/src/transformers/distributed/mixin.py b/src/transformers/distributed/mixin.py index 6f4e96b04d28..821eb8f2a733 100644 --- a/src/transformers/distributed/mixin.py +++ b/src/transformers/distributed/mixin.py @@ -88,7 +88,7 @@ def tp_plan(self) -> dict[str, str]: if hasattr(self.config, "distributed_config") and self.config.distributed_config.enable_expert_parallel: if not self._ep_plan: raise ValueError( - f"Expert parallelism was requested (`enable_expert_parallel=True`), but " + f"Expert parallelism was requested (`ep_size > 1`), but " f"`{self.__class__.__name__}` does not define an expert-parallel plan. Add a " f"`base_model_ep_plan` to its config, or disable expert parallelism." ) diff --git a/src/transformers/distributed/utils.py b/src/transformers/distributed/utils.py index c06182d86b2e..aa80eae9c3f5 100644 --- a/src/transformers/distributed/utils.py +++ b/src/transformers/distributed/utils.py @@ -135,6 +135,10 @@ class MeshManager: """Named access to dense and expert parallel axes without exposing their view selection.""" def __init__(self, dense_mesh: DeviceMesh, expert_mesh: DeviceMesh): + """ + dense_mesh: (pp, fsdp, tp) -> attention, dense MLPs, embeddings, lm_heads + expert_mesh: (pp, efsdp, ep) -> experts + """ self._dense_mesh = dense_mesh self._expert_mesh = expert_mesh diff --git a/tests/models/deepseek_v4/test_modeling_deepseek_v4.py b/tests/models/deepseek_v4/test_modeling_deepseek_v4.py index aecaa6584ac1..64f6e89e7d61 100644 --- a/tests/models/deepseek_v4/test_modeling_deepseek_v4.py +++ b/tests/models/deepseek_v4/test_modeling_deepseek_v4.py @@ -454,7 +454,7 @@ def main() -> int: attn_implementation="eager", experts_implementation=LOADTIME_DISPATCH, distributed_config=DistributedConfig( - tp_size=int(os.environ["WORLD_SIZE"]), enable_expert_parallel=True + tp_size=int(os.environ["WORLD_SIZE"]), ep_size=int(os.environ["WORLD_SIZE"]) ), ) model.eval() @@ -533,7 +533,7 @@ def main() -> int: attn_implementation="eager", experts_implementation=LOADTIME_DISPATCH, distributed_config=DistributedConfig( - tp_size=int(os.environ["WORLD_SIZE"]), enable_expert_parallel=True + tp_size=int(os.environ["WORLD_SIZE"]), ep_size=int(os.environ["WORLD_SIZE"]) ), ) model.eval() diff --git a/tests/test_fsdp_mixin.py b/tests/test_fsdp_mixin.py index b59108cee436..efcaed91a5f1 100644 --- a/tests/test_fsdp_mixin.py +++ b/tests/test_fsdp_mixin.py @@ -536,7 +536,7 @@ def _test_fsdp2_expert_parallel_2d_vs_ddp_impl(rank, config_class, config_dict, model = AutoModelForCausalLM.from_pretrained( init_model_dir, torch_dtype=dtype, - distributed_config=DistributedConfig(tp_size=2, fsdp_size=dp, enable_expert_parallel=True), + distributed_config=DistributedConfig(tp_size=2, fsdp_size=dp, ep_size=2), ) assert model.tp_size == 2 and model.fsdp_size == dp assert model._device_mesh.mesh_dim_names == ("pp", "fsdp", "tp") diff --git a/tests/test_tensor_parallel_mixin.py b/tests/test_tensor_parallel_mixin.py index 22d92bdd0948..38f78a0163f3 100644 --- a/tests/test_tensor_parallel_mixin.py +++ b/tests/test_tensor_parallel_mixin.py @@ -394,7 +394,7 @@ def _load_ep_and_reference_models(model_path, model_class): """Load EP model and non-EP reference model for comparison.""" model_ep = model_class.from_pretrained( model_path, - distributed_config=DistributedConfig(tp_size=dist.get_world_size(), enable_expert_parallel=True), + distributed_config=DistributedConfig(tp_size=dist.get_world_size(), ep_size=dist.get_world_size()), ) dist.barrier() From 083e25326a0ed6f2da5a7f3ec4451afeb5d83d2b Mon Sep 17 00:00:00 2001 From: 3outeille Date: Wed, 16 Sep 2026 12:19:49 +0000 Subject: [PATCH 05/86] clean doc --- docs/source/en/expert_parallelism.md | 16 ---------------- 1 file changed, 16 deletions(-) diff --git a/docs/source/en/expert_parallelism.md b/docs/source/en/expert_parallelism.md index 8d23ca6c9feb..74687635dd94 100644 --- a/docs/source/en/expert_parallelism.md +++ b/docs/source/en/expert_parallelism.md @@ -66,8 +66,6 @@ distributed_config = DistributedConfig( model = AutoModelForCausalLM.from_pretrained("Qwen/Qwen3-30B-A3B", distributed_config=distributed_config) ``` -The model is loaded on a `(pp, fsdp, tp)` device mesh with `pp_size=1`, and `tp_size * fsdp_size` must equal the number of processes. The expert parallel plan shards the experts across `tp`, then FSDP2 shards every parameter, experts included, across `fsdp` and owns their gradient reduction. Each `fsdp` rank trains on its own part of the batch. - Load the model as usual, then train with [`Trainer`]. It takes the gradient norm across both meshes and gives each mesh its own optimizer param group. [`~Trainer.save_model`] gathers sharded weights into a regular checkpoint. This requires `accelerate>=1.12` so the `Trainer` can mirror `tp_size` and `fsdp_size` into [`~Accelerate.ParallelismConfig`]. The table below compares EP-only training with 2D EP+FSDP2 on 8xH100 GPUs. The workload is full fine-tuning of Qwen3-30B-A3B in bf16 at sequence length 2048. More FSDP shards cut peak memory, and tokens/s drop some because FSDP2 all-gathers and reduce-scatters the experts across `fsdp`. @@ -81,20 +79,6 @@ The table below compares EP-only training with 2D EP+FSDP2 on 8xH100 GPUs. The w > [!WARNING] > Resuming from a checkpoint is not supported yet for models sharded at load time, so the [`Trainer`] only accepts `save_only_model=True` or `save_strategy="no"` for them. -## Mesh views - -The legacy `enable_expert_parallel=True` flag is a deprecated alias for `ep_size=tp_size` when `ep_size` -is omitted, and will be removed in v5.20. It emits a `FutureWarning` and leaves `tp_size` and `fsdp_size` -unchanged. An explicit `ep_size` takes precedence over the flag, and `ep_size=1` disables EP. - -Internally, a mesh manager provides two views of the same ranks: `(pp, fsdp, tp)` for dense layers -and `(pp, efsdp, ep)` for experts, where `efsdp_size = fsdp_size * tp_size // ep_size`. -Both retain size-one axes, so callers can select dimensions by name. The mesh builder supports -`ep_size` values that are multiples of `tp_size` and divide `fsdp_size * tp_size`. -Model loading currently rejects enabled EP layouts with `ep_size != tp_size` because all-reduce -requires identical tokens within each expert group. Expert sharding and FSDP continue to use the -`tp` and `fsdp` axes of the dense view. - ## API reference [[autodoc]] DistributedConfig From 2e6829699544a2a76d76b8b5486232c2f695e9b3 Mon Sep 17 00:00:00 2001 From: 3outeille Date: Wed, 16 Sep 2026 13:30:18 +0000 Subject: [PATCH 06/86] Use ep_size in the Kimi K2.5 expert parallel example --- docs/source/en/model_doc/kimi_k25.md | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/docs/source/en/model_doc/kimi_k25.md b/docs/source/en/model_doc/kimi_k25.md index 8fc5d85b15b0..20e22ab0cf92 100644 --- a/docs/source/en/model_doc/kimi_k25.md +++ b/docs/source/en/model_doc/kimi_k25.md @@ -49,7 +49,7 @@ import torch from transformers import AutoProcessor, AutoTokenizer, AutoModelForImageTextToText from transformers.distributed.configuration_utils import DistributedConfig -distributed_config = DistributedConfig(enable_expert_parallel=True) +distributed_config = DistributedConfig(tp_size=int(os.environ["WORLD_SIZE"]), ep_size=int(os.environ["WORLD_SIZE"])) processor = AutoProcessor.from_pretrained('moonshotai/Kimi-K2.6') model = AutoModelForImageTextToText.from_pretrained( From 29eb5a360dd4cfd40e13e310f1a52f61a0d1318d Mon Sep 17 00:00:00 2001 From: 3outeille Date: Wed, 16 Sep 2026 13:32:37 +0000 Subject: [PATCH 07/86] Enable expert parallelism in the Mega MoE example --- docs/source/en/experts_interface.md | 7 +++++-- 1 file changed, 5 insertions(+), 2 deletions(-) diff --git a/docs/source/en/experts_interface.md b/docs/source/en/experts_interface.md index baada16c77d7..65c45f532772 100644 --- a/docs/source/en/experts_interface.md +++ b/docs/source/en/experts_interface.md @@ -153,14 +153,17 @@ This backend requires: - A Blackwell GPU (compute capability ≥ 10.0) with a CUDA toolkit (`nvcc`) 12.9 or later. - FP4-packed expert weights paired with UE8M0 weight scales (the pre-quantized checkpoint typically declares `expert_dtype="fp4"` and `scale_fmt="ue8m0"` in its config). -- A `torch.distributed` process group for the expert-parallel group, which the tensor-parallel wrapping supplies automatically. +- A `torch.distributed` process group for the expert-parallel group, which the expert-parallel wrapping supplies automatically when `ep_size > 1`. ```py import os from transformers import AutoModelForCausalLM, DistributedConfig -distributed_config = DistributedConfig(tp_size=int(os.environ["WORLD_SIZE"])) +distributed_config = DistributedConfig( + tp_size=int(os.environ["WORLD_SIZE"]), + ep_size=int(os.environ["WORLD_SIZE"]), +) model = AutoModelForCausalLM.from_pretrained( "deepseek-ai/DeepSeek-V4", experts_implementation="deepgemm_megamoe", From d6090dfd9042eea724a08429e08e254f6885b0ec Mon Sep 17 00:00:00 2001 From: 3outeille Date: Fri, 18 Sep 2026 15:17:34 +0000 Subject: [PATCH 08/86] warn once --- .../distributed/configuration_utils.py | 14 ++++++++------ 1 file changed, 8 insertions(+), 6 deletions(-) diff --git a/src/transformers/distributed/configuration_utils.py b/src/transformers/distributed/configuration_utils.py index ad33f96d39b7..05a0e405264c 100644 --- a/src/transformers/distributed/configuration_utils.py +++ b/src/transformers/distributed/configuration_utils.py @@ -18,6 +18,7 @@ from dataclasses import asdict, dataclass from typing import Literal +from .utils import _get_torch_distributed_rank @dataclass class DistributedConfig: @@ -92,12 +93,13 @@ def _resolve_parallelism(self): if self.enable_expert_parallel and self.ep_size is None: self.ep_size = self.tp_size - warnings.warn( - f"`enable_expert_parallel` without `ep_size` is deprecated and will be removed in v5.20. " - f"Use ep_size={self.ep_size} instead.", - FutureWarning, - stacklevel=4, - ) + if _get_torch_distributed_rank() == 0: + warnings.warn( + f"`enable_expert_parallel` without `ep_size` is deprecated and will be removed in v5.20. " + f"Use ep_size={self.ep_size} instead.", + FutureWarning, + stacklevel=4, + ) if self.ep_size is None: self.ep_size = 1 From 76ff6f0e13101b949ada0cf68a349d09dc88e067 Mon Sep 17 00:00:00 2001 From: 3outeille Date: Wed, 16 Sep 2026 11:05:41 +0000 Subject: [PATCH 09/86] Add dense and expert device mesh views with a MeshManager `initialize_distributed_mesh` now builds two named views of the same ranks: `(pp, fsdp, tp)` for dense layers and `(pp, efsdp, ep)` for experts, both keeping size-one axes so callers select dimensions by name. `MeshManager` routes `ep`/`efsdp` lookups to the expert view and everything else to the dense view. `DistributedConfig` gains `ep_size` (defaults to `tp_size` when `enable_expert_parallel=True`) and `efsdp_size`, with size validation. Model execution is unchanged: expert sharding and FSDP still use the `tp` and `fsdp` axes, and loading rejects `ep_size != tp_size` until the all-to-all dispatcher lands. --- docs/source/en/expert_parallelism.md | 14 ++ tests/test_distributed_config.py | 232 +++++++++++++++++++++++++++ 2 files changed, 246 insertions(+) create mode 100644 tests/test_distributed_config.py diff --git a/docs/source/en/expert_parallelism.md b/docs/source/en/expert_parallelism.md index 74687635dd94..62a34b648322 100644 --- a/docs/source/en/expert_parallelism.md +++ b/docs/source/en/expert_parallelism.md @@ -79,6 +79,20 @@ The table below compares EP-only training with 2D EP+FSDP2 on 8xH100 GPUs. The w > [!WARNING] > Resuming from a checkpoint is not supported yet for models sharded at load time, so the [`Trainer`] only accepts `save_only_model=True` or `save_strategy="no"` for them. +## Mesh views + +`DistributedConfig` also accepts an explicit `ep_size`. For the current all-reduce implementation, +set `ep_size=tp_size`; `DistributedConfig(tp_size=4, ep_size=4)` is equivalent to +`DistributedConfig(tp_size=4, enable_expert_parallel=True)`. An explicit `ep_size=1` disables EP. + +Internally, a mesh manager provides two views of the same ranks: `(pp, fsdp, tp)` for dense layers +and `(pp, efsdp, ep)` for experts, where `efsdp_size = fsdp_size * tp_size // ep_size`. +Both retain size-one axes, so callers can select dimensions by name. The mesh builder supports +`ep_size` values that are multiples of `tp_size` and divide `fsdp_size * tp_size`. +Model loading currently rejects enabled EP layouts with `ep_size != tp_size` because all-reduce +requires identical tokens within each expert group. Expert sharding and FSDP continue to use the +`tp` and `fsdp` axes of the dense view. + ## API reference [[autodoc]] DistributedConfig diff --git a/tests/test_distributed_config.py b/tests/test_distributed_config.py new file mode 100644 index 000000000000..ba9f9c507f66 --- /dev/null +++ b/tests/test_distributed_config.py @@ -0,0 +1,232 @@ +# Copyright 2026 The HuggingFace Team. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +import os +import tempfile +import unittest +from datetime import timedelta +from unittest.mock import patch + +from transformers.distributed import DistributedConfig +from transformers.testing_utils import require_torch, require_torch_greater_or_equal +from transformers.utils import is_torch_available + + +if is_torch_available(): + import torch + import torch.distributed as dist + import torch.multiprocessing as mp + + from transformers.distributed.mixin import DistributedMixin + from transformers.distributed.utils import initialize_distributed_mesh + + +class DistributedConfigTest(unittest.TestCase): + def test_defaults_and_round_trip(self): + for kwargs in ({}, {"tp_size": 4}, {"fsdp_size": 4}, {"pp_size": 4}, {"tp_size": 2, "fsdp_size": 2}): + with self.subTest(kwargs=kwargs): + config = DistributedConfig(**kwargs) + self.assertEqual(config.ep_size, 1) + self.assertFalse(config.enable_expert_parallel) + self.assertEqual(config.efsdp_size, config.fsdp_size * config.tp_size) + self.assertEqual(DistributedConfig.from_dict(config.to_dict()), config) + + def test_legacy_and_explicit_ep_sizes(self): + legacy = DistributedConfig(tp_size=4, fsdp_size=2, enable_expert_parallel=True) + explicit = DistributedConfig(tp_size=4, fsdp_size=2, ep_size=4) + self.assertEqual(legacy, explicit) + self.assertEqual(DistributedConfig.from_dict(explicit.to_dict()), explicit) + for ep_size in (1, 4, 8): + with self.subTest(ep_size=ep_size): + config = DistributedConfig(tp_size=4, fsdp_size=2, ep_size=ep_size, enable_expert_parallel=True) + self.assertEqual(config.ep_size, ep_size) + self.assertEqual(config.enable_expert_parallel, ep_size > 1) + self.assertEqual((config.tp_size, config.fsdp_size), (4, 2)) + + def test_inferred_tp_size(self): + with patch.dict(os.environ, {"WORLD_SIZE": "8"}): + config = DistributedConfig(tp_plan="auto", fsdp_size=2, enable_expert_parallel=True) + self.assertEqual((config.tp_size, config.ep_size, config.efsdp_size), (4, 4, 2)) + + def test_expert_mesh_sizes(self): + for fsdp, tp, ep, efsdp in ((8, 1, 4, 2), (2, 2, 4, 1), (4, 2, 4, 2), (2, 2, 2, 2), (1, 4, 4, 1)): + with self.subTest(fsdp=fsdp, tp=tp, ep=ep): + config = DistributedConfig(fsdp_size=fsdp, tp_size=tp, ep_size=ep) + self.assertEqual(config.efsdp_size, efsdp) + self.assertEqual(DistributedConfig.from_dict(config.to_dict()), config) + + def test_invalid_sizes(self): + for name in ("tp_size", "fsdp_size", "pp_size", "ep_size"): + for value in (0, -1): + with self.subTest(name=name, value=value), self.assertRaisesRegex(ValueError, "must be >= 1"): + DistributedConfig(**{name: value}) + for kwargs, message in ( + ({"tp_size": 4, "ep_size": 2}, "multiple"), + ({"fsdp_size": 4, "ep_size": 3}, "must divide"), + ({"fsdp_size": 2, "pp_size": 2}, "pipeline parallelism"), + ({"ep_size": 2}, "must divide"), + ): + with self.subTest(kwargs=kwargs), self.assertRaisesRegex(ValueError, message): + DistributedConfig(**kwargs) + + +@require_torch +class DistributedMeshValidationTest(unittest.TestCase): + def test_disabled_mesh_does_not_initialize_distributed(self): + with patch("transformers.distributed.utils._ensure_torch_distributed") as initialize: + self.assertEqual(initialize_distributed_mesh(DistributedConfig()), (None, None)) + config, device_map, meshes = DistributedMixin.prepare_distribute_model({}, device_map="cpu") + self.assertEqual(config, DistributedConfig()) + self.assertEqual(device_map, "cpu") + self.assertIsNone(meshes) + initialize.assert_not_called() + + def test_model_loading_rejects_unsupported_ep_layout_before_initialization(self): + with patch("transformers.distributed.mixin.initialize_distributed_mesh") as initialize: + with self.assertRaisesRegex(ValueError, "ep_size=tp_size"): + DistributedMixin.prepare_distribute_model(DistributedConfig(fsdp_size=4, ep_size=2)) + initialize.assert_not_called() + + def test_world_size_mismatch(self): + with ( + patch("transformers.distributed.utils._ensure_torch_distributed"), + patch("torch._C._get_accelerator", return_value=torch.device("cpu")), + patch("torch.distributed.get_world_size", return_value=2), + self.assertRaisesRegex(RuntimeError, "requires 4 processes"), + ): + initialize_distributed_mesh(DistributedConfig(tp_size=4)) + + +def _mesh_worker(rank, rendezvous): + world_size = 4 + dist.init_process_group( + "gloo", + init_method=f"file://{rendezvous}", + rank=rank, + world_size=world_size, + timeout=timedelta(seconds=120), + ) + os.environ["LOCAL_RANK"] = str(rank) + try: + configs = [ + DistributedConfig(fsdp_size=4), + DistributedConfig(tp_size=4), + DistributedConfig(pp_size=4), + DistributedConfig(tp_size=2, fsdp_size=2), + DistributedConfig(tp_size=2, pp_size=2), + ] + configs += [DistributedConfig(fsdp_size=4, ep_size=ep) for ep in (2, 4)] + configs += [DistributedConfig(fsdp_size=2, tp_size=2, ep_size=ep) for ep in (2, 4)] + configs += [DistributedConfig(fsdp_size=1, tp_size=4, ep_size=4)] + for config in configs: + with patch("torch._C._get_accelerator", return_value=torch.device("cpu")): + _, meshes = initialize_distributed_mesh(config) + assert meshes.get_mesh(("pp", "fsdp", "tp")).mesh_dim_names == ("pp", "fsdp", "tp") + for axes in (("pp", "fsdp", "tp"), ("pp", "efsdp", "ep")): + assert meshes.get_mesh(axes).size() == world_size + assert meshes.get_mesh(axes).mesh_dim_names == axes + for name in axes: + assert meshes.get_mesh(name).size() == getattr(config, name + "_size") + assert meshes.get_mesh(("fsdp", "tp")).mesh_dim_names == ("fsdp", "tp") + assert meshes.get_mesh(("efsdp", "ep")).mesh_dim_names == ("efsdp", "ep") + for invalid in ("missing", ("tp", "ep"), ("fsdp", "efsdp")): + try: + meshes.get_mesh(invalid) + except KeyError: + pass + else: + raise AssertionError(f"Accepted invalid mesh dimensions: {invalid}") + stage_size = config.fsdp_size * config.tp_size + stage_start = rank // stage_size * stage_size + expert_rank = (rank - stage_start) % config.ep_size + ep_start = rank // config.ep_size * config.ep_size + assert dist.get_process_group_ranks(meshes.get_mesh("ep").get_group()) == list( + range(ep_start, ep_start + config.ep_size) + ) + assert dist.get_process_group_ranks(meshes.get_mesh("efsdp").get_group()) == list( + range(stage_start + expert_rank, stage_start + stage_size, config.ep_size) + ) + tp_start = rank // config.tp_size * config.tp_size + assert dist.get_process_group_ranks(meshes.get_mesh("tp").get_group()) == list( + range(tp_start, tp_start + config.tp_size) + ) + assert dist.get_process_group_ranks(meshes.get_mesh("fsdp").get_group()) == list( + range(stage_start + rank % config.tp_size, stage_start + stage_size, config.tp_size) + ) + assert dist.get_process_group_ranks(meshes.get_mesh("pp").get_group()) == list( + range(rank % stage_size, world_size, stage_size) + ) + assert meshes.get_mesh("ep").get_group() is meshes.get_mesh("ep").get_group() + finally: + dist.destroy_process_group() + + +def _dense_load_worker(rank, rendezvous): + from transformers import Qwen2Config, Qwen2ForCausalLM + + os.environ.update(RANK=str(rank), LOCAL_RANK=str(rank), WORLD_SIZE="2", LOCAL_WORLD_SIZE="2") + dist.init_process_group( + "gloo", init_method=f"file://{rendezvous}", rank=rank, world_size=2, timeout=timedelta(seconds=120) + ) + try: + torch.manual_seed(42) + config = Qwen2Config( + vocab_size=32, + hidden_size=8, + intermediate_size=8, + num_hidden_layers=2, + num_attention_heads=2, + num_key_value_heads=2, + ) + reference = Qwen2ForCausalLM(config).eval() + source = rendezvous + "_model" + if rank == 0: + reference.save_pretrained(source) + dist.barrier() + inputs = torch.tensor([[1, 2, 3]]) + generation_kwargs = { + "max_new_tokens": 2, + "do_sample": False, + "output_logits": True, + "return_dict_in_generate": True, + } + expected = reference.generate(inputs, **generation_kwargs) + for distributed_config in (DistributedConfig(tp_size=2), DistributedConfig(pp_size=2)): + with patch("torch._C._get_accelerator", return_value=torch.device("cpu")): + model = Qwen2ForCausalLM.from_pretrained(source, distributed_config=distributed_config).eval() + assert model._device_mesh is model._mesh_manager.get_mesh(("pp", "fsdp", "tp")) + actual = model.generate(inputs, **generation_kwargs) + torch.testing.assert_close(actual.sequences, expected.sequences) + torch.testing.assert_close(torch.stack(actual.logits), torch.stack(expected.logits)) + if distributed_config.tp_size > 1: + destination = rendezvous + "_saved" + model.save_pretrained(destination) + dist.barrier() + restored = Qwen2ForCausalLM.from_pretrained(destination).eval() + for name, param in restored.named_parameters(): + torch.testing.assert_close(param, dict(reference.named_parameters())[name], atol=0, rtol=0) + finally: + dist.destroy_process_group() + + +@require_torch +@require_torch_greater_or_equal("2.5") +class DistributedMeshTest(unittest.TestCase): + def test_mesh_groups(self): + with tempfile.TemporaryDirectory() as directory: + mp.spawn(_mesh_worker, args=(os.path.join(directory, "init"),), nprocs=4, join=True) + + def test_dense_load_generate_and_save(self): + with tempfile.TemporaryDirectory() as directory: + mp.spawn(_dense_load_worker, args=(os.path.join(directory, "init"),), nprocs=2, join=True) From 9a1b525f6aa6b31e73a9eee66686987ad0f1f032 Mon Sep 17 00:00:00 2001 From: 3outeille Date: Wed, 16 Sep 2026 11:05:42 +0000 Subject: [PATCH 10/86] cleaning --- tests/test_distributed_config.py | 232 ------------------------------- 1 file changed, 232 deletions(-) delete mode 100644 tests/test_distributed_config.py diff --git a/tests/test_distributed_config.py b/tests/test_distributed_config.py deleted file mode 100644 index ba9f9c507f66..000000000000 --- a/tests/test_distributed_config.py +++ /dev/null @@ -1,232 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import os -import tempfile -import unittest -from datetime import timedelta -from unittest.mock import patch - -from transformers.distributed import DistributedConfig -from transformers.testing_utils import require_torch, require_torch_greater_or_equal -from transformers.utils import is_torch_available - - -if is_torch_available(): - import torch - import torch.distributed as dist - import torch.multiprocessing as mp - - from transformers.distributed.mixin import DistributedMixin - from transformers.distributed.utils import initialize_distributed_mesh - - -class DistributedConfigTest(unittest.TestCase): - def test_defaults_and_round_trip(self): - for kwargs in ({}, {"tp_size": 4}, {"fsdp_size": 4}, {"pp_size": 4}, {"tp_size": 2, "fsdp_size": 2}): - with self.subTest(kwargs=kwargs): - config = DistributedConfig(**kwargs) - self.assertEqual(config.ep_size, 1) - self.assertFalse(config.enable_expert_parallel) - self.assertEqual(config.efsdp_size, config.fsdp_size * config.tp_size) - self.assertEqual(DistributedConfig.from_dict(config.to_dict()), config) - - def test_legacy_and_explicit_ep_sizes(self): - legacy = DistributedConfig(tp_size=4, fsdp_size=2, enable_expert_parallel=True) - explicit = DistributedConfig(tp_size=4, fsdp_size=2, ep_size=4) - self.assertEqual(legacy, explicit) - self.assertEqual(DistributedConfig.from_dict(explicit.to_dict()), explicit) - for ep_size in (1, 4, 8): - with self.subTest(ep_size=ep_size): - config = DistributedConfig(tp_size=4, fsdp_size=2, ep_size=ep_size, enable_expert_parallel=True) - self.assertEqual(config.ep_size, ep_size) - self.assertEqual(config.enable_expert_parallel, ep_size > 1) - self.assertEqual((config.tp_size, config.fsdp_size), (4, 2)) - - def test_inferred_tp_size(self): - with patch.dict(os.environ, {"WORLD_SIZE": "8"}): - config = DistributedConfig(tp_plan="auto", fsdp_size=2, enable_expert_parallel=True) - self.assertEqual((config.tp_size, config.ep_size, config.efsdp_size), (4, 4, 2)) - - def test_expert_mesh_sizes(self): - for fsdp, tp, ep, efsdp in ((8, 1, 4, 2), (2, 2, 4, 1), (4, 2, 4, 2), (2, 2, 2, 2), (1, 4, 4, 1)): - with self.subTest(fsdp=fsdp, tp=tp, ep=ep): - config = DistributedConfig(fsdp_size=fsdp, tp_size=tp, ep_size=ep) - self.assertEqual(config.efsdp_size, efsdp) - self.assertEqual(DistributedConfig.from_dict(config.to_dict()), config) - - def test_invalid_sizes(self): - for name in ("tp_size", "fsdp_size", "pp_size", "ep_size"): - for value in (0, -1): - with self.subTest(name=name, value=value), self.assertRaisesRegex(ValueError, "must be >= 1"): - DistributedConfig(**{name: value}) - for kwargs, message in ( - ({"tp_size": 4, "ep_size": 2}, "multiple"), - ({"fsdp_size": 4, "ep_size": 3}, "must divide"), - ({"fsdp_size": 2, "pp_size": 2}, "pipeline parallelism"), - ({"ep_size": 2}, "must divide"), - ): - with self.subTest(kwargs=kwargs), self.assertRaisesRegex(ValueError, message): - DistributedConfig(**kwargs) - - -@require_torch -class DistributedMeshValidationTest(unittest.TestCase): - def test_disabled_mesh_does_not_initialize_distributed(self): - with patch("transformers.distributed.utils._ensure_torch_distributed") as initialize: - self.assertEqual(initialize_distributed_mesh(DistributedConfig()), (None, None)) - config, device_map, meshes = DistributedMixin.prepare_distribute_model({}, device_map="cpu") - self.assertEqual(config, DistributedConfig()) - self.assertEqual(device_map, "cpu") - self.assertIsNone(meshes) - initialize.assert_not_called() - - def test_model_loading_rejects_unsupported_ep_layout_before_initialization(self): - with patch("transformers.distributed.mixin.initialize_distributed_mesh") as initialize: - with self.assertRaisesRegex(ValueError, "ep_size=tp_size"): - DistributedMixin.prepare_distribute_model(DistributedConfig(fsdp_size=4, ep_size=2)) - initialize.assert_not_called() - - def test_world_size_mismatch(self): - with ( - patch("transformers.distributed.utils._ensure_torch_distributed"), - patch("torch._C._get_accelerator", return_value=torch.device("cpu")), - patch("torch.distributed.get_world_size", return_value=2), - self.assertRaisesRegex(RuntimeError, "requires 4 processes"), - ): - initialize_distributed_mesh(DistributedConfig(tp_size=4)) - - -def _mesh_worker(rank, rendezvous): - world_size = 4 - dist.init_process_group( - "gloo", - init_method=f"file://{rendezvous}", - rank=rank, - world_size=world_size, - timeout=timedelta(seconds=120), - ) - os.environ["LOCAL_RANK"] = str(rank) - try: - configs = [ - DistributedConfig(fsdp_size=4), - DistributedConfig(tp_size=4), - DistributedConfig(pp_size=4), - DistributedConfig(tp_size=2, fsdp_size=2), - DistributedConfig(tp_size=2, pp_size=2), - ] - configs += [DistributedConfig(fsdp_size=4, ep_size=ep) for ep in (2, 4)] - configs += [DistributedConfig(fsdp_size=2, tp_size=2, ep_size=ep) for ep in (2, 4)] - configs += [DistributedConfig(fsdp_size=1, tp_size=4, ep_size=4)] - for config in configs: - with patch("torch._C._get_accelerator", return_value=torch.device("cpu")): - _, meshes = initialize_distributed_mesh(config) - assert meshes.get_mesh(("pp", "fsdp", "tp")).mesh_dim_names == ("pp", "fsdp", "tp") - for axes in (("pp", "fsdp", "tp"), ("pp", "efsdp", "ep")): - assert meshes.get_mesh(axes).size() == world_size - assert meshes.get_mesh(axes).mesh_dim_names == axes - for name in axes: - assert meshes.get_mesh(name).size() == getattr(config, name + "_size") - assert meshes.get_mesh(("fsdp", "tp")).mesh_dim_names == ("fsdp", "tp") - assert meshes.get_mesh(("efsdp", "ep")).mesh_dim_names == ("efsdp", "ep") - for invalid in ("missing", ("tp", "ep"), ("fsdp", "efsdp")): - try: - meshes.get_mesh(invalid) - except KeyError: - pass - else: - raise AssertionError(f"Accepted invalid mesh dimensions: {invalid}") - stage_size = config.fsdp_size * config.tp_size - stage_start = rank // stage_size * stage_size - expert_rank = (rank - stage_start) % config.ep_size - ep_start = rank // config.ep_size * config.ep_size - assert dist.get_process_group_ranks(meshes.get_mesh("ep").get_group()) == list( - range(ep_start, ep_start + config.ep_size) - ) - assert dist.get_process_group_ranks(meshes.get_mesh("efsdp").get_group()) == list( - range(stage_start + expert_rank, stage_start + stage_size, config.ep_size) - ) - tp_start = rank // config.tp_size * config.tp_size - assert dist.get_process_group_ranks(meshes.get_mesh("tp").get_group()) == list( - range(tp_start, tp_start + config.tp_size) - ) - assert dist.get_process_group_ranks(meshes.get_mesh("fsdp").get_group()) == list( - range(stage_start + rank % config.tp_size, stage_start + stage_size, config.tp_size) - ) - assert dist.get_process_group_ranks(meshes.get_mesh("pp").get_group()) == list( - range(rank % stage_size, world_size, stage_size) - ) - assert meshes.get_mesh("ep").get_group() is meshes.get_mesh("ep").get_group() - finally: - dist.destroy_process_group() - - -def _dense_load_worker(rank, rendezvous): - from transformers import Qwen2Config, Qwen2ForCausalLM - - os.environ.update(RANK=str(rank), LOCAL_RANK=str(rank), WORLD_SIZE="2", LOCAL_WORLD_SIZE="2") - dist.init_process_group( - "gloo", init_method=f"file://{rendezvous}", rank=rank, world_size=2, timeout=timedelta(seconds=120) - ) - try: - torch.manual_seed(42) - config = Qwen2Config( - vocab_size=32, - hidden_size=8, - intermediate_size=8, - num_hidden_layers=2, - num_attention_heads=2, - num_key_value_heads=2, - ) - reference = Qwen2ForCausalLM(config).eval() - source = rendezvous + "_model" - if rank == 0: - reference.save_pretrained(source) - dist.barrier() - inputs = torch.tensor([[1, 2, 3]]) - generation_kwargs = { - "max_new_tokens": 2, - "do_sample": False, - "output_logits": True, - "return_dict_in_generate": True, - } - expected = reference.generate(inputs, **generation_kwargs) - for distributed_config in (DistributedConfig(tp_size=2), DistributedConfig(pp_size=2)): - with patch("torch._C._get_accelerator", return_value=torch.device("cpu")): - model = Qwen2ForCausalLM.from_pretrained(source, distributed_config=distributed_config).eval() - assert model._device_mesh is model._mesh_manager.get_mesh(("pp", "fsdp", "tp")) - actual = model.generate(inputs, **generation_kwargs) - torch.testing.assert_close(actual.sequences, expected.sequences) - torch.testing.assert_close(torch.stack(actual.logits), torch.stack(expected.logits)) - if distributed_config.tp_size > 1: - destination = rendezvous + "_saved" - model.save_pretrained(destination) - dist.barrier() - restored = Qwen2ForCausalLM.from_pretrained(destination).eval() - for name, param in restored.named_parameters(): - torch.testing.assert_close(param, dict(reference.named_parameters())[name], atol=0, rtol=0) - finally: - dist.destroy_process_group() - - -@require_torch -@require_torch_greater_or_equal("2.5") -class DistributedMeshTest(unittest.TestCase): - def test_mesh_groups(self): - with tempfile.TemporaryDirectory() as directory: - mp.spawn(_mesh_worker, args=(os.path.join(directory, "init"),), nprocs=4, join=True) - - def test_dense_load_generate_and_save(self): - with tempfile.TemporaryDirectory() as directory: - mp.spawn(_dense_load_worker, args=(os.path.join(directory, "init"),), nprocs=2, join=True) From ce8657f9437d101e0a0fa4e1587e5ce5f328bf29 Mon Sep 17 00:00:00 2001 From: 3outeille Date: Wed, 16 Sep 2026 12:15:08 +0000 Subject: [PATCH 11/86] update --- docs/source/en/expert_parallelism.md | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/docs/source/en/expert_parallelism.md b/docs/source/en/expert_parallelism.md index 62a34b648322..64c67a661799 100644 --- a/docs/source/en/expert_parallelism.md +++ b/docs/source/en/expert_parallelism.md @@ -81,9 +81,9 @@ The table below compares EP-only training with 2D EP+FSDP2 on 8xH100 GPUs. The w ## Mesh views -`DistributedConfig` also accepts an explicit `ep_size`. For the current all-reduce implementation, -set `ep_size=tp_size`; `DistributedConfig(tp_size=4, ep_size=4)` is equivalent to -`DistributedConfig(tp_size=4, enable_expert_parallel=True)`. An explicit `ep_size=1` disables EP. +The legacy `enable_expert_parallel=True` flag is a deprecated alias for `ep_size=tp_size` when `ep_size` +is omitted, and will be removed in v5.20. It emits a `FutureWarning` and leaves `tp_size` and `fsdp_size` +unchanged. An explicit `ep_size` takes precedence over the flag, and `ep_size=1` disables EP. Internally, a mesh manager provides two views of the same ranks: `(pp, fsdp, tp)` for dense layers and `(pp, efsdp, ep)` for experts, where `efsdp_size = fsdp_size * tp_size // ep_size`. From ea3aacba78e34817871b9fe2168af03966c63cf4 Mon Sep 17 00:00:00 2001 From: 3outeille Date: Wed, 16 Sep 2026 12:19:49 +0000 Subject: [PATCH 12/86] clean doc --- docs/source/en/expert_parallelism.md | 14 -------------- 1 file changed, 14 deletions(-) diff --git a/docs/source/en/expert_parallelism.md b/docs/source/en/expert_parallelism.md index 64c67a661799..74687635dd94 100644 --- a/docs/source/en/expert_parallelism.md +++ b/docs/source/en/expert_parallelism.md @@ -79,20 +79,6 @@ The table below compares EP-only training with 2D EP+FSDP2 on 8xH100 GPUs. The w > [!WARNING] > Resuming from a checkpoint is not supported yet for models sharded at load time, so the [`Trainer`] only accepts `save_only_model=True` or `save_strategy="no"` for them. -## Mesh views - -The legacy `enable_expert_parallel=True` flag is a deprecated alias for `ep_size=tp_size` when `ep_size` -is omitted, and will be removed in v5.20. It emits a `FutureWarning` and leaves `tp_size` and `fsdp_size` -unchanged. An explicit `ep_size` takes precedence over the flag, and `ep_size=1` disables EP. - -Internally, a mesh manager provides two views of the same ranks: `(pp, fsdp, tp)` for dense layers -and `(pp, efsdp, ep)` for experts, where `efsdp_size = fsdp_size * tp_size // ep_size`. -Both retain size-one axes, so callers can select dimensions by name. The mesh builder supports -`ep_size` values that are multiples of `tp_size` and divide `fsdp_size * tp_size`. -Model loading currently rejects enabled EP layouts with `ep_size != tp_size` because all-reduce -requires identical tokens within each expert group. Expert sharding and FSDP continue to use the -`tp` and `fsdp` axes of the dense view. - ## API reference [[autodoc]] DistributedConfig From d498e802be6b7b7fb9b6319157cdd2f206d8264a Mon Sep 17 00:00:00 2001 From: Ferdinand Mom <47445085+3outeille@users.noreply.github.com> Date: Sat, 19 Sep 2026 00:07:32 +0900 Subject: [PATCH 13/86] Update src/transformers/distributed/mixin.py Co-authored-by: Arthur <48595927+ArthurZucker@users.noreply.github.com> --- src/transformers/distributed/mixin.py | 1 + 1 file changed, 1 insertion(+) diff --git a/src/transformers/distributed/mixin.py b/src/transformers/distributed/mixin.py index 821eb8f2a733..224dddb63eeb 100644 --- a/src/transformers/distributed/mixin.py +++ b/src/transformers/distributed/mixin.py @@ -154,6 +154,7 @@ def prepare_distribute_model( if distributed_config.ep_size > 1 and distributed_config.ep_size != distributed_config.tp_size: raise ValueError( "All-reduce expert parallelism requires `ep_size=tp_size` and identical tokens per EP group." + "The token dispatch version allowing `ep_size != tp_size` is coming soon!" ) if distributed_config.tp_size == 1 and distributed_config.fsdp_size == 1 and distributed_config.pp_size == 1: From ce3f57604d1e6c3ebb8e36f7a67d88b71006c41c Mon Sep 17 00:00:00 2001 From: 3outeille Date: Fri, 18 Sep 2026 15:42:27 +0000 Subject: [PATCH 14/86] style: ruff import block formatting in distributed config --- src/transformers/distributed/configuration_utils.py | 1 + 1 file changed, 1 insertion(+) diff --git a/src/transformers/distributed/configuration_utils.py b/src/transformers/distributed/configuration_utils.py index 05a0e405264c..9012271ec545 100644 --- a/src/transformers/distributed/configuration_utils.py +++ b/src/transformers/distributed/configuration_utils.py @@ -20,6 +20,7 @@ from .utils import _get_torch_distributed_rank + @dataclass class DistributedConfig: """ From f8415e73595640f0fa62fe70166e1e135d633b37 Mon Sep 17 00:00:00 2001 From: 3outeille Date: Wed, 30 Sep 2026 14:43:29 +0000 Subject: [PATCH 15/86] add more comments --- src/transformers/distributed/mixin.py | 8 +++--- src/transformers/distributed/utils.py | 37 +++++++++++++++++++++------ 2 files changed, 33 insertions(+), 12 deletions(-) diff --git a/src/transformers/distributed/mixin.py b/src/transformers/distributed/mixin.py index 224dddb63eeb..75b799f57e15 100644 --- a/src/transformers/distributed/mixin.py +++ b/src/transformers/distributed/mixin.py @@ -29,7 +29,7 @@ gather_state_dict_for_save, ) from .utils import ( - MeshManager, + TransformersDeviceMesh, _distributed_barrier, _get_torch_distributed_rank, _is_torch_distributed_initialized, @@ -50,7 +50,7 @@ class DistributedMixin: """Distributed orchestration and save/load hooks for [`PreTrainedModel`].""" _device_mesh = None - _mesh_manager: MeshManager | None = None + _mesh_manager: TransformersDeviceMesh | None = None _tp_plan: dict[str, str] | None = None _ep_plan: dict[str, str] | None = None _tp_size = None @@ -144,7 +144,7 @@ def prepare_distribute_model( cls, distributed_config: DistributedConfig | dict | None, device_map=None, - ) -> tuple[DistributedConfig | None, object, MeshManager | None]: + ) -> tuple[DistributedConfig | None, object, TransformersDeviceMesh | None]: if distributed_config is None: return None, device_map, None @@ -174,7 +174,7 @@ def maybe_distribute_model( cls, model: nn.Module, distributed_config: DistributedConfig | None, - mesh_manager: MeshManager | None, + mesh_manager: TransformersDeviceMesh | None, ): """Apply TP or FSDP2 after model init, before weight loading.""" if mesh_manager is not None: diff --git a/src/transformers/distributed/utils.py b/src/transformers/distributed/utils.py index aa80eae9c3f5..39cd8256bd8f 100644 --- a/src/transformers/distributed/utils.py +++ b/src/transformers/distributed/utils.py @@ -131,14 +131,35 @@ def _distributed_barrier(): torch.distributed.barrier() -class MeshManager: - """Named access to dense and expert parallel axes without exposing their view selection.""" +class TransformersDeviceMesh: + """ + Holds the device meshes used by a model. + + dense layers and experts are sharded differently, so they need different views of the same ranks. + + dense : (pp, fsdp, tp) attention, dense MLPs, embeddings, lm_heads + expert : (pp, efsdp, ep) experts + + Both views cover the same world, so pp * fsdp * tp == pp * efsdp * ep. + efsdp is not something you pick, it is whatever is left once ep is fixed: + efsdp = fsdp * tp / ep. It is the FSDP axis for expert weights same role `fsdp` plays for the dense params. + + There is no etp (expert tensor parallel) axis yet meaning experts are never tensor-sharded here. + If one were ever added, the identity would become pp * efsdp * ep * etp == pp * fsdp * tp and efsdp would shrink by etp + (efsdp = fsdp * tp / (ep * etp)) + + Why not reuse fsdp mesh ? When ep_size == tp_size, efsdp == fsdp, both in size and in which + ranks are grouped together, so the fsdp axis of the dense mesh would work for experts too. + As soon as ep_size != tp_size the two group different ranks and you need a separate axis. + + Regarding ep value, We decide to default it to node width (8 on most machines) so all-to-all never leaves the node. + - On a single node, ep == fsdp * tp thus efsdp = 1, the axis does nothing. + - On several nodes, we still keep ep at node width, since all-to-all across nodes is expensive. + However, each node then holds a full copy of the expert group and efsdp is the number of copies, which is where FSDP happens for the experts + i.e: 2 nodes x 8 GPUs -> efsdp = 16 / 8 = 2, one EP group per node, two copies, sharded over efsdp. + """ def __init__(self, dense_mesh: DeviceMesh, expert_mesh: DeviceMesh): - """ - dense_mesh: (pp, fsdp, tp) -> attention, dense MLPs, embeddings, lm_heads - expert_mesh: (pp, efsdp, ep) -> experts - """ self._dense_mesh = dense_mesh self._expert_mesh = expert_mesh @@ -242,7 +263,7 @@ def initialize_fully_sharded_data_parallelism(distributed_config: DistributedCon def initialize_distributed_mesh( distributed_config: DistributedConfig, -) -> tuple[torch.device | None, MeshManager | None]: +) -> tuple[torch.device | None, TransformersDeviceMesh | None]: """Create a device mesh containing every configured parallel dimension.""" mesh_shape = (distributed_config.pp_size, distributed_config.fsdp_size, distributed_config.tp_size) if mesh_shape == (1, 1, 1): @@ -277,7 +298,7 @@ def initialize_distributed_mesh( (distributed_config.pp_size, distributed_config.efsdp_size, distributed_config.ep_size), mesh_dim_names=("pp", "efsdp", "ep"), ) - return device_map, MeshManager(dense_mesh, expert_mesh) + return device_map, TransformersDeviceMesh(dense_mesh, expert_mesh) def gather_full_state_dict(model) -> dict[str, torch.Tensor]: From a115da8ebd77cd665cee9c87eb58b760e1d30f13 Mon Sep 17 00:00:00 2001 From: 3outeille Date: Fri, 18 Sep 2026 17:22:06 +0000 Subject: [PATCH 16/86] Decouple tp_plan and ep_plan for expert parallelism. Keep TP and EP plans separate on the model, resolve overrides via resolve_parallel_plans, and apply both through tensor parallel sharding on the tp mesh when ep_size matches tp_size. --- docs/source/en/expert_parallelism.md | 25 +- .../distributed/configuration_utils.py | 26 ++- src/transformers/distributed/mixin.py | 83 ++++--- .../distributed/tensor_parallel.py | 88 +++++-- src/transformers/modeling_utils.py | 4 +- tests/tensor_parallel/test_tensor_parallel.py | 214 +++++++++++++++++- tests/test_modeling_common.py | 7 +- tests/test_tensor_parallel_mixin.py | 7 +- 8 files changed, 376 insertions(+), 78 deletions(-) diff --git a/docs/source/en/expert_parallelism.md b/docs/source/en/expert_parallelism.md index 74687635dd94..785c7f7be2f2 100644 --- a/docs/source/en/expert_parallelism.md +++ b/docs/source/en/expert_parallelism.md @@ -39,10 +39,12 @@ model = AutoModelForCausalLM.from_pretrained( ) ``` -> [!TIP] -> Expert parallelism automatically enables [tensor parallelism](./perf_infer_gpu_multi) for attention layers. +Each MoE model defines two plans in its config: `base_model_tp_plan` for the dense modules and `base_model_ep_plan` for the experts. They are exposed on the loaded model as `model.tp_plan` and `model.ep_plan`. With `tp_size > 1` and `ep_size > 1`, both apply: the [tensor parallel](./perf_infer_gpu_multi) plan shards attention and the dense MLPs, and the expert parallel plan shards the experts. EP rules take precedence over TP rules for the same modules, so expert weights are sharded once, by the EP plan. In the EP plan, the [`GroupedGemmParallel`] style splits the expert weights along the expert dimension so each rank loads only its local experts, and `ep_router` masks the experts that live on other ranks before an all-reduce combines the expert outputs. + +`tp_plan` is applied only when `tp_size > 1`, and `ep_plan` only when `ep_size > 1`. With TP enabled and EP disabled, the full TP plan applies, expert rules included. -Setting `ep_size > 1` switches to the `ep_plan` (expert parallel plan) defined in each MoE model's config file. The [`GroupedGemmParallel`] class splits expert weights so each device loads only its local experts. The `ep_router` routes tokens to experts and an all-reduce operation combines their outputs. +> [!TIP] +> `enable_expert_parallel=True` is a deprecated alias for `ep_size=tp_size`, used only when `ep_size` is omitted, and emits a `FutureWarning`. Launch your inference script with [torchrun](https://pytorch.org/docs/stable/elastic/run.html) and specify how many devices to use. The number of devices must evenly divide the total number of experts. @@ -50,9 +52,24 @@ Launch your inference script with [torchrun](https://pytorch.org/docs/stable/ela torchrun --nproc-per-node 8 your_script.py ``` +### Overriding the plans + +Pass `tp_plan={...}` or `ep_plan={...}` to [`DistributedConfig`] to override individual rules of the predefined plans. Unspecified rules are kept, and the merged plans are stored on the model. Each key must match a module, a parameter, or an existing plan entry; otherwise loading raises a `ValueError` before anything is sharded. Use the full path as seen from the loaded model, so `model.layers.*` for a causal LM and `layers.*` for its base model. + +```py +distributed_config = DistributedConfig( + tp_size=4, + ep_size=4, + tp_plan={"model.layers.*.self_attn.q_proj": "colwise_rep"}, + ep_plan={"model.layers.*.mlp.experts.down_proj": "grouped_gemm"}, +) +``` + +Providing a plan does not infer parallel sizes: set `tp_size` and `ep_size` explicitly. + ## Combining with FSDP2 -Expert parallelism only shards the experts. Everything else (attention, embeddings, norms) and its optimizer state is replicated on every expert-parallel rank, which limits how large a model you can train. Add [FSDP2](./fsdp) on a second mesh dimension with `fsdp_size`, and keep `ep_size=tp_size` for the expert parallel width. +Tensor and expert parallelism shard the weights across `tp`, but the optimizer state and the modules without a rule are still replicated on every rank of the group, which limits how large a model you can train. Add [FSDP2](./fsdp) on a second mesh dimension with `fsdp_size`, and keep `ep_size=tp_size` for the expert parallel width. ```py from transformers import AutoModelForCausalLM diff --git a/src/transformers/distributed/configuration_utils.py b/src/transformers/distributed/configuration_utils.py index 9012271ec545..91926709897b 100644 --- a/src/transformers/distributed/configuration_utils.py +++ b/src/transformers/distributed/configuration_utils.py @@ -18,8 +18,6 @@ from dataclasses import asdict, dataclass from typing import Literal -from .utils import _get_torch_distributed_rank - @dataclass class DistributedConfig: @@ -32,7 +30,8 @@ class DistributedConfig: `WORLD_SIZE // (other_parallel_size)`. If `None` and no `tp_plan` is set, defaults to 1. tp_plan (`dict[str, str]` or `"auto"`, *optional*): Tensor parallel sharding plan. Pass `"auto"`, or leave as `None` when `tp_size` is set, to use the - model's predefined `base_model_tp_plan`. Pass a dictionary to override the predefined plan. + model's predefined `base_model_tp_plan`. Pass a dictionary to override individual rules of that plan; + unspecified rules are kept. enable_sequence_parallel (`bool`, *optional*, defaults to `False`): Reserved for sequence parallelism. Not wired up yet. enable_expert_parallel (`bool`, *optional*, defaults to `False`): @@ -49,6 +48,10 @@ class DistributedConfig: ep_size (`int`, *optional*): Number of devices owning distinct expert shards. Defaults to 1. Set it explicitly to enable EP. Model execution currently requires `ep_size=tp_size` when EP is enabled. + ep_plan (`dict[str, str]`, *optional*): + Expert parallel sharding plan. Leave as `None` to use the model's predefined `base_model_ep_plan`. Pass a + dictionary to override individual rules of that plan; unspecified rules are kept. Applied only when + `ep_size > 1`, and its rules take precedence over `tp_plan` rules for the same modules. """ tp_size: int | None = None @@ -60,6 +63,7 @@ class DistributedConfig: fsdp_mixed_precision: bool = False pp_size: int | None = None ep_size: int | None = None + ep_plan: dict[str, str] | None = None @property def efsdp_size(self) -> int: @@ -94,13 +98,12 @@ def _resolve_parallelism(self): if self.enable_expert_parallel and self.ep_size is None: self.ep_size = self.tp_size - if _get_torch_distributed_rank() == 0: - warnings.warn( - f"`enable_expert_parallel` without `ep_size` is deprecated and will be removed in v5.20. " - f"Use ep_size={self.ep_size} instead.", - FutureWarning, - stacklevel=4, - ) + warnings.warn( + f"`enable_expert_parallel` without `ep_size` is deprecated and will be removed in v5.20. " + f"Use ep_size={self.ep_size} instead.", + FutureWarning, + stacklevel=4, + ) if self.ep_size is None: self.ep_size = 1 @@ -109,6 +112,9 @@ def _resolve_parallelism(self): def _validate_mesh_config(self): """Validate mesh sizes before the model's expert plan is available.""" + if self.ep_plan is not None and not isinstance(self.ep_plan, dict): + raise ValueError("`ep_plan` must be a dictionary or None.") + if self.ep_size > 1: if self.ep_size % self.tp_size: raise ValueError("`ep_size` must be a multiple of `tp_size`.") diff --git a/src/transformers/distributed/mixin.py b/src/transformers/distributed/mixin.py index 75b799f57e15..4ae276547185 100644 --- a/src/transformers/distributed/mixin.py +++ b/src/transformers/distributed/mixin.py @@ -24,9 +24,10 @@ from .fsdp import apply_fully_sharded_data_parallelism, is_fsdp_managed_module from .pipeline_parallel import apply_pipeline_parallelism from .tensor_parallel import ( - _validate_tp_plan_styles, + _validate_parallel_plan_styles, apply_tensor_parallelism, gather_state_dict_for_save, + resolve_parallel_plans, ) from .utils import ( TransformersDeviceMesh, @@ -84,16 +85,13 @@ def init_parallel_plans(self) -> None: @property def tp_plan(self) -> dict[str, str]: - """The full tp plan for the model's modules.""" - if hasattr(self.config, "distributed_config") and self.config.distributed_config.enable_expert_parallel: - if not self._ep_plan: - raise ValueError( - f"Expert parallelism was requested (`ep_size > 1`), but " - f"`{self.__class__.__name__}` does not define an expert-parallel plan. Add a " - f"`base_model_ep_plan` to its config, or disable expert parallelism." - ) - return self._ep_plan - return self._tp_plan + """The full tensor parallel plan for the model's modules.""" + return self._tp_plan if self._tp_plan is not None else {} + + @property + def ep_plan(self) -> dict[str, str]: + """The full expert parallel plan for the model's modules, kept separate from `tp_plan`.""" + return self._ep_plan if self._ep_plan is not None else {} @property def fsdp_plan(self) -> dict[str, str]: @@ -111,7 +109,7 @@ def tp_plan(self, plan: dict[str, str] | None): if not isinstance(plan, dict): raise ValueError("Can only set a dictionary as `tp_plan`") - _validate_tp_plan_styles(plan) + _validate_parallel_plan_styles(plan) model_param_names = [name for name, _ in self.named_parameters()] for layer_pattern in plan.keys(): @@ -129,6 +127,17 @@ def tp_plan(self, plan: dict[str, str] | None): self._tp_plan = plan + @ep_plan.setter + def ep_plan(self, plan: dict[str, str] | None): + if plan is None: + self._ep_plan = {} + return + if not isinstance(plan, dict): + raise ValueError("Can only set a dictionary as `ep_plan`") + + _validate_parallel_plan_styles(plan) + self._ep_plan = plan + @pp_plan.setter def pp_plan(self, plan: dict[str, tuple[str, str]] | None): if plan is None: @@ -154,7 +163,6 @@ def prepare_distribute_model( if distributed_config.ep_size > 1 and distributed_config.ep_size != distributed_config.tp_size: raise ValueError( "All-reduce expert parallelism requires `ep_size=tp_size` and identical tokens per EP group." - "The token dispatch version allowing `ep_size != tp_size` is coming soon!" ) if distributed_config.tp_size == 1 and distributed_config.fsdp_size == 1 and distributed_config.pp_size == 1: @@ -176,29 +184,32 @@ def maybe_distribute_model( distributed_config: DistributedConfig | None, mesh_manager: TransformersDeviceMesh | None, ): - """Apply TP or FSDP2 after model init, before weight loading.""" - if mesh_manager is not None: - model.config.distributed_config = distributed_config - model._mesh_manager = mesh_manager - model._device_mesh = mesh_manager.get_mesh(("pp", "fsdp", "tp")) - model._tp_size = distributed_config.tp_size - model._fsdp_size = distributed_config.fsdp_size - - if distributed_config.pp_size > 1: - pp_mesh = mesh_manager.get_mesh("pp") - model = apply_pipeline_parallelism(model, pp_mesh) - - # Both may apply: the tensor/expert parallel plan shards across `tp` first, then FSDP2 - # shards every parameter (the `tp`-sharded ones included) across `fsdp`. - if distributed_config.tp_size > 1: - tp_mesh = mesh_manager.get_mesh("tp") - if isinstance(distributed_config.tp_plan, dict): - model.tp_plan = distributed_config.tp_plan - model = apply_tensor_parallelism(model, tp_mesh) - - if distributed_config.fsdp_size > 1: - fsdp_mesh = mesh_manager.get_mesh("fsdp") - model = apply_fully_sharded_data_parallelism(model, fsdp_mesh) + """Apply pipeline, tensor and expert parallelism, then FSDP2, after model init and before weight loading.""" + if mesh_manager is None: + return model + + model.config.distributed_config = distributed_config + model._mesh_manager = mesh_manager + model._device_mesh = mesh_manager.get_mesh(("pp", "fsdp", "tp")) + model._tp_size = distributed_config.tp_size + model._fsdp_size = distributed_config.fsdp_size + + # Resolve both plans before sharding anything: overrides are merged into `model.tp_plan` / `model.ep_plan`, + # and the experts named by the EP plan are removed from the TP plan so they are sharded once. + tp_plan, ep_plan = resolve_parallel_plans(model, distributed_config) + + if distributed_config.pp_size > 1: + model = apply_pipeline_parallelism(model, mesh_manager.get_mesh("pp")) + + tp_mesh = mesh_manager.get_mesh("tp") + if tp_plan: + model = apply_tensor_parallelism(model, tp_mesh, tp_plan) + if ep_plan: + # Legacy masked EP: the EP group is the TP group, every rank keeps every token. + model = apply_tensor_parallelism(model, tp_mesh, ep_plan) + + if distributed_config.fsdp_size > 1: + model = apply_fully_sharded_data_parallelism(model, mesh_manager.get_mesh("fsdp")) return model def should_save_on_this_rank(self, is_main_process: bool) -> bool: diff --git a/src/transformers/distributed/tensor_parallel.py b/src/transformers/distributed/tensor_parallel.py index 3151f6264a12..99973c5abf09 100644 --- a/src/transformers/distributed/tensor_parallel.py +++ b/src/transformers/distributed/tensor_parallel.py @@ -15,12 +15,21 @@ import contextlib import re +from fnmatch import fnmatchcase +from typing import TYPE_CHECKING from ..utils import logging from ..utils.generic import GeneralInterface from ..utils.import_utils import is_torch_available, is_torch_distributed_available +if TYPE_CHECKING: + from torch import nn + from torch.distributed.device_mesh import DeviceMesh + + from .configuration_utils import DistributedConfig + + logger = logging.get_logger(__name__) if is_torch_available(): @@ -71,21 +80,21 @@ def verify_tp_plan(expected_keys: list[str], tp_plan: dict[str, str] | None): logger.warning(f"The following layers were not sharded: {', '.join(unsharded_layers)}") -def _get_parameter_tp_plan(parameter_name: str, tp_plan: dict[str, str], is_weight=True) -> str | None: +def _get_parameter_plan(parameter_name: str, plan: dict[str, str], is_weight=True) -> str | None: """ - Get the TP style for a parameter from the TP plan. + Get the parallel style for a parameter or module from a TP or EP plan. - The TP plan is a dictionary that maps parameter names to TP styles. + The plan is a dictionary that maps parameter or module names to parallel styles. The parameter name can be a generic name with wildcards (e.g. "*.weight") or a specific name (e.g. "layer_1.weight"). The `is_weight` is important because for weights, we want to support `.weights` and `.bias` cases seamlessly! but not parent classes for `post_init` calls """ generic_param_name = replace_layer_number_by_wildcard(parameter_name) - if generic_param_name in tp_plan: - return tp_plan[generic_param_name] - elif is_weight and "." in generic_param_name and (module_name := generic_param_name.rsplit(".", 1)[0]) in tp_plan: - return tp_plan[module_name] + if generic_param_name in plan: + return plan[generic_param_name] + elif is_weight and "." in generic_param_name and (module_name := generic_param_name.rsplit(".", 1)[0]) in plan: + return plan[module_name] return None @@ -792,32 +801,77 @@ class ParallelInterface(GeneralInterface): ALL_PARALLEL_STYLES: ParallelInterface = ParallelInterface() -def _validate_tp_plan_styles(tp_plan: dict[str, str] | None) -> None: - unsupported_styles = {style for style in (tp_plan or {}).values() if style not in ALL_PARALLEL_STYLES} +def _validate_parallel_plan_styles(plan: dict[str, str] | None) -> None: + unsupported_styles = {style for style in (plan or {}).values() if style not in ALL_PARALLEL_STYLES} if unsupported_styles: raise ValueError( - f"Unsupported tensor parallel styles: {unsupported_styles}. " - f"Supported styles are {list(ALL_PARALLEL_STYLES.keys())}" + f"Unsupported parallel styles: {unsupported_styles}. Supported styles are {list(ALL_PARALLEL_STYLES.keys())}" + ) + + +def resolve_parallel_plans( + model: nn.Module, distributed_config: DistributedConfig +) -> tuple[dict[str, str], dict[str, str]]: + """Merge the `DistributedConfig` overrides into the model's plans and split them between TP and EP. + + Returns the TP plan to apply to the dense modules and the EP plan to apply to the experts. Each plan is empty + when its parallel size is 1. EP owns every module it names, so TP rules for those modules and their children + are dropped: expert weights are sharded once, by the EP plan. + """ + # Reject invalid paths before merging, e.g. "layers.*" when the model uses "model.layers.*". + layer_names = {name for name, _ in model.named_modules()} | {name for name, _ in model.named_parameters()} + layer_names |= {replace_layer_number_by_wildcard(name) for name in layer_names} + for plan_name in ("tp_plan", "ep_plan"): + override = getattr(distributed_config, plan_name) + if isinstance(override, dict): + valid_names = layer_names | set(getattr(model, plan_name)) + for pattern in override: + if pattern not in valid_names: + raise ValueError( + f"The `{plan_name}` pattern {pattern!r} does not match any module, parameter, " + f"or existing plan entry in {type(model).__name__}. " + "Check the full path, including any 'model.' prefix." + ) + + if isinstance(distributed_config.tp_plan, dict): + model._tp_plan = model.tp_plan | distributed_config.tp_plan + if isinstance(distributed_config.ep_plan, dict): + model._ep_plan = model.ep_plan | distributed_config.ep_plan + + tp_plan = dict(model.tp_plan) if distributed_config.tp_size > 1 else {} + ep_plan = dict(model.ep_plan) if distributed_config.ep_size > 1 else {} + if distributed_config.ep_size > 1 and not ep_plan: + raise ValueError( + f"Expert parallelism was requested (`ep_size={distributed_config.ep_size}`), but `{type(model).__name__}` " + "does not define an expert-parallel plan. Pass `ep_plan` in `DistributedConfig`, add a " + "`base_model_ep_plan` to the model's config, or disable expert parallelism." ) + def is_expert_path(name: str) -> bool: + return any(fnmatchcase(name, path) or fnmatchcase(name, path + ".*") for path in ep_plan) + + tp_plan = {name: style for name, style in tp_plan.items() if not is_expert_path(name)} + _validate_parallel_plan_styles(tp_plan) + _validate_parallel_plan_styles(ep_plan) + return tp_plan, ep_plan -def apply_tensor_parallelism(model, tp_mesh): - """DTensor backend: shard params as placeholders and install TP forward hooks.""" - _validate_tp_plan_styles(model.tp_plan) +def apply_tensor_parallelism(model: nn.Module, tp_mesh: DeviceMesh, plan: dict[str, str] | None = None): + plan = model.tp_plan if plan is None else plan + _validate_parallel_plan_styles(plan) for name, module in model.named_modules(): # Create DTensor placeholders so the loader knows which shard belongs to this rank. for p_name, _ in list(module.named_parameters(recurse=False)): full = f"{name}.{p_name}" if name else p_name - style_name = _get_parameter_tp_plan(parameter_name=full, tp_plan=model.tp_plan, is_weight=True) + style_name = _get_parameter_plan(parameter_name=full, plan=plan, is_weight=True) if style_name is not None and style_name in ALL_PARALLEL_STYLES: style = ALL_PARALLEL_STYLES[style_name] style.validate_param(module, p_name, tp_mesh, parameter_name=full) style.shard_param(module, p_name, tp_mesh) - # Install the input/output transforms required by this module's TP style. - style_name = _get_parameter_tp_plan(parameter_name=name, tp_plan=model.tp_plan, is_weight=False) + # Install the input/output transforms required by this module's style. + style_name = _get_parameter_plan(parameter_name=name, plan=plan, is_weight=False) if style_name is not None and style_name in ALL_PARALLEL_STYLES: if style_name == "mla_kv_a_proj": # MLA needs to know the qk_rope_head_dim to split the projection output into KV and RoPE parts. diff --git a/src/transformers/modeling_utils.py b/src/transformers/modeling_utils.py index 07291b2870f6..ccb9eb334ad4 100644 --- a/src/transformers/modeling_utils.py +++ b/src/transformers/modeling_utils.py @@ -57,7 +57,7 @@ from .distributed import DistributedConfig from .distributed.mixin import DistributedMixin from .distributed.sharding_utils import _dtensor_from_local_like -from .distributed.tensor_parallel import _get_parameter_tp_plan, verify_tp_plan +from .distributed.tensor_parallel import _get_parameter_plan, verify_tp_plan from .distributed.utils import ( _get_torch_distributed_world_size, _is_torch_distributed_initialized, @@ -5013,7 +5013,7 @@ def get_total_byte_count( param_byte_count = param.numel() * dtype_size if len(tp_plan) > 0: - is_part_of_plan = _get_parameter_tp_plan(param_name, tp_plan, is_weight=True) is not None + is_part_of_plan = _get_parameter_plan(param_name, tp_plan, is_weight=True) is not None param_byte_count //= _get_torch_distributed_world_size() if is_part_of_plan else 1 total_byte_count[device] += param_byte_count diff --git a/tests/tensor_parallel/test_tensor_parallel.py b/tests/tensor_parallel/test_tensor_parallel.py index c83741361fd5..c39930ef2307 100644 --- a/tests/tensor_parallel/test_tensor_parallel.py +++ b/tests/tensor_parallel/test_tensor_parallel.py @@ -16,8 +16,9 @@ import torch -from transformers import AutoModelForCausalLM +from transformers import AutoModelForCausalLM, Qwen3MoeConfig, Qwen3MoeForCausalLM, Qwen3MoeModel from transformers.distributed import tensor_parallel +from transformers.distributed.configuration_utils import DistributedConfig from transformers.distributed.sharding_utils import DtensorShardOperation from transformers.distributed.tensor_parallel import ( ALL_PARALLEL_STYLES, @@ -26,7 +27,216 @@ PackedRowwiseParallel, RowwiseParallel, ) -from transformers.testing_utils import TestCasePlus, is_tensor_parallel_test +from transformers.testing_utils import TestCasePlus, is_tensor_parallel_test, require_torch + + +# Qwen3 MoE's predefined plans, as resolved on `Qwen3MoeModel` (no `model.` prefix). +DENSE_TP_PLAN = { + "layers.*.self_attn.q_proj": "colwise", + "layers.*.self_attn.k_proj": "colwise", + "layers.*.self_attn.v_proj": "colwise", + "layers.*.self_attn.q_norm": "replicated_with_grad_allreduce", + "layers.*.self_attn.k_norm": "replicated_with_grad_allreduce", + "layers.*.self_attn.o_proj": "rowwise", + "layers.*.mlp.gate_proj": "colwise", + "layers.*.mlp.up_proj": "colwise", + "layers.*.mlp.down_proj": "rowwise", +} +EXPERT_TP_PLAN = { + "layers.*.mlp.experts.gate_up_proj": "packed_colwise", + "layers.*.mlp.experts.down_proj": "rowwise", + "layers.*.mlp.experts": "moe_tp_experts", +} +EP_PLAN = { + "layers.*.mlp.gate": "ep_router", + "layers.*.mlp.experts.gate_up_proj": "grouped_gemm", + "layers.*.mlp.experts.down_proj": "grouped_gemm", + "layers.*.mlp.experts": "moe_tp_experts", +} + + +@require_torch +class TestParallelPlanResolution(TestCasePlus): + def setUp(self): + super().setUp() + self.config = Qwen3MoeConfig( + vocab_size=32, + hidden_size=16, + intermediate_size=32, + moe_intermediate_size=8, + num_hidden_layers=1, + num_attention_heads=4, + num_key_value_heads=4, + head_dim=4, + num_experts=4, + num_experts_per_tok=2, + ) + with torch.device("meta"): + self.model = Qwen3MoeModel(self.config) + + def test_ep_plan_setter(self): + self.model.ep_plan = None + self.assertEqual(self.model.ep_plan, {}) + with self.assertRaisesRegex(ValueError, "Can only set a dictionary"): + self.model.ep_plan = "auto" + with self.assertRaisesRegex(ValueError, "Unsupported parallel styles"): + self.model.ep_plan = {"layers.*.mlp.experts": "invalid_style"} + self.model.ep_plan = EP_PLAN + self.assertEqual(self.model.ep_plan, EP_PLAN) + + def test_disabled_parallelism_has_no_plans(self): + for config in (DistributedConfig(), DistributedConfig(fsdp_size=8), DistributedConfig(pp_size=2)): + with self.subTest(config=config): + self.assertEqual(tensor_parallel.resolve_parallel_plans(self.model, config), ({}, {})) + + def test_tp_only_keeps_experts_in_tp_plan(self): + tp_plan, ep_plan = tensor_parallel.resolve_parallel_plans(self.model, DistributedConfig(tp_size=4)) + self.assertEqual(tp_plan, DENSE_TP_PLAN | EXPERT_TP_PLAN) + self.assertEqual(ep_plan, {}) + + def test_ep_takes_experts_and_router_out_of_tp_plan(self): + for config in ( + DistributedConfig(tp_size=4, ep_size=4), + DistributedConfig(tp_size=2, fsdp_size=2, ep_size=2), + ): + with self.subTest(config=config): + tp_plan, ep_plan = tensor_parallel.resolve_parallel_plans(self.model, config) + self.assertEqual(tp_plan, DENSE_TP_PLAN) + self.assertEqual(ep_plan, EP_PLAN) + + def test_legacy_flag_is_an_alias_for_ep_size(self): + with warnings.catch_warnings(record=True) as caught: + warnings.simplefilter("always") + legacy = DistributedConfig(tp_size=4, enable_expert_parallel=True) + self.assertEqual([w.category for w in caught], [FutureWarning]) + self.assertIn("Use ep_size=4 instead", str(caught[0].message)) + with warnings.catch_warnings(record=True) as caught: + warnings.simplefilter("always") + explicit = DistributedConfig(tp_size=4, ep_size=4) + disabled = DistributedConfig(tp_size=4, ep_size=1, enable_expert_parallel=True) + self.assertEqual(caught, []) + self.assertEqual(legacy, explicit) + self.assertFalse(disabled.enable_expert_parallel) + self.assertEqual( + tensor_parallel.resolve_parallel_plans(self.model, legacy), + tensor_parallel.resolve_parallel_plans(self.model, explicit), + ) + + def test_ep_plan_is_a_dict_and_round_trips(self): + with self.assertRaisesRegex(ValueError, "`ep_plan` must be a dictionary or None"): + DistributedConfig(tp_size=4, ep_size=4, ep_plan="auto") + config = DistributedConfig(tp_size=4, ep_size=4, ep_plan={"layers.*.mlp.gate": "ep_router"}) + self.assertEqual(config.to_dict()["ep_plan"], {"layers.*.mlp.gate": "ep_router"}) + self.assertEqual(DistributedConfig.from_dict(config.to_dict()), config) + + def test_overrides_merge_into_the_predefined_plans(self): + config = DistributedConfig( + tp_size=4, + ep_size=4, + tp_plan={"layers.*.self_attn.q_proj": "colwise_rep"}, + ep_plan={"layers.*.mlp.experts.down_proj": "rowwise"}, + ) + tp_plan, ep_plan = tensor_parallel.resolve_parallel_plans(self.model, config) + self.assertEqual(tp_plan, DENSE_TP_PLAN | {"layers.*.self_attn.q_proj": "colwise_rep"}) + self.assertEqual(ep_plan, EP_PLAN | {"layers.*.mlp.experts.down_proj": "rowwise"}) + # The merged plans are stored on the model, the config defaults are untouched. + self.assertEqual(self.model.tp_plan["layers.*.self_attn.q_proj"], "colwise_rep") + self.assertEqual(self.model.ep_plan["layers.*.mlp.experts.down_proj"], "rowwise") + self.assertEqual(self.model.config.base_model_tp_plan["layers.*.self_attn.q_proj"], "colwise") + self.assertEqual(self.model.config.base_model_ep_plan["layers.*.mlp.experts.down_proj"], "grouped_gemm") + # The overrides are not rewritten with the merged plans. + self.assertEqual(config.tp_plan, {"layers.*.self_attn.q_proj": "colwise_rep"}) + self.assertEqual(config.ep_plan, {"layers.*.mlp.experts.down_proj": "rowwise"}) + # The merged EP plan stays on the model but is not applied while EP is disabled. + tp_plan, ep_plan = tensor_parallel.resolve_parallel_plans(self.model, DistributedConfig(tp_size=4)) + self.assertEqual(tp_plan, DENSE_TP_PLAN | EXPERT_TP_PLAN | {"layers.*.self_attn.q_proj": "colwise_rep"}) + self.assertEqual(ep_plan, {}) + self.assertEqual(self.model.ep_plan["layers.*.mlp.experts.down_proj"], "rowwise") + + def test_ep_rules_take_precedence_over_tp_rules_for_the_same_modules(self): + config = DistributedConfig( + tp_size=4, + ep_size=4, + tp_plan={"layers.*.mlp.experts.gate_up_proj": "packed_rowwise", "layers.*.mlp.gate": "colwise"}, + ) + tp_plan, ep_plan = tensor_parallel.resolve_parallel_plans(self.model, config) + self.assertEqual(tp_plan, DENSE_TP_PLAN) + self.assertEqual(ep_plan, EP_PLAN) + # The custom TP rules are kept on the model and apply as soon as EP is disabled. + tp_plan, ep_plan = tensor_parallel.resolve_parallel_plans(self.model, DistributedConfig(tp_size=4)) + self.assertEqual(tp_plan["layers.*.mlp.experts.gate_up_proj"], "packed_rowwise") + self.assertEqual(tp_plan["layers.*.mlp.gate"], "colwise") + self.assertEqual(ep_plan, {}) + + def test_ep_requires_an_expert_plan(self): + self.model.ep_plan = None + with self.assertRaisesRegex(ValueError, "does not define an expert-parallel plan"): + tensor_parallel.resolve_parallel_plans(self.model, DistributedConfig(tp_size=4, ep_size=4)) + config = DistributedConfig(tp_size=4, ep_size=4, ep_plan=EP_PLAN) + self.assertEqual(tensor_parallel.resolve_parallel_plans(self.model, config), (DENSE_TP_PLAN, EP_PLAN)) + + def test_unmatched_override_keys_raise_without_changing_plans(self): + original_tp_plan, original_ep_plan = self.model.tp_plan.copy(), self.model.ep_plan.copy() + for plan_name in ("tp_plan", "ep_plan"): + for key in ("layers.*.mlp.experst", "layers.*.mlp.experts.missing_weight", "model.layers.*.mlp.experts"): + with self.subTest(plan_name=plan_name, key=key): + config = DistributedConfig(tp_size=4, ep_size=4, **{plan_name: {key: "grouped_gemm"}}) + with self.assertRaisesRegex(ValueError, f"The `{plan_name}` pattern .* does not match") as error: + tensor_parallel.resolve_parallel_plans(self.model, config) + self.assertIn(key, str(error.exception)) + self.assertIn("Qwen3MoeModel", str(error.exception)) + self.assertEqual(self.model.tp_plan, original_tp_plan) + self.assertEqual(self.model.ep_plan, original_ep_plan) + + def test_override_keys_can_match_modules_parameters_or_existing_plan_keys(self): + for plan_name in ("tp_plan", "ep_plan"): + # `gate_proj` is in the predefined TP plan even though this MoE model has no such module. + for key in ("layers.*.mlp", "layers.0.self_attn.q_proj.weight", "layers.*.mlp.gate_proj"): + with self.subTest(plan_name=plan_name, key=key): + original = getattr(self.model, plan_name).copy() + if plan_name == "ep_plan": + self.model.ep_plan = original | {"layers.*.mlp.gate_proj": "colwise"} + config = DistributedConfig(tp_size=4, **{plan_name: {key: "colwise_rep"}}) + tensor_parallel.resolve_parallel_plans(self.model, config) + self.assertEqual(getattr(self.model, plan_name)[key], "colwise_rep") + setattr(self.model, plan_name, original) + + def test_head_model_overrides_need_the_model_prefix(self): + with torch.device("meta"): + model = Qwen3MoeForCausalLM(self.config) + config = DistributedConfig(tp_size=4, ep_size=4, ep_plan={"layers.*.mlp.gate": "ep_router"}) + with self.assertRaisesRegex(ValueError, "including any 'model.' prefix"): + tensor_parallel.resolve_parallel_plans(model, config) + + config = DistributedConfig( + tp_size=4, + ep_size=4, + tp_plan={"model.layers.*.self_attn.q_proj": "colwise_rep"}, + ep_plan={"model.layers.*.mlp.gate": "ep_router"}, + ) + tp_plan, ep_plan = tensor_parallel.resolve_parallel_plans(model, config) + expected_tp_plan = {f"model.{k}": v for k, v in DENSE_TP_PLAN.items()} | {"lm_head": "colwise_gather_output"} + self.assertEqual(tp_plan, expected_tp_plan | config.tp_plan) + self.assertEqual(ep_plan, {f"model.{k}": v for k, v in EP_PLAN.items()}) + + def test_masked_ep_shards_and_installs_hooks_on_the_tp_mesh(self): + tp_mesh = object() + _, ep_plan = tensor_parallel.resolve_parallel_plans(self.model, DistributedConfig(tp_size=4, ep_size=4)) + experts, router = self.model.layers[0].mlp.experts, self.model.layers[0].mlp.gate + with ( + patch.object(ALL_PARALLEL_STYLES["grouped_gemm"], "validate_param") as validate, + patch.object(ALL_PARALLEL_STYLES["grouped_gemm"], "shard_param") as shard, + patch.object(ALL_PARALLEL_STYLES["moe_tp_experts"], "install_forward") as install_experts, + patch.object(ALL_PARALLEL_STYLES["ep_router"], "install_forward") as install_router, + ): + result = tensor_parallel.apply_tensor_parallelism(self.model, tp_mesh, ep_plan) + self.assertIs(result, self.model) + self.assertEqual(shard.call_count, 2) + for name in ("gate_up_proj", "down_proj"): + validate.assert_any_call(experts, name, tp_mesh, parameter_name=f"layers.0.mlp.experts.{name}") + shard.assert_any_call(experts, name, tp_mesh) + install_experts.assert_called_once_with(experts, tp_mesh) + install_router.assert_called_once_with(router, tp_mesh) @is_tensor_parallel_test diff --git a/tests/test_modeling_common.py b/tests/test_modeling_common.py index 99e9456487ae..38df5eca748d 100644 --- a/tests/test_modeling_common.py +++ b/tests/test_modeling_common.py @@ -136,7 +136,7 @@ from torch import nn from transformers import MODEL_MAPPING - from transformers.distributed.tensor_parallel import _get_parameter_tp_plan + from transformers.distributed.tensor_parallel import _get_parameter_plan from transformers.integrations.accelerate import compute_module_sizes from transformers.modeling_utils import load_state_dict from transformers.pytorch_utils import id_tensor_storage @@ -4874,10 +4874,9 @@ def test_tp_plan_matches_params(self): for pattern in tp_plan: # Check if this given pattern matches any param or module (the value attributed to the pattern does not matter) pattern_usage[pattern] = any( - _get_parameter_tp_plan(param, {pattern: ""}, is_weight=True) is not None for param in param_names + _get_parameter_plan(param, {pattern: ""}, is_weight=True) is not None for param in param_names ) or any( - _get_parameter_tp_plan(module, {pattern: ""}, is_weight=False) is not None - for module in module_names + _get_parameter_plan(module, {pattern: ""}, is_weight=False) is not None for module in module_names ) unused_entries = {k for k, v in pattern_usage.items() if not v} diff --git a/tests/test_tensor_parallel_mixin.py b/tests/test_tensor_parallel_mixin.py index 38f78a0163f3..c82e1de4671a 100644 --- a/tests/test_tensor_parallel_mixin.py +++ b/tests/test_tensor_parallel_mixin.py @@ -20,7 +20,7 @@ from transformers import TorchAoConfig, set_seed from transformers.distributed.configuration_utils import DistributedConfig -from transformers.distributed.tensor_parallel import _get_parameter_tp_plan +from transformers.distributed.tensor_parallel import _get_parameter_plan from transformers.testing_utils import ( is_tensor_parallel_test, is_torch_available, @@ -182,7 +182,7 @@ def _verify_tp_sharding(rank, model_tp, model_ref): # Verify sharding is correct for dim in range(param.ndim): if param.size(dim) != param_full.size(dim): - param_plan = _get_parameter_tp_plan(name, model_tp._tp_plan, is_weight=True) + param_plan = _get_parameter_plan(name, model_tp._tp_plan, is_weight=True) if param_plan in ("packed_colwise", "packed_rowwise"): expected_size = param_full.size(dim) // world_size assert param.size(dim) == expected_size, ( @@ -268,7 +268,7 @@ def _test_tp_backward_impl(rank, model_path, model_class, atol, rtol): if grad.shape != grad_tp.shape: for dim in range(grad.ndim): if grad.size(dim) != grad_tp.size(dim): - param_plan = _get_parameter_tp_plan(name, model_tp._tp_plan, is_weight=True) + param_plan = _get_parameter_plan(name, model_tp._tp_plan, is_weight=True) if param_plan in ("packed_colwise", "packed_rowwise"): # interleaved slicing grad = get_packed_grad_shard(grad, world_size, rank, dim) @@ -392,6 +392,7 @@ def _test_tp_generation_quantized_impl(_rank, model_path, model_class, max_new_t def _load_ep_and_reference_models(model_path, model_class): """Load EP model and non-EP reference model for comparison.""" + # All-reduce EP: every rank sees the same tokens, so TP and EP span the same ranks. model_ep = model_class.from_pretrained( model_path, distributed_config=DistributedConfig(tp_size=dist.get_world_size(), ep_size=dist.get_world_size()), From 271de167c289a07c0ce1f1e48fe7589cc89f2017 Mon Sep 17 00:00:00 2001 From: 3outeille Date: Tue, 22 Sep 2026 09:22:27 +0000 Subject: [PATCH 17/86] Match expert paths with a regex and keep only expert rules under token dispatch Replace the fnmatch check in resolve_parallel_plans with a plan-pattern regex so plan keys are matched literally except for `*`. When the EP plan uses `ep_dispatch_experts`, drop the router masking rules from the EP plan and only take the expert modules and their parameters out of the TP plan. --- .../distributed/tensor_parallel.py | 24 +++++++++++++++---- 1 file changed, 19 insertions(+), 5 deletions(-) diff --git a/src/transformers/distributed/tensor_parallel.py b/src/transformers/distributed/tensor_parallel.py index 99973c5abf09..e5296f72d126 100644 --- a/src/transformers/distributed/tensor_parallel.py +++ b/src/transformers/distributed/tensor_parallel.py @@ -15,7 +15,6 @@ import contextlib import re -from fnmatch import fnmatchcase from typing import TYPE_CHECKING from ..utils import logging @@ -51,6 +50,14 @@ def replace_layer_number_by_wildcard(name: str) -> str: return re.sub(r"\.\d+(\.|$)", lambda m: ".*" + m.group(1), name) +def _plan_pattern_to_regex(pattern: str) -> str: + """ + Translate a plan key into a regex, where `*` stands for any run of characters (typically a layer index, e.g. + `"model.layers.*.mlp.experts"`). Every other character is matched literally. + """ + return ".*".join(re.escape(part) for part in pattern.split("*")) + + def verify_tp_plan(expected_keys: list[str], tp_plan: dict[str, str] | None): """ Verify the TP plan of the model, log a warning if the layers that were not sharded and the rules that were not applied. @@ -847,10 +854,17 @@ def resolve_parallel_plans( "`base_model_ep_plan` to the model's config, or disable expert parallelism." ) - def is_expert_path(name: str) -> bool: - return any(fnmatchcase(name, path) or fnmatchcase(name, path + ".*") for path in ep_plan) - - tp_plan = {name: style for name, style in tp_plan.items() if not is_expert_path(name)} + def is_expert_path(name: str, paths: list[str]) -> bool: + # An EP path also owns its children, e.g. `...experts.gate_up_proj` under `...experts`. + return any(re.fullmatch(rf"{_plan_pattern_to_regex(path)}(\..*)?", name) for path in paths) + + expert_paths = list(ep_plan) + if "ep_dispatch_experts" in ep_plan.values(): + # Dispatch finds each expert's owner from the global expert ids, so the router masking hooks + # (`ep_router`) must not run: keep only the expert modules and their parameter rules. + expert_paths = [name for name, style in ep_plan.items() if style in ("moe_tp_experts", "ep_dispatch_experts")] + ep_plan = {name: style for name, style in ep_plan.items() if is_expert_path(name, expert_paths)} + tp_plan = {name: style for name, style in tp_plan.items() if not is_expert_path(name, expert_paths)} _validate_parallel_plan_styles(tp_plan) _validate_parallel_plan_styles(ep_plan) return tp_plan, ep_plan From cc5200f940d61dfaf03a70c063df2609316a7fa1 Mon Sep 17 00:00:00 2001 From: 3outeille Date: Wed, 30 Sep 2026 17:12:44 +0000 Subject: [PATCH 18/86] revert change on warning --- .../distributed/configuration_utils.py | 15 +++++++++------ 1 file changed, 9 insertions(+), 6 deletions(-) diff --git a/src/transformers/distributed/configuration_utils.py b/src/transformers/distributed/configuration_utils.py index 91926709897b..77800f12a64d 100644 --- a/src/transformers/distributed/configuration_utils.py +++ b/src/transformers/distributed/configuration_utils.py @@ -18,6 +18,8 @@ from dataclasses import asdict, dataclass from typing import Literal +from .utils import _get_torch_distributed_rank + @dataclass class DistributedConfig: @@ -98,12 +100,13 @@ def _resolve_parallelism(self): if self.enable_expert_parallel and self.ep_size is None: self.ep_size = self.tp_size - warnings.warn( - f"`enable_expert_parallel` without `ep_size` is deprecated and will be removed in v5.20. " - f"Use ep_size={self.ep_size} instead.", - FutureWarning, - stacklevel=4, - ) + if _get_torch_distributed_rank() == 0: + warnings.warn( + f"`enable_expert_parallel` without `ep_size` is deprecated and will be removed in v5.20. " + f"Use ep_size={self.ep_size} instead.", + FutureWarning, + stacklevel=4, + ) if self.ep_size is None: self.ep_size = 1 From 17f7b659784c40a64f9d78475eb6dc98ca4b0b4b Mon Sep 17 00:00:00 2001 From: Ferdinand Mom <47445085+3outeille@users.noreply.github.com> Date: Thu, 1 Oct 2026 05:34:19 +0900 Subject: [PATCH 19/86] Apply suggestion from @ArthurZucker Co-authored-by: Arthur <48595927+ArthurZucker@users.noreply.github.com> --- .../distributed/tensor_parallel.py | 20 +++++-------------- 1 file changed, 5 insertions(+), 15 deletions(-) diff --git a/src/transformers/distributed/tensor_parallel.py b/src/transformers/distributed/tensor_parallel.py index e5296f72d126..2a57300287c6 100644 --- a/src/transformers/distributed/tensor_parallel.py +++ b/src/transformers/distributed/tensor_parallel.py @@ -826,24 +826,14 @@ def resolve_parallel_plans( are dropped: expert weights are sharded once, by the EP plan. """ # Reject invalid paths before merging, e.g. "layers.*" when the model uses "model.layers.*". - layer_names = {name for name, _ in model.named_modules()} | {name for name, _ in model.named_parameters()} - layer_names |= {replace_layer_number_by_wildcard(name) for name in layer_names} + names = {replace_layer_number_by_wildcard(n) for n, _ in chain(model.named_modules(), model.named_parameters())} for plan_name in ("tp_plan", "ep_plan"): override = getattr(distributed_config, plan_name) if isinstance(override, dict): - valid_names = layer_names | set(getattr(model, plan_name)) - for pattern in override: - if pattern not in valid_names: - raise ValueError( - f"The `{plan_name}` pattern {pattern!r} does not match any module, parameter, " - f"or existing plan entry in {type(model).__name__}. " - "Check the full path, including any 'model.' prefix." - ) - - if isinstance(distributed_config.tp_plan, dict): - model._tp_plan = model.tp_plan | distributed_config.tp_plan - if isinstance(distributed_config.ep_plan, dict): - model._ep_plan = model.ep_plan | distributed_config.ep_plan + plan = getattr(model, plan_name) + if unknown := override.keys() - names - plan.keys(): + raise ValueError(f"`{plan_name}` keys {sorted(unknown)} match nothing in {type(model).__name__}.") + setattr(model, f"_{plan_name}", plan | override) tp_plan = dict(model.tp_plan) if distributed_config.tp_size > 1 else {} ep_plan = dict(model.ep_plan) if distributed_config.ep_size > 1 else {} From ef837266c0740d67c507d77234bc047b42d662f5 Mon Sep 17 00:00:00 2001 From: 3outeille Date: Wed, 30 Sep 2026 21:21:54 +0000 Subject: [PATCH 20/86] cleaning --- .../distributed/tensor_parallel.py | 27 ++++++------------- tests/tensor_parallel/test_tensor_parallel.py | 6 ++--- 2 files changed, 11 insertions(+), 22 deletions(-) diff --git a/src/transformers/distributed/tensor_parallel.py b/src/transformers/distributed/tensor_parallel.py index 2a57300287c6..5864520fac99 100644 --- a/src/transformers/distributed/tensor_parallel.py +++ b/src/transformers/distributed/tensor_parallel.py @@ -15,6 +15,7 @@ import contextlib import re +from itertools import chain from typing import TYPE_CHECKING from ..utils import logging @@ -50,14 +51,6 @@ def replace_layer_number_by_wildcard(name: str) -> str: return re.sub(r"\.\d+(\.|$)", lambda m: ".*" + m.group(1), name) -def _plan_pattern_to_regex(pattern: str) -> str: - """ - Translate a plan key into a regex, where `*` stands for any run of characters (typically a layer index, e.g. - `"model.layers.*.mlp.experts"`). Every other character is matched literally. - """ - return ".*".join(re.escape(part) for part in pattern.split("*")) - - def verify_tp_plan(expected_keys: list[str], tp_plan: dict[str, str] | None): """ Verify the TP plan of the model, log a warning if the layers that were not sharded and the rules that were not applied. @@ -844,17 +837,13 @@ def resolve_parallel_plans( "`base_model_ep_plan` to the model's config, or disable expert parallelism." ) - def is_expert_path(name: str, paths: list[str]) -> bool: - # An EP path also owns its children, e.g. `...experts.gate_up_proj` under `...experts`. - return any(re.fullmatch(rf"{_plan_pattern_to_regex(path)}(\..*)?", name) for path in paths) - - expert_paths = list(ep_plan) - if "ep_dispatch_experts" in ep_plan.values(): - # Dispatch finds each expert's owner from the global expert ids, so the router masking hooks - # (`ep_router`) must not run: keep only the expert modules and their parameter rules. - expert_paths = [name for name, style in ep_plan.items() if style in ("moe_tp_experts", "ep_dispatch_experts")] - ep_plan = {name: style for name, style in ep_plan.items() if is_expert_path(name, expert_paths)} - tp_plan = {name: style for name, style in tp_plan.items() if not is_expert_path(name, expert_paths)} + if "ep_dispatch_experts" in ep_plan.values() and "ep_router" in ep_plan.values(): + raise ValueError("`ep_dispatch_experts` routes tokens itself; remove the `ep_router` rules from `ep_plan`.") + + # EP rules take precedence: drop TP rules on EP modules and their children. + is_expert = re.compile(rf"(?:{'|'.join(map(re.escape, ep_plan))})(?:\..+)?").fullmatch + tp_plan = {name: style for name, style in tp_plan.items() if not is_expert(name)} + _validate_parallel_plan_styles(tp_plan) _validate_parallel_plan_styles(ep_plan) return tp_plan, ep_plan diff --git a/tests/tensor_parallel/test_tensor_parallel.py b/tests/tensor_parallel/test_tensor_parallel.py index c39930ef2307..a265cbd4087f 100644 --- a/tests/tensor_parallel/test_tensor_parallel.py +++ b/tests/tensor_parallel/test_tensor_parallel.py @@ -181,7 +181,7 @@ def test_unmatched_override_keys_raise_without_changing_plans(self): for key in ("layers.*.mlp.experst", "layers.*.mlp.experts.missing_weight", "model.layers.*.mlp.experts"): with self.subTest(plan_name=plan_name, key=key): config = DistributedConfig(tp_size=4, ep_size=4, **{plan_name: {key: "grouped_gemm"}}) - with self.assertRaisesRegex(ValueError, f"The `{plan_name}` pattern .* does not match") as error: + with self.assertRaisesRegex(ValueError, f"`{plan_name}` keys .* match nothing in") as error: tensor_parallel.resolve_parallel_plans(self.model, config) self.assertIn(key, str(error.exception)) self.assertIn("Qwen3MoeModel", str(error.exception)) @@ -191,7 +191,7 @@ def test_unmatched_override_keys_raise_without_changing_plans(self): def test_override_keys_can_match_modules_parameters_or_existing_plan_keys(self): for plan_name in ("tp_plan", "ep_plan"): # `gate_proj` is in the predefined TP plan even though this MoE model has no such module. - for key in ("layers.*.mlp", "layers.0.self_attn.q_proj.weight", "layers.*.mlp.gate_proj"): + for key in ("layers.*.mlp", "layers.*.self_attn.q_proj.weight", "layers.*.mlp.gate_proj"): with self.subTest(plan_name=plan_name, key=key): original = getattr(self.model, plan_name).copy() if plan_name == "ep_plan": @@ -205,7 +205,7 @@ def test_head_model_overrides_need_the_model_prefix(self): with torch.device("meta"): model = Qwen3MoeForCausalLM(self.config) config = DistributedConfig(tp_size=4, ep_size=4, ep_plan={"layers.*.mlp.gate": "ep_router"}) - with self.assertRaisesRegex(ValueError, "including any 'model.' prefix"): + with self.assertRaisesRegex(ValueError, "match nothing in Qwen3MoeForCausalLM"): tensor_parallel.resolve_parallel_plans(model, config) config = DistributedConfig( From e7bfac9b3a951bd6c6357a83f7c755dfc411f177 Mon Sep 17 00:00:00 2001 From: 3outeille Date: Fri, 18 Sep 2026 16:59:48 +0000 Subject: [PATCH 21/86] Fix tied embeddings for models with an EP-only base plan --- src/transformers/configuration_utils.py | 6 ++++-- 1 file changed, 4 insertions(+), 2 deletions(-) diff --git a/src/transformers/configuration_utils.py b/src/transformers/configuration_utils.py index 23534549c957..815a5f794103 100755 --- a/src/transformers/configuration_utils.py +++ b/src/transformers/configuration_utils.py @@ -388,9 +388,11 @@ def __post_init__(self, **kwargs): self.per_layer_config = per_layer_config # TODO: to support models whose input embedding module is not named `embed_tokens` (e.g. GPT-NeoX's `embed_in`). - if getattr(self, "tie_word_embeddings", False) and self.base_model_tp_plan is not None: + if getattr(self, "tie_word_embeddings", False) and ( + self.base_model_tp_plan is not None or self.base_model_ep_plan is not None + ): self.base_model_tp_plan = { - **self.base_model_tp_plan, + **(self.base_model_tp_plan or {}), "embed_tokens": "embedding_rowwise", } From 4cc84e4248763ed08088f924ff6b5bd16d16ba12 Mon Sep 17 00:00:00 2001 From: 3outeille Date: Fri, 18 Sep 2026 16:59:48 +0000 Subject: [PATCH 22/86] [distributed] Add expert-parallel token dispatch, default for Qwen3 MoE --- docs/source/en/expert_parallelism.md | 108 +++++++++- .../distributed/configuration_utils.py | 24 ++- src/transformers/distributed/fsdp.py | 23 +- src/transformers/distributed/mixin.py | 33 ++- .../distributed/tensor_parallel.py | 199 +++++++++++++++++- .../qwen3_moe/configuration_qwen3_moe.py | 10 +- tests/tensor_parallel/test_tensor_parallel.py | 43 ++-- 7 files changed, 390 insertions(+), 50 deletions(-) diff --git a/docs/source/en/expert_parallelism.md b/docs/source/en/expert_parallelism.md index 785c7f7be2f2..61d646c56f23 100644 --- a/docs/source/en/expert_parallelism.md +++ b/docs/source/en/expert_parallelism.md @@ -19,7 +19,7 @@ rendered properly in your Markdown viewer. ## DistributedConfig -Enable expert parallelism with the [`DistributedConfig`] class and the `ep_size` argument. The current all-reduce implementation requires `ep_size=tp_size`, so every rank in an expert group receives the same tokens. +Enable expert parallelism with the [`DistributedConfig`] class and the `ep_size` argument. Most models route with masking and all-reduce, which requires `ep_size=tp_size` so every rank in an expert group receives the same tokens. Models whose plan uses [token dispatch](#token-dispatch), such as Qwen3 MoE, can set `ep_size` independently of `tp_size`. ```py import os @@ -43,10 +43,17 @@ Each MoE model defines two plans in its config: `base_model_tp_plan` for the den `tp_plan` is applied only when `tp_size > 1`, and `ep_plan` only when `ep_size > 1`. With TP enabled and EP disabled, the full TP plan applies, expert rules included. +The expert forward rule in `ep_plan` selects how tokens reach the experts: + +| rule | mechanism | layout | +| :--- | :--- | :--- | +| `"moe_tp_experts"` with `"ep_router"` on the router | masking and all-reduce: every rank runs its local experts on the whole batch, the router masks the others, and an all-reduce combines the outputs | `ep_size=tp_size` | +| `"ep_dispatch_experts"` | [token dispatch](#token-dispatch): each rank keeps its own tokens and only exchanges the routed (token, expert) pairs with two all-to-all collectives | `ep_size` a multiple of `tp_size` that divides `fsdp_size * tp_size` | + > [!TIP] > `enable_expert_parallel=True` is a deprecated alias for `ep_size=tp_size`, used only when `ep_size` is omitted, and emits a `FutureWarning`. -Launch your inference script with [torchrun](https://pytorch.org/docs/stable/elastic/run.html) and specify how many devices to use. The number of devices must evenly divide the total number of experts. +Launch your inference script with [torchrun](https://pytorch.org/docs/stable/elastic/run.html). The number of processes must equal `tp_size * fsdp_size * pp_size`, and `ep_size` must evenly divide the number of experts. ```zsh torchrun --nproc-per-node 8 your_script.py @@ -67,9 +74,98 @@ distributed_config = DistributedConfig( Providing a plan does not infer parallel sizes: set `tp_size` and `ep_size` explicitly. +Qwen3 MoE defaults to `"ep_dispatch_experts"`. To use masking and all-reduce instead, set `ep_size=tp_size` and override both the router and the expert forward rules: + +```py +distributed_config = DistributedConfig( + tp_size=4, + ep_size=4, + ep_plan={ + "model.layers.*.mlp.gate": "ep_router", + "model.layers.*.mlp.experts": "moe_tp_experts", + }, +) +``` + +Conversely, override the expert forward rule of a model whose plan uses masking with `"ep_dispatch_experts"` to use token dispatch. The router rule is then ignored, since dispatch needs the global expert ids to find each expert's owner. + +## Token dispatch + +With token dispatch, each rank trains on its own part of the batch. At every MoE layer, a rank routes its tokens, sends each (token, expert) pair to the rank that owns the expert with an all-to-all, runs its local experts on what it receives, gets the results back with a second all-to-all and combines them with the routing weights. Only the routed activations and expert outputs travel, and no rank computes experts for tokens it does not own. + +```py +from transformers import AutoModelForCausalLM +from transformers.distributed import DistributedConfig + +distributed_config = DistributedConfig( + tp_size=1, + fsdp_size=8, + ep_size=4, +) +model = AutoModelForCausalLM.from_pretrained("Qwen/Qwen3-30B-A3B", distributed_config=distributed_config) +``` + +With `tp_size=1`, `ep_size` must divide `fsdp_size` and the number of experts, and attention is not bound by `num_key_value_heads`. For the rest of the model: + +- The parameters outside the experts are sharded with [FSDP2](./fsdp) across `fsdp`, which reduces their gradients. +- The experts are sharded across `ep` and, when `efsdp_size = fsdp_size * tp_size // ep_size` is larger than one, additionally FSDP-sharded across `efsdp`. They are always FSDP-wrapped, so `fsdp_mixed_precision` and `fsdp_cpu_offload` apply to them too and [`~PreTrainedModel.save_pretrained`] gathers them like any other parameter. +- An expert parallel group holds `ep_size / tp_size` batches, so an expert's gradient is a sum over that many batches. The `efsdp` reduction divides by `fsdp_size` instead of its group size, which gives the same per-batch average FSDP2 takes for the dense modules. +- Every local expert also processes one zero pad row per layer. A rank whose experts received no tokens still joins the reverse all-to-all and the expert gradient reduction. + +### How the sizes combine + +- `tp_size * fsdp_size` is the number of processes. `ep_size` adds none: it regroups the same ranks for the expert weights only. +- `ep_size` cuts the expert list into `ep_size` blocks. Each rank computes `num_experts / ep_size` experts, and `ep_size` consecutive ranks hold one complete set. That set of ranks is the group the all-to-all runs in. +- `efsdp_size = fsdp_size * tp_size / ep_size` is how many complete copies of the expert set exist. Ranks at the same position in different copies shard those experts for memory and average their gradients, like FSDP does for the dense modules. +- The batch a rank holds depends on `fsdp` only. Consecutive ranks form a TP group and get the same batch; the `fsdp_size` groups get different batches. + +Eight processes, `tp_size=2, fsdp_size=4`, eight experts: + +```text +rank 0 1 2 3 4 5 6 7 +batch [====B0====] [====B1====] [====B2====] [====B3====] one batch per TP pair +tp 0 1 0 1 0 1 0 1 + +ep_size=2 E0-3 E4-7 E0-3 E4-7 E0-3 E4-7 E0-3 E4-7 group = a TP pair, efsdp_size=4 +ep_size=4 E0E1 E2E3 E4E5 E6E7 E0E1 E2E3 E4E5 E6E7 group = two pairs, efsdp_size=2 +ep_size=8 E0 E1 E2 E3 E4 E5 E6 E7 group = all ranks, efsdp_size=1 +``` + +Two numbers follow from the picture: + +- Inside an EP group, each token exists `tp_size` times, once per rank of the pair that holds its batch. This does not depend on `ep_size`. +- An EP group holds `ep_size / tp_size` different batches. This is the count an expert's gradient sums over. + +### With tensor parallelism + +Set `tp_size > 1` to shard the dense modules with the TP plan while the experts use dispatch. On eight processes: + +```py +distributed_config = DistributedConfig( + tp_size=2, + fsdp_size=4, + ep_size=4, +) +``` + +Each pair of TP ranks receives the same batch, because tensor parallelism replicates the activations inside the pair. If both ranks dispatched all of their tokens, the rank owning an expert would receive every token twice, compute it twice, and its weight gradient would double. Experts are whole on one rank, so the duplicate cannot be split by weights, and the owner is usually another rank, so it cannot be resolved by ownership as masking does. The pair therefore splits the rows: each TP rank dispatches a disjoint `1 / tp_size` of the tokens, results come back to the rank that sent them, and an all-reduce over the pair of the zero-padded halves restores the replicated output the next layer expects. The split is by `tp_size`, not `ep_size`, since only the ranks that hold a batch can send it. Expert groups span four ranks and each expert is FSDP-sharded across `efsdp_size = 2` ranks, while the trunk's FSDP group spans four ranks. The model's usual TP constraints, such as attention-head divisibility, still apply to the dense modules. Token slices may be uneven or empty, including during single-token decoding. + +Token dispatch cannot be combined yet with pipeline parallelism yet; use `pp_size=1` (not tested yet) +These configurations each use eight GPUs: + +| Configuration | Result | +| :--- | :--- | +| `DistributedConfig(tp_size=4, fsdp_size=2, ep_size=4)` | Dispatch with TP groups of four, each slicing its batch in four (Qwen3 MoE default plan); unless you specify a ep_plan to use the legacy masked EP | +| `DistributedConfig(tp_size=1, fsdp_size=8, ep_size=4)` | Dispatch with an independent batch on each rank, no slicing, and experts FSDP-sharded across pairs of ranks. | +| `DistributedConfig(tp_size=2, fsdp_size=4, ep_size=4)` | Dispatch with a TP pair per batch, each pair slicing its batch in two; two batches per expert group. | +| `DistributedConfig(tp_size=8, ep_size=8)` | Dispatch with every rank sharing one batch, sliced in eight, or masking and all-reduce for a masked plan. | + +> [!WARNING] +> The [`Trainer`] does not account for token dispatch yet: batch and token counting assume the all-reduce layout, where the ranks of a TP group share a batch and `fsdp_size` data-parallel shards exist. Trainer support for dispatch comes in a follow-up. + ## Combining with FSDP2 -Tensor and expert parallelism shard the weights across `tp`, but the optimizer state and the modules without a rule are still replicated on every rank of the group, which limits how large a model you can train. Add [FSDP2](./fsdp) on a second mesh dimension with `fsdp_size`, and keep `ep_size=tp_size` for the expert parallel width. +Tensor and expert parallelism shard the weights across `tp`, but the optimizer state and the modules without a rule are still replicated on every rank of the group, which limits how large a model you can train. Add [FSDP2](./fsdp) on a second mesh dimension with `fsdp_size`. With masking and all-reduce, keep `ep_size=tp_size` and pass `ep_plan={"layers.*.mlp.gate": "ep_router", "layers.*.mlp.experts": "moe_tp_experts"}``. ```py from transformers import AutoModelForCausalLM @@ -77,8 +173,12 @@ from transformers.distributed import DistributedConfig distributed_config = DistributedConfig( tp_size=4, - ep_size=4, # expert parallel size, must match tp_size + ep_size=4, # expert parallel size, must match tp_size with masking and all-reduce fsdp_size=2, # data parallel shards + ep_plan={ + "model.layers.*.mlp.gate": "ep_router", + "model.layers.*.mlp.experts": "moe_tp_experts", + }, ) model = AutoModelForCausalLM.from_pretrained("Qwen/Qwen3-30B-A3B", distributed_config=distributed_config) ``` diff --git a/src/transformers/distributed/configuration_utils.py b/src/transformers/distributed/configuration_utils.py index 77800f12a64d..a2b3c14248ac 100644 --- a/src/transformers/distributed/configuration_utils.py +++ b/src/transformers/distributed/configuration_utils.py @@ -18,6 +18,7 @@ from dataclasses import asdict, dataclass from typing import Literal +from ..utils import is_torch_greater_or_equal from .utils import _get_torch_distributed_rank @@ -48,12 +49,14 @@ class DistributedConfig: pp_size (`int`, *optional*): Number of devices for pipeline parallelism. If `None` and another parallel mode is set, defaults to 1. ep_size (`int`, *optional*): - Number of devices owning distinct expert shards. Defaults to 1. Set it explicitly to enable EP. - Model execution currently requires `ep_size=tp_size` when EP is enabled. + Number of devices owning distinct expert shards. Defaults to 1. Set it explicitly to enable EP. Must be + a multiple of `tp_size` and divide `fsdp_size * tp_size`. All-reduce expert plans require + `ep_size=tp_size`; token dispatch (`"ep_dispatch_experts"`) also allows `ep_size > tp_size`. ep_plan (`dict[str, str]`, *optional*): Expert parallel sharding plan. Leave as `None` to use the model's predefined `base_model_ep_plan`. Pass a dictionary to override individual rules of that plan; unspecified rules are kept. Applied only when - `ep_size > 1`, and its rules take precedence over `tp_plan` rules for the same modules. + `ep_size > 1`, and its rules take precedence over `tp_plan` rules for the same modules. An + `"ep_dispatch_experts"` rule selects all-to-all token dispatch instead of router masking and all-reduce. """ tp_size: int | None = None @@ -130,6 +133,21 @@ def _validate_mesh_config(self): "Use DistributedConfig(tp_size=N, fsdp_size=M), or combine TP and PP." ) + def _validate_resolved_ep_plan(self, ep_plan: dict[str, str]): + """Validate the layout against the resolved EP plan, once the model's defaults and overrides are merged.""" + if self.ep_size <= 1 or not ep_plan: + return + + if "ep_dispatch_experts" in ep_plan.values(): + if self.pp_size > 1: + raise ValueError("Combining token dispatch with pipeline parallelism is not supported/tested yet.") + if not is_torch_greater_or_equal("2.7"): + raise OSError("Expert-parallel token dispatch requires `torch>=2.7`.") + elif {"ep_router", "moe_tp_experts"}.issubset(ep_plan.values()) and self.ep_size != self.tp_size: + raise ValueError( + "All-reduce expert parallelism requires `ep_size=tp_size`, so every rank of an expert group sees the same tokens" + ) + @classmethod def from_dict(cls, config_dict: dict, **kwargs) -> "DistributedConfig": merged = {**config_dict, **kwargs} diff --git a/src/transformers/distributed/fsdp.py b/src/transformers/distributed/fsdp.py index aef3a5979a8c..d651ae61f398 100644 --- a/src/transformers/distributed/fsdp.py +++ b/src/transformers/distributed/fsdp.py @@ -27,6 +27,7 @@ import torch.nn as nn from .configuration_utils import DistributedConfig + from .utils import MeshManager if is_torch_available(): import torch @@ -184,15 +185,15 @@ def verify_fsdp_plan(module_names: list[str], fsdp_plan: dict[str, str] | None) logger.warning(f"The following FSDP rules were not applied to any module: {unused_rules}") -def apply_fully_sharded_data_parallelism( - model: nn.Module, fsdp_mesh: torch.distributed.device_mesh.DeviceMesh -) -> nn.Module: +def apply_fully_sharded_data_parallelism(model: nn.Module, mesh_manager: MeshManager) -> nn.Module: """ - Apply FSDP2 (fully_shard) to a model. + Apply FSDP2 (fully_shard) to a model: dispatched experts on `efsdp` and the rest of the modules on `fsdp`. Torch availability, distributed initialization and the version requirement are asserted upstream by `initialize_distributed_mesh`. """ + distributed_config = model.config.distributed_config + fsdp_mesh = mesh_manager.get_mesh("fsdp") fsdp_plan = dict(getattr(model, "_fsdp_plan", None) or {}) if not fsdp_plan: raise ValueError( @@ -206,6 +207,20 @@ def apply_fully_sharded_data_parallelism( adapted_fsdp_plan = _resolve_tied_embed_lm_head_plan(fsdp_plan, model) reshard_targets, no_reshard_targets = expand_fsdp_plan(model, adapted_fsdp_plan) + fsdp_policy_kwargs = _get_fsdp_policy_kwargs(distributed_config) + if distributed_config.ep_size > 1 and "ep_dispatch_experts" in model.ep_plan.values(): + expert_mesh = mesh_manager.get_mesh("efsdp") + for module in model.modules(): + if getattr(module, "_is_expert_parallel", False): + fully_shard(module, mesh=expert_mesh, reshard_after_forward=True, **fsdp_policy_kwargs) + # An expert group spans several data-parallel batches, so an expert's gradient sums over + # all of them. FSDP2 would divide by the efsdp group size; dividing by fsdp_size instead + # gives the same per-batch average the dense modules get on the fsdp mesh, even when efsdp has a single rank. + module.set_gradient_divide_factor(float(distributed_config.fsdp_size)) + if torch.distributed.get_backend(expert_mesh.get_group()) != "nccl": + # Non-NCCL backends need to sum first, then apply the division otherwise it runtime error. + module.set_force_sum_reduction_for_comms(True) + for module_name, module in reshard_targets: fully_shard(module, mesh=fsdp_mesh, reshard_after_forward=True, **fsdp_policy_kwargs) logger.debug(f"Applied fully_shard to {module_name} (reshard=True)") diff --git a/src/transformers/distributed/mixin.py b/src/transformers/distributed/mixin.py index 4ae276547185..584e2e5a6fb4 100644 --- a/src/transformers/distributed/mixin.py +++ b/src/transformers/distributed/mixin.py @@ -25,6 +25,7 @@ from .pipeline_parallel import apply_pipeline_parallelism from .tensor_parallel import ( _validate_parallel_plan_styles, + apply_expert_parallelism, apply_tensor_parallelism, gather_state_dict_for_save, resolve_parallel_plans, @@ -197,19 +198,29 @@ def maybe_distribute_model( # Resolve both plans before sharding anything: overrides are merged into `model.tp_plan` / `model.ep_plan`, # and the experts named by the EP plan are removed from the TP plan so they are sharded once. tp_plan, ep_plan = resolve_parallel_plans(model, distributed_config) + distributed_config._validate_resolved_ep_plan(ep_plan) if distributed_config.pp_size > 1: model = apply_pipeline_parallelism(model, mesh_manager.get_mesh("pp")) - tp_mesh = mesh_manager.get_mesh("tp") if tp_plan: - model = apply_tensor_parallelism(model, tp_mesh, tp_plan) + model = apply_tensor_parallelism(model, mesh_manager.get_mesh("tp"), tp_plan) + if ep_plan: - # Legacy masked EP: the EP group is the TP group, every rank keeps every token. - model = apply_tensor_parallelism(model, tp_mesh, ep_plan) + tp_mesh = mesh_manager.get_mesh("tp") + ep_mesh = mesh_manager.get_mesh("ep") + + if {"ep_router", "moe_tp_experts"}.issubset(ep_plan.values()): + # Legacy masked EP: the EP group is the TP group, every rank keeps every token. + model = apply_tensor_parallelism(model, tp_mesh, ep_plan) + elif "ep_dispatch_experts" in ep_plan.values(): + # EP + DP with tp_size >= 1: the ranks of a TP group share the same batch. If we want a specific token, + # we will have to slice the batch here in order to avoid computing tp_size times the same batch. + model = apply_expert_parallelism(model, ep_mesh, tp_mesh, ep_plan) + + if distributed_config.fsdp_size > 1 or "ep_dispatch_experts" in ep_plan.values(): + model = apply_fully_sharded_data_parallelism(model, mesh_manager) - if distributed_config.fsdp_size > 1: - model = apply_fully_sharded_data_parallelism(model, mesh_manager.get_mesh("fsdp")) return model def should_save_on_this_rank(self, is_main_process: bool) -> bool: @@ -270,9 +281,11 @@ def gather_sharded_state_dict_for_save( if distributed_config is None: return state_dict - if distributed_config.fsdp_size > 1: - # Also covers the 2-D (fsdp, tp) mesh: every parameter is FSDP-managed, and the full - # state dict is only materialized on rank 0. + if distributed_config.fsdp_size > 1 or ( + distributed_config.ep_size > 1 and "ep_dispatch_experts" in self.ep_plan.values() + ): + # Also covers the 2-D (fsdp, tp) mesh and token dispatch: every parameter is FSDP-managed, and the + # full state dict is only materialized on rank 0. if not _is_torch_distributed_initialized(): raise ValueError( "Saving an FSDP-wrapped model requires torch.distributed to be initialized. " @@ -294,5 +307,5 @@ def barrier_after_gathered_checkpoint_save(self, distributed_config: Distributed """Barrier so non-writer ranks wait for rank 0 to finish gathered checkpoint writes.""" if distributed_config is None: return - if distributed_config.tp_size > 1 or distributed_config.fsdp_size > 1: + if distributed_config.tp_size > 1 or distributed_config.fsdp_size > 1 or distributed_config.ep_size > 1: _distributed_barrier() diff --git a/src/transformers/distributed/tensor_parallel.py b/src/transformers/distributed/tensor_parallel.py index 5864520fac99..d6c14907efe6 100644 --- a/src/transformers/distributed/tensor_parallel.py +++ b/src/transformers/distributed/tensor_parallel.py @@ -15,6 +15,7 @@ import contextlib import re +from collections.abc import Callable from itertools import chain from typing import TYPE_CHECKING @@ -37,6 +38,7 @@ if is_torch_distributed_available(): import torch.distributed as dist + from torch.distributed.nn.functional import all_to_all_single from torch.distributed.tensor import DTensor, Partial, Replicate, Shard, distribute_tensor from torch.distributed.tensor.placement_types import _StridedShard @@ -748,6 +750,165 @@ def transform_output_post_forward(self, module, output, mesh): return output +class EpDispatchExpertsParallel(MoeExpertsParallel): + """ + Dispatch disjoint TP token slices to the experts' owners, then replicate the combined output on TP. + + Example: + Let's say we have 8 experts [E0, E7] with DistributedConfig(tp_size=2, fsdp_size=4, ep_size=4). That imply: + - Since fsdp_size=4, we have 4 batches B denoted [B0, B3] + - Because we have tp_size=2, that means *both ranks share the same batch* + - efsdp = (fsdp_size * tp_size) / ep_size = 4 * 2 / 4 = 2 + + GPU 0 1 2 3 4 5 6 7 + | | | | | | | | + dense view --------------------------------------------------------- + batch [====B0====] [====B1====] [====B2====] [====B3====] + tp_size [___________ 0 ___________] [___________ 1 ___________] + fsdp_size 0 1 2 3 + + expert view --------------------------------------------------------- + + experts E0E1 E2E3 E4E5 E6E7 E0E1 E2E3 E4E5 E6E7 + ep_size 0 1 2 3 0 1 2 3 + efsdp_size [___________ 0 ___________] [___________ 1 ___________] + + Assume token 1 in B0 chose E4, which lives on another rank, so it must travel by all-to-all. B0 sits on 2 ranks (cf diagram), so if both sent it, E4 would compute it twice. + We need to make sure that token 1 is in rank 0 range, so rank 0 sends it and rank 1 does not have it in its slice. + """ + + def _dispatch_tokens( + self, + hidden_states: torch.Tensor, + top_k_index: torch.Tensor, + num_local_experts: int, + ep_group, + ep_size: int, + ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, list[int], list[int]]: + """Send each selected (token, expert) pair to the rank that owns the expert. + + Also returns the sort order and the per-rank split sizes that `_combine_tokens` needs to reverse the exchange. + """ + hidden_dim = hidden_states.size(-1) + num_top_k = top_k_index.size(-1) + + # Sorting the selected pairs by expert groups them by owner rank, since each rank owns a contiguous range of + # experts, and the per-expert counts tell every receiver which expert each token it gets is for. The split + # sizes are the one host sync of the layer. + expert_ids = top_k_index.reshape(-1) + order = torch.argsort(expert_ids) + send_tokens = hidden_states[order // num_top_k] + send_counts = torch.zeros(num_local_experts * ep_size, dtype=torch.long, device=hidden_states.device) + send_counts = send_counts.scatter_add_(0, expert_ids, torch.ones_like(expert_ids)).view( + ep_size, num_local_experts + ) + recv_counts = torch.empty_like(send_counts) + torch.distributed.all_to_all_single(recv_counts, send_counts, group=ep_group) + send_sizes, recv_sizes = torch.stack([send_counts.sum(dim=1), recv_counts.sum(dim=1)]).tolist() + recv_tokens = all_to_all_single( + send_tokens.new_empty(sum(recv_sizes), hidden_dim), + send_tokens, + output_split_sizes=recv_sizes, + input_split_sizes=send_sizes, + group=ep_group, + ) + recv_expert_ids = torch.arange(num_local_experts, device=hidden_states.device).repeat(ep_size) + recv_expert_ids = recv_expert_ids.repeat_interleave(recv_counts.reshape(-1), output_size=sum(recv_sizes)) + return recv_tokens, recv_expert_ids, order, send_sizes, recv_sizes + + def _run_local_experts( + self, + experts_forward: Callable, + tokens: torch.Tensor, + expert_ids: torch.Tensor, + num_local_experts: int, + ) -> torch.Tensor: + """Run local experts with top-1 routing and unit weights; apply routing weights after combine.""" + # One zero row per expert keeps tokens and all expert weights connected to backward. Without it, + # empty eager experts can skip the reverse all-to-all and FSDP reduction, leaving other ranks waiting. + num_tokens, hidden_dim = tokens.shape + local_expert_ids = torch.arange(num_local_experts, device=tokens.device) + tokens = torch.cat([tokens, tokens.new_zeros(num_local_experts, hidden_dim)]) + expert_ids = torch.cat([expert_ids, local_expert_ids]).unsqueeze(-1) + weights = torch.ones_like(expert_ids, dtype=tokens.dtype) + return experts_forward(tokens, expert_ids, weights)[:num_tokens] + + def _combine_tokens( + self, + expert_output: torch.Tensor, + top_k_weights: torch.Tensor, + order: torch.Tensor, + send_sizes: list[int], + recv_sizes: list[int], + ep_group, + ) -> torch.Tensor: + """Return expert outputs to the token owners and combine them with routing weights.""" + num_tokens, num_top_k = top_k_weights.shape + hidden_dim = expert_output.size(-1) + recv_out = all_to_all_single( + expert_output.new_empty(order.numel(), hidden_dim), + expert_output, + output_split_sizes=send_sizes, + input_split_sizes=recv_sizes, + group=ep_group, + ) + # Restore the original (token, top-k slot) order, then apply routing weights. + token_outputs = torch.empty_like(recv_out) + token_outputs[order] = recv_out + token_outputs = token_outputs.view(num_tokens, num_top_k, hidden_dim) + return (token_outputs * top_k_weights.unsqueeze(-1)).sum(dim=1) + + def transform_inputs_pre_forward(self, module, args, kwargs, mesh, *, tp_mesh=None): + hidden_states, top_k_index, top_k_weights = args + if isinstance(hidden_states, DTensor): + hidden_states = hidden_states.to_local() + if isinstance(top_k_weights, DTensor): + top_k_weights = top_k_weights.to_local() + if tp_mesh is None or tp_mesh.size() == 1: + return (hidden_states, top_k_index, top_k_weights), kwargs + + tp_group = tp_mesh.get_group() + hidden_states = _AllReduceBackward.apply(hidden_states, tp_group) + top_k_weights = _AllReduceBackward.apply(top_k_weights, tp_group) + # TP ranks share the same batch. Slice tokens here so each is dispatched only once, then restore the full output in the post hook + num_tokens, tp_rank, tp_size = hidden_states.size(0), tp_mesh.get_local_rank(), tp_mesh.size() + rows = slice(num_tokens * tp_rank // tp_size, num_tokens * (tp_rank + 1) // tp_size) + return (hidden_states[rows], top_k_index[rows], top_k_weights[rows]), kwargs + + def transform_output_post_forward(self, module, output, mesh, *, tp_mesh=None, num_tokens=None): + if tp_mesh is None or tp_mesh.size() == 1: + return output + tp_rank, tp_size = tp_mesh.get_local_rank(), tp_mesh.size() + rows = slice(num_tokens * tp_rank // tp_size, num_tokens * (tp_rank + 1) // tp_size) + full_output = output.new_zeros(num_tokens, output.size(-1)) + full_output[rows] = output + return _AllReduceForward.apply(full_output, tp_mesh.get_group()) + + def install_forward(self, module, ep_mesh, *, tp_mesh=None): + experts_forward = module.forward + ep_group, ep_size = ep_mesh.get_group(), ep_mesh.size() + + def tp_forward(hidden_states, top_k_index, top_k_weights): + # Read the full token count before the pre hook slices the inputs on TP. + num_tokens = hidden_states.size(0) + (hidden_states, top_k_index, top_k_weights), _ = self.transform_inputs_pre_forward( + module, (hidden_states, top_k_index, top_k_weights), {}, ep_mesh, tp_mesh=tp_mesh + ) + with self.context_around_forward(module, ep_mesh): + tokens, expert_ids, order, send_sizes, recv_sizes = self._dispatch_tokens( + hidden_states, top_k_index, module.num_experts, ep_group, ep_size + ) + expert_output = self._run_local_experts(experts_forward, tokens, expert_ids, module.num_experts) + output = self._combine_tokens( + expert_output, top_k_weights, order, send_sizes, recv_sizes, ep_group + ).to(hidden_states.dtype) + + return self.transform_output_post_forward(module, output, ep_mesh, tp_mesh=tp_mesh, num_tokens=num_tokens) + + module.forward = tp_forward + return module + + class MoeTensorParalellMegaMoeExperts(MoeExpertsParallel): """TP layer for DeepGEMM Mega MoE experts. @@ -785,6 +946,7 @@ class ParallelInterface(GeneralInterface): "sequence_parallel": SequenceParallel(use_local_output=True), "grouped_gemm": MoEParamShard(Shard(0), shards_expert_dim=True), "ep_router": EpRouterParallel(), + "ep_dispatch_experts": EpDispatchExpertsParallel(), "megamoe_router": RouterParallelMegaMoe(), "moe_tp_experts": MoeExpertsParallel(), "megamoe_experts": MoeTensorParalellMegaMoeExperts(), @@ -849,22 +1011,23 @@ def resolve_parallel_plans( return tp_plan, ep_plan -def apply_tensor_parallelism(model: nn.Module, tp_mesh: DeviceMesh, plan: dict[str, str] | None = None): - plan = model.tp_plan if plan is None else plan - _validate_parallel_plan_styles(plan) +def apply_tensor_parallelism(model, tp_mesh, tp_plan=None): + """Apply parameter sharding and forward hooks on the TP mesh.""" + tp_plan = model.tp_plan if tp_plan is None else tp_plan + _validate_parallel_plan_styles(tp_plan) for name, module in model.named_modules(): # Create DTensor placeholders so the loader knows which shard belongs to this rank. for p_name, _ in list(module.named_parameters(recurse=False)): full = f"{name}.{p_name}" if name else p_name - style_name = _get_parameter_plan(parameter_name=full, plan=plan, is_weight=True) + style_name = _get_parameter_plan(parameter_name=full, plan=tp_plan, is_weight=True) if style_name is not None and style_name in ALL_PARALLEL_STYLES: style = ALL_PARALLEL_STYLES[style_name] style.validate_param(module, p_name, tp_mesh, parameter_name=full) style.shard_param(module, p_name, tp_mesh) - # Install the input/output transforms required by this module's style. - style_name = _get_parameter_plan(parameter_name=name, plan=plan, is_weight=False) + # Install the input/output transforms required by this module's TP style. + style_name = _get_parameter_plan(parameter_name=name, plan=tp_plan, is_weight=False) if style_name is not None and style_name in ALL_PARALLEL_STYLES: if style_name == "mla_kv_a_proj": # MLA needs to know the qk_rope_head_dim to split the projection output into KV and RoPE parts. @@ -876,6 +1039,30 @@ def apply_tensor_parallelism(model: nn.Module, tp_mesh: DeviceMesh, plan: dict[s return model +def apply_expert_parallelism(model: nn.Module, ep_mesh: DeviceMesh, tp_mesh: DeviceMesh, plan: dict[str, str]): + """Shard experts on EP; use TP to split shared tokens and reconstruct outputs around dispatch.""" + for name, module in model.named_modules(): + for p_name, _ in list(module.named_parameters(recurse=False)): + full = f"{name}.{p_name}" if name else p_name + style_name = _get_parameter_plan(parameter_name=full, plan=plan, is_weight=True) + if style_name is not None and style_name in ALL_PARALLEL_STYLES: + style = ALL_PARALLEL_STYLES[style_name] + style.validate_param(module, p_name, ep_mesh, parameter_name=full) + style.shard_param(module, p_name, ep_mesh) + + # Dispatch hooks need both meshes to redistribute tokens between TP and EP ranks. + style_name = _get_parameter_plan(parameter_name=name, plan=plan, is_weight=False) + if style_name is not None and style_name in ALL_PARALLEL_STYLES: + style = ALL_PARALLEL_STYLES[style_name] + if style_name == "ep_dispatch_experts": + style.install_forward(module, ep_mesh=ep_mesh, tp_mesh=tp_mesh) + else: + style.install_forward(module, ep_mesh) + module._is_hooked = True + + return model + + def gather_state_dict_for_save( state_dict: dict[str, torch.Tensor], _tp_plan: dict[str, str], diff --git a/src/transformers/models/qwen3_moe/configuration_qwen3_moe.py b/src/transformers/models/qwen3_moe/configuration_qwen3_moe.py index 9a7d4b4c8b5b..62cf801edfcf 100644 --- a/src/transformers/models/qwen3_moe/configuration_qwen3_moe.py +++ b/src/transformers/models/qwen3_moe/configuration_qwen3_moe.py @@ -67,14 +67,14 @@ class Qwen3MoeConfig(PreTrainedConfig): "layers.*.mlp.up_proj": "colwise", "layers.*.mlp.down_proj": "rowwise", } - # Expert-only EP plan: only shards MoE experts, not attention. - # Attention is left unsharded — FSDP2 handles attention weight distribution. - # This allows EP to scale beyond num_kv_heads (not constrained by 4 for Qwen3-30B). + # Token dispatch by default, so `ep_size` can exceed `tp_size` (bounded by `num_key_value_heads`, 4 on + # Qwen3-30B): with `tp_size=1` the attention is left to FSDP2. For router masking with all-reduce, set + # `ep_size=tp_size` and pass `ep_plan={"layers.*.mlp.gate": "ep_router", "layers.*.mlp.experts": "moe_tp_experts"}` + # (prefixed with `model.` on the causal LM). base_model_ep_plan = { - "layers.*.mlp.gate": "ep_router", "layers.*.mlp.experts.gate_up_proj": "grouped_gemm", "layers.*.mlp.experts.down_proj": "grouped_gemm", - "layers.*.mlp.experts": "moe_tp_experts", + "layers.*.mlp.experts": "ep_dispatch_experts", } base_model_pp_plan = { "embed_tokens": (["input_ids"], ["inputs_embeds"]), diff --git a/tests/tensor_parallel/test_tensor_parallel.py b/tests/tensor_parallel/test_tensor_parallel.py index a265cbd4087f..3fed8106c98d 100644 --- a/tests/tensor_parallel/test_tensor_parallel.py +++ b/tests/tensor_parallel/test_tensor_parallel.py @@ -23,6 +23,7 @@ from transformers.distributed.tensor_parallel import ( ALL_PARALLEL_STYLES, ColwiseParallel, + EpDispatchExpertsParallel, PackedColwiseParallel, PackedRowwiseParallel, RowwiseParallel, @@ -31,7 +32,7 @@ # Qwen3 MoE's predefined plans, as resolved on `Qwen3MoeModel` (no `model.` prefix). -DENSE_TP_PLAN = { +TP_DENSE_PLAN = { "layers.*.self_attn.q_proj": "colwise", "layers.*.self_attn.k_proj": "colwise", "layers.*.self_attn.v_proj": "colwise", @@ -42,17 +43,17 @@ "layers.*.mlp.up_proj": "colwise", "layers.*.mlp.down_proj": "rowwise", } -EXPERT_TP_PLAN = { +TP_EXPERT_PLAN = { "layers.*.mlp.experts.gate_up_proj": "packed_colwise", "layers.*.mlp.experts.down_proj": "rowwise", "layers.*.mlp.experts": "moe_tp_experts", } EP_PLAN = { - "layers.*.mlp.gate": "ep_router", "layers.*.mlp.experts.gate_up_proj": "grouped_gemm", "layers.*.mlp.experts.down_proj": "grouped_gemm", - "layers.*.mlp.experts": "moe_tp_experts", + "layers.*.mlp.experts": "ep_dispatch_experts", } +EP_PLAN_MASKED = EP_PLAN | {"layers.*.mlp.gate": "ep_router", "layers.*.mlp.experts": "moe_tp_experts"} @require_torch @@ -74,6 +75,10 @@ def setUp(self): with torch.device("meta"): self.model = Qwen3MoeModel(self.config) + def _reset_plans(self): + self.model.tp_plan = TP_DENSE_PLAN | TP_EXPERT_PLAN + self.model.ep_plan = EP_PLAN.copy() + def test_ep_plan_setter(self): self.model.ep_plan = None self.assertEqual(self.model.ep_plan, {}) @@ -91,7 +96,7 @@ def test_disabled_parallelism_has_no_plans(self): def test_tp_only_keeps_experts_in_tp_plan(self): tp_plan, ep_plan = tensor_parallel.resolve_parallel_plans(self.model, DistributedConfig(tp_size=4)) - self.assertEqual(tp_plan, DENSE_TP_PLAN | EXPERT_TP_PLAN) + self.assertEqual(tp_plan, TP_DENSE_PLAN | TP_EXPERT_PLAN) self.assertEqual(ep_plan, {}) def test_ep_takes_experts_and_router_out_of_tp_plan(self): @@ -101,7 +106,7 @@ def test_ep_takes_experts_and_router_out_of_tp_plan(self): ): with self.subTest(config=config): tp_plan, ep_plan = tensor_parallel.resolve_parallel_plans(self.model, config) - self.assertEqual(tp_plan, DENSE_TP_PLAN) + self.assertEqual(tp_plan, TP_DENSE_PLAN | TP_EXPERT_PLAN) self.assertEqual(ep_plan, EP_PLAN) def test_legacy_flag_is_an_alias_for_ep_size(self): @@ -125,8 +130,8 @@ def test_legacy_flag_is_an_alias_for_ep_size(self): def test_ep_plan_is_a_dict_and_round_trips(self): with self.assertRaisesRegex(ValueError, "`ep_plan` must be a dictionary or None"): DistributedConfig(tp_size=4, ep_size=4, ep_plan="auto") - config = DistributedConfig(tp_size=4, ep_size=4, ep_plan={"layers.*.mlp.gate": "ep_router"}) - self.assertEqual(config.to_dict()["ep_plan"], {"layers.*.mlp.gate": "ep_router"}) + config = DistributedConfig(tp_size=4, ep_size=4, ep_plan=EP_PLAN_MASKED) + self.assertEqual(config.to_dict()["ep_plan"], EP_PLAN_MASKED) self.assertEqual(DistributedConfig.from_dict(config.to_dict()), config) def test_overrides_merge_into_the_predefined_plans(self): @@ -137,7 +142,7 @@ def test_overrides_merge_into_the_predefined_plans(self): ep_plan={"layers.*.mlp.experts.down_proj": "rowwise"}, ) tp_plan, ep_plan = tensor_parallel.resolve_parallel_plans(self.model, config) - self.assertEqual(tp_plan, DENSE_TP_PLAN | {"layers.*.self_attn.q_proj": "colwise_rep"}) + self.assertEqual(tp_plan, TP_DENSE_PLAN | {"layers.*.self_attn.q_proj": "colwise_rep"}) self.assertEqual(ep_plan, EP_PLAN | {"layers.*.mlp.experts.down_proj": "rowwise"}) # The merged plans are stored on the model, the config defaults are untouched. self.assertEqual(self.model.tp_plan["layers.*.self_attn.q_proj"], "colwise_rep") @@ -149,7 +154,7 @@ def test_overrides_merge_into_the_predefined_plans(self): self.assertEqual(config.ep_plan, {"layers.*.mlp.experts.down_proj": "rowwise"}) # The merged EP plan stays on the model but is not applied while EP is disabled. tp_plan, ep_plan = tensor_parallel.resolve_parallel_plans(self.model, DistributedConfig(tp_size=4)) - self.assertEqual(tp_plan, DENSE_TP_PLAN | EXPERT_TP_PLAN | {"layers.*.self_attn.q_proj": "colwise_rep"}) + self.assertEqual(tp_plan, TP_DENSE_PLAN | TP_EXPERT_PLAN | {"layers.*.self_attn.q_proj": "colwise_rep"}) self.assertEqual(ep_plan, {}) self.assertEqual(self.model.ep_plan["layers.*.mlp.experts.down_proj"], "rowwise") @@ -158,10 +163,11 @@ def test_ep_rules_take_precedence_over_tp_rules_for_the_same_modules(self): tp_size=4, ep_size=4, tp_plan={"layers.*.mlp.experts.gate_up_proj": "packed_rowwise", "layers.*.mlp.gate": "colwise"}, + ep_plan=EP_PLAN_MASKED, ) tp_plan, ep_plan = tensor_parallel.resolve_parallel_plans(self.model, config) - self.assertEqual(tp_plan, DENSE_TP_PLAN) - self.assertEqual(ep_plan, EP_PLAN) + self.assertEqual(tp_plan, TP_DENSE_PLAN) + self.assertEqual(ep_plan, EP_PLAN_MASKED) # The custom TP rules are kept on the model and apply as soon as EP is disabled. tp_plan, ep_plan = tensor_parallel.resolve_parallel_plans(self.model, DistributedConfig(tp_size=4)) self.assertEqual(tp_plan["layers.*.mlp.experts.gate_up_proj"], "packed_rowwise") @@ -173,7 +179,7 @@ def test_ep_requires_an_expert_plan(self): with self.assertRaisesRegex(ValueError, "does not define an expert-parallel plan"): tensor_parallel.resolve_parallel_plans(self.model, DistributedConfig(tp_size=4, ep_size=4)) config = DistributedConfig(tp_size=4, ep_size=4, ep_plan=EP_PLAN) - self.assertEqual(tensor_parallel.resolve_parallel_plans(self.model, config), (DENSE_TP_PLAN, EP_PLAN)) + self.assertEqual(tensor_parallel.resolve_parallel_plans(self.model, config), (TP_DENSE_PLAN, EP_PLAN)) def test_unmatched_override_keys_raise_without_changing_plans(self): original_tp_plan, original_ep_plan = self.model.tp_plan.copy(), self.model.ep_plan.copy() @@ -204,7 +210,7 @@ def test_override_keys_can_match_modules_parameters_or_existing_plan_keys(self): def test_head_model_overrides_need_the_model_prefix(self): with torch.device("meta"): model = Qwen3MoeForCausalLM(self.config) - config = DistributedConfig(tp_size=4, ep_size=4, ep_plan={"layers.*.mlp.gate": "ep_router"}) + config = DistributedConfig(tp_size=4, ep_size=4, ep_plan=EP_PLAN_MASKED) with self.assertRaisesRegex(ValueError, "match nothing in Qwen3MoeForCausalLM"): tensor_parallel.resolve_parallel_plans(model, config) @@ -212,16 +218,17 @@ def test_head_model_overrides_need_the_model_prefix(self): tp_size=4, ep_size=4, tp_plan={"model.layers.*.self_attn.q_proj": "colwise_rep"}, - ep_plan={"model.layers.*.mlp.gate": "ep_router"}, + ep_plan={f"model.{k}": v for k, v in EP_PLAN_MASKED.items()}, ) tp_plan, ep_plan = tensor_parallel.resolve_parallel_plans(model, config) - expected_tp_plan = {f"model.{k}": v for k, v in DENSE_TP_PLAN.items()} | {"lm_head": "colwise_gather_output"} + expected_tp_plan = {f"model.{k}": v for k, v in TP_DENSE_PLAN.items()} | {"lm_head": "colwise_gather_output"} self.assertEqual(tp_plan, expected_tp_plan | config.tp_plan) - self.assertEqual(ep_plan, {f"model.{k}": v for k, v in EP_PLAN.items()}) + self.assertEqual(ep_plan, {f"model.{k}": v for k, v in EP_PLAN_MASKED.items()}) def test_masked_ep_shards_and_installs_hooks_on_the_tp_mesh(self): tp_mesh = object() - _, ep_plan = tensor_parallel.resolve_parallel_plans(self.model, DistributedConfig(tp_size=4, ep_size=4)) + config = DistributedConfig(tp_size=4, ep_size=4, ep_plan=EP_PLAN_MASKED) + _, ep_plan = tensor_parallel.resolve_parallel_plans(self.model, config) experts, router = self.model.layers[0].mlp.experts, self.model.layers[0].mlp.gate with ( patch.object(ALL_PARALLEL_STYLES["grouped_gemm"], "validate_param") as validate, From b756980d837e9a107dafd60d07441f7bb32d605b Mon Sep 17 00:00:00 2001 From: 3outeille Date: Fri, 18 Sep 2026 17:17:11 +0000 Subject: [PATCH 23/86] clean --- tests/tensor_parallel/test_tensor_parallel.py | 1 - 1 file changed, 1 deletion(-) diff --git a/tests/tensor_parallel/test_tensor_parallel.py b/tests/tensor_parallel/test_tensor_parallel.py index 3fed8106c98d..90d4302bee2b 100644 --- a/tests/tensor_parallel/test_tensor_parallel.py +++ b/tests/tensor_parallel/test_tensor_parallel.py @@ -23,7 +23,6 @@ from transformers.distributed.tensor_parallel import ( ALL_PARALLEL_STYLES, ColwiseParallel, - EpDispatchExpertsParallel, PackedColwiseParallel, PackedRowwiseParallel, RowwiseParallel, From 4d5ef4cf441630c49ad2d2cb57cc7ea7f58f4f1e Mon Sep 17 00:00:00 2001 From: 3outeille Date: Sat, 19 Sep 2026 15:42:33 +0000 Subject: [PATCH 24/86] fix --- tests/tensor_parallel/test_tensor_parallel.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/tests/tensor_parallel/test_tensor_parallel.py b/tests/tensor_parallel/test_tensor_parallel.py index 90d4302bee2b..b9e8a3349254 100644 --- a/tests/tensor_parallel/test_tensor_parallel.py +++ b/tests/tensor_parallel/test_tensor_parallel.py @@ -105,7 +105,7 @@ def test_ep_takes_experts_and_router_out_of_tp_plan(self): ): with self.subTest(config=config): tp_plan, ep_plan = tensor_parallel.resolve_parallel_plans(self.model, config) - self.assertEqual(tp_plan, TP_DENSE_PLAN | TP_EXPERT_PLAN) + self.assertEqual(tp_plan, TP_DENSE_PLAN) self.assertEqual(ep_plan, EP_PLAN) def test_legacy_flag_is_an_alias_for_ep_size(self): From 556f4b9e4e716b3aa5b64cd7ecccaa3bfddcf7c5 Mon Sep 17 00:00:00 2001 From: 3outeille Date: Mon, 28 Sep 2026 15:43:24 +0000 Subject: [PATCH 25/86] make check repo --- src/transformers/models/mellum/configuration_mellum.py | 10 +++++----- 1 file changed, 5 insertions(+), 5 deletions(-) diff --git a/src/transformers/models/mellum/configuration_mellum.py b/src/transformers/models/mellum/configuration_mellum.py index 49deff0e2e69..9b3374607297 100644 --- a/src/transformers/models/mellum/configuration_mellum.py +++ b/src/transformers/models/mellum/configuration_mellum.py @@ -64,14 +64,14 @@ class MellumConfig(PreTrainedConfig): "layers.*.mlp.up_proj": "colwise", "layers.*.mlp.down_proj": "rowwise", } - # Expert-only EP plan: only shards MoE experts, not attention. - # Attention is left unsharded — FSDP2 handles attention weight distribution. - # This allows EP to scale beyond num_kv_heads (not constrained by 4 for Qwen3-30B). + # Token dispatch by default, so `ep_size` can exceed `tp_size` (bounded by `num_key_value_heads`, 4 on + # Qwen3-30B): with `tp_size=1` the attention is left to FSDP2. For router masking with all-reduce, set + # `ep_size=tp_size` and pass `ep_plan={"layers.*.mlp.gate": "ep_router", "layers.*.mlp.experts": "moe_tp_experts"}` + # (prefixed with `model.` on the causal LM). base_model_ep_plan = { - "layers.*.mlp.gate": "ep_router", "layers.*.mlp.experts.gate_up_proj": "grouped_gemm", "layers.*.mlp.experts.down_proj": "grouped_gemm", - "layers.*.mlp.experts": "moe_tp_experts", + "layers.*.mlp.experts": "ep_dispatch_experts", } base_model_pp_plan = { "embed_tokens": (["input_ids"], ["inputs_embeds"]), From 59578e01118e1eb3982fad55fc950c93e8303544 Mon Sep 17 00:00:00 2001 From: 3outeille Date: Mon, 28 Sep 2026 18:17:45 +0000 Subject: [PATCH 26/86] Drop the load-time ep_size == tp_size guard that rejects token dispatch The check ran before the expert plan was resolved, so it rejected every token-dispatch layout (ep_size != tp_size). DistributedConfig._validate_resolved_ep_plan already enforces ep_size == tp_size for the masked plan once the plan is known. --- src/transformers/distributed/mixin.py | 5 ----- 1 file changed, 5 deletions(-) diff --git a/src/transformers/distributed/mixin.py b/src/transformers/distributed/mixin.py index 584e2e5a6fb4..36a5b9c246f8 100644 --- a/src/transformers/distributed/mixin.py +++ b/src/transformers/distributed/mixin.py @@ -161,11 +161,6 @@ def prepare_distribute_model( if isinstance(distributed_config, dict): distributed_config = DistributedConfig.from_dict(distributed_config) - if distributed_config.ep_size > 1 and distributed_config.ep_size != distributed_config.tp_size: - raise ValueError( - "All-reduce expert parallelism requires `ep_size=tp_size` and identical tokens per EP group." - ) - if distributed_config.tp_size == 1 and distributed_config.fsdp_size == 1 and distributed_config.pp_size == 1: return distributed_config, device_map, None From 16fbcca98fed2b4040977cab22b071e71e28d8f2 Mon Sep 17 00:00:00 2001 From: 3outeille Date: Tue, 29 Sep 2026 12:53:15 +0000 Subject: [PATCH 27/86] remove all to all warning by using functional version --- .../distributed/tensor_parallel.py | 19 +++---------------- 1 file changed, 3 insertions(+), 16 deletions(-) diff --git a/src/transformers/distributed/tensor_parallel.py b/src/transformers/distributed/tensor_parallel.py index d6c14907efe6..c8f5fbecc577 100644 --- a/src/transformers/distributed/tensor_parallel.py +++ b/src/transformers/distributed/tensor_parallel.py @@ -38,7 +38,7 @@ if is_torch_distributed_available(): import torch.distributed as dist - from torch.distributed.nn.functional import all_to_all_single + from torch.distributed._functional_collectives import all_to_all_single from torch.distributed.tensor import DTensor, Partial, Replicate, Shard, distribute_tensor from torch.distributed.tensor.placement_types import _StridedShard @@ -789,7 +789,6 @@ def _dispatch_tokens( Also returns the sort order and the per-rank split sizes that `_combine_tokens` needs to reverse the exchange. """ - hidden_dim = hidden_states.size(-1) num_top_k = top_k_index.size(-1) # Sorting the selected pairs by expert groups them by owner rank, since each rank owns a contiguous range of @@ -805,13 +804,7 @@ def _dispatch_tokens( recv_counts = torch.empty_like(send_counts) torch.distributed.all_to_all_single(recv_counts, send_counts, group=ep_group) send_sizes, recv_sizes = torch.stack([send_counts.sum(dim=1), recv_counts.sum(dim=1)]).tolist() - recv_tokens = all_to_all_single( - send_tokens.new_empty(sum(recv_sizes), hidden_dim), - send_tokens, - output_split_sizes=recv_sizes, - input_split_sizes=send_sizes, - group=ep_group, - ) + recv_tokens = all_to_all_single(send_tokens, recv_sizes, send_sizes, ep_group) recv_expert_ids = torch.arange(num_local_experts, device=hidden_states.device).repeat(ep_size) recv_expert_ids = recv_expert_ids.repeat_interleave(recv_counts.reshape(-1), output_size=sum(recv_sizes)) return recv_tokens, recv_expert_ids, order, send_sizes, recv_sizes @@ -845,13 +838,7 @@ def _combine_tokens( """Return expert outputs to the token owners and combine them with routing weights.""" num_tokens, num_top_k = top_k_weights.shape hidden_dim = expert_output.size(-1) - recv_out = all_to_all_single( - expert_output.new_empty(order.numel(), hidden_dim), - expert_output, - output_split_sizes=send_sizes, - input_split_sizes=recv_sizes, - group=ep_group, - ) + recv_out = all_to_all_single(expert_output, send_sizes, recv_sizes, ep_group) # Restore the original (token, top-k slot) order, then apply routing weights. token_outputs = torch.empty_like(recv_out) token_outputs[order] = recv_out From 60fd0822fbb9792537acaf5a96dfbfa37bf0fb1e Mon Sep 17 00:00:00 2001 From: 3outeille Date: Wed, 30 Sep 2026 22:50:55 +0000 Subject: [PATCH 28/86] use ep plan instead of group_gemm is_expert --- src/transformers/distributed/fsdp.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/src/transformers/distributed/fsdp.py b/src/transformers/distributed/fsdp.py index d651ae61f398..60366b6ceb87 100644 --- a/src/transformers/distributed/fsdp.py +++ b/src/transformers/distributed/fsdp.py @@ -19,7 +19,7 @@ from ..utils import is_torch_available, is_torch_distributed_available, is_torch_greater_or_equal, logging, strtobool from ..utils.quantization_config import QuantizationMethod -from .tensor_parallel import replace_layer_number_by_wildcard +from .tensor_parallel import _get_parameter_plan, replace_layer_number_by_wildcard from .utils import _is_torch_distributed_initialized @@ -210,8 +210,8 @@ def apply_fully_sharded_data_parallelism(model: nn.Module, mesh_manager: MeshMan fsdp_policy_kwargs = _get_fsdp_policy_kwargs(distributed_config) if distributed_config.ep_size > 1 and "ep_dispatch_experts" in model.ep_plan.values(): expert_mesh = mesh_manager.get_mesh("efsdp") - for module in model.modules(): - if getattr(module, "_is_expert_parallel", False): + for module_name, module in model.named_modules(): + if _get_parameter_plan(module_name, model.ep_plan, is_weight=False) == "ep_dispatch_experts": fully_shard(module, mesh=expert_mesh, reshard_after_forward=True, **fsdp_policy_kwargs) # An expert group spans several data-parallel batches, so an expert's gradient sums over # all of them. FSDP2 would divide by the efsdp group size; dividing by fsdp_size instead From fc42fb9cbc24d9fc37d8d397a7eb7a25ce17e24e Mon Sep 17 00:00:00 2001 From: 3outeille Date: Wed, 30 Sep 2026 23:35:39 +0000 Subject: [PATCH 29/86] use dtensor instead of manuall all_reduce --- .../distributed/tensor_parallel.py | 36 ++++++++++--------- 1 file changed, 20 insertions(+), 16 deletions(-) diff --git a/src/transformers/distributed/tensor_parallel.py b/src/transformers/distributed/tensor_parallel.py index c8f5fbecc577..c7d7b12bf134 100644 --- a/src/transformers/distributed/tensor_parallel.py +++ b/src/transformers/distributed/tensor_parallel.py @@ -853,30 +853,34 @@ def transform_inputs_pre_forward(self, module, args, kwargs, mesh, *, tp_mesh=No top_k_weights = top_k_weights.to_local() if tp_mesh is None or tp_mesh.size() == 1: return (hidden_states, top_k_index, top_k_weights), kwargs - - tp_group = tp_mesh.get_group() - hidden_states = _AllReduceBackward.apply(hidden_states, tp_group) - top_k_weights = _AllReduceBackward.apply(top_k_weights, tp_group) - # TP ranks share the same batch. Slice tokens here so each is dispatched only once, then restore the full output in the post hook - num_tokens, tp_rank, tp_size = hidden_states.size(0), tp_mesh.get_local_rank(), tp_mesh.size() - rows = slice(num_tokens * tp_rank // tp_size, num_tokens * (tp_rank + 1) // tp_size) - return (hidden_states[rows], top_k_index[rows], top_k_weights[rows]), kwargs + # TP ranks share the same batch, so keep only this rank's rows and each token is dispatched once. + # Replicate -> Shard(0) is a local chunk (no communication); its backward all-gathers the row gradients. + hidden_states = DTensor.from_local(hidden_states, tp_mesh, [Replicate()], run_check=False) + top_k_index = DTensor.from_local(top_k_index, tp_mesh, [Replicate()], run_check=False) + top_k_weights = DTensor.from_local(top_k_weights, tp_mesh, [Replicate()], run_check=False) + + hidden_states = hidden_states.redistribute(tp_mesh, [Shard(0)]).to_local() + top_k_index = top_k_index.redistribute(tp_mesh, [Shard(0)]).to_local() + top_k_weights = top_k_weights.redistribute(tp_mesh, [Shard(0)]).to_local() + return (hidden_states, top_k_index, top_k_weights), kwargs def transform_output_post_forward(self, module, output, mesh, *, tp_mesh=None, num_tokens=None): if tp_mesh is None or tp_mesh.size() == 1: return output - tp_rank, tp_size = tp_mesh.get_local_rank(), tp_mesh.size() - rows = slice(num_tokens * tp_rank // tp_size, num_tokens * (tp_rank + 1) // tp_size) - full_output = output.new_zeros(num_tokens, output.size(-1)) - full_output[rows] = output - return _AllReduceForward.apply(full_output, tp_mesh.get_group()) + # Shard(0) -> Replicate is one all-gather of the row slices, which also handles uneven and empty slices. + hidden_dim = output.size(-1) + output = DTensor.from_local( + output.contiguous(), tp_mesh, [Shard(0)], shape=(num_tokens, hidden_dim), stride=(hidden_dim, 1) + ) + return output.full_tensor() def install_forward(self, module, ep_mesh, *, tp_mesh=None): + """Experts stay whole (Shard(0) on `ep_mesh`); `tp_mesh` is only the group of ranks holding the same batch.""" experts_forward = module.forward ep_group, ep_size = ep_mesh.get_group(), ep_mesh.size() - def tp_forward(hidden_states, top_k_index, top_k_weights): - # Read the full token count before the pre hook slices the inputs on TP. + def ep_forward(hidden_states, top_k_index, top_k_weights): + # Read the full token count before the pre hook slices the inputs across the batch replicas. num_tokens = hidden_states.size(0) (hidden_states, top_k_index, top_k_weights), _ = self.transform_inputs_pre_forward( module, (hidden_states, top_k_index, top_k_weights), {}, ep_mesh, tp_mesh=tp_mesh @@ -892,7 +896,7 @@ def tp_forward(hidden_states, top_k_index, top_k_weights): return self.transform_output_post_forward(module, output, ep_mesh, tp_mesh=tp_mesh, num_tokens=num_tokens) - module.forward = tp_forward + module.forward = ep_forward return module From 1478d28813a5fdf853e958889a6547f77584b281 Mon Sep 17 00:00:00 2001 From: 3outeille Date: Thu, 1 Oct 2026 00:10:58 +0000 Subject: [PATCH 30/86] better comment --- src/transformers/distributed/fsdp.py | 10 +++++++--- 1 file changed, 7 insertions(+), 3 deletions(-) diff --git a/src/transformers/distributed/fsdp.py b/src/transformers/distributed/fsdp.py index 60366b6ceb87..7d9a73223101 100644 --- a/src/transformers/distributed/fsdp.py +++ b/src/transformers/distributed/fsdp.py @@ -213,9 +213,13 @@ def apply_fully_sharded_data_parallelism(model: nn.Module, mesh_manager: MeshMan for module_name, module in model.named_modules(): if _get_parameter_plan(module_name, model.ep_plan, is_weight=False) == "ep_dispatch_experts": fully_shard(module, mesh=expert_mesh, reshard_after_forward=True, **fsdp_policy_kwargs) - # An expert group spans several data-parallel batches, so an expert's gradient sums over - # all of them. FSDP2 would divide by the efsdp group size; dividing by fsdp_size instead - # gives the same per-batch average the dense modules get on the fsdp mesh, even when efsdp has a single rank. + # Dense parameters on `fsdp` get a per-batch average: gradients summed over `fsdp_size` batches, + # then divided by `fsdp_size`. Experts must match that scale: + # - an EP group holds `ep_size / tp_size` distinct batches (TP ranks share a batch, and each token is + # dispatched exactly once), so an expert's local gradient already sums over that many batches; + # - the `efsdp` reduce then sums `efsdp_size` copies of that expert, one per EP group. + # The expert gradient therefore covers `ep_size / tp_size * efsdp_size = fsdp_size` batches, so divide + # by `fsdp_size` rather than `efsdp_size`. module.set_gradient_divide_factor(float(distributed_config.fsdp_size)) if torch.distributed.get_backend(expert_mesh.get_group()) != "nccl": # Non-NCCL backends need to sum first, then apply the division otherwise it runtime error. From 51472f9305d724d8eeff4647ad84eaf8c8b7e7b5 Mon Sep 17 00:00:00 2001 From: Ferdinand Mom <47445085+3outeille@users.noreply.github.com> Date: Thu, 1 Oct 2026 00:12:20 +0100 Subject: [PATCH 31/86] [`distributed`]: Fix DeepSeek-V4 tied embeddings with combined TP and EP (#48824) Fix tied embeddings for models with an EP-only base plan From ea814ebb52eb0316617a81bc90f9149470908e9a Mon Sep 17 00:00:00 2001 From: 3outeille Date: Thu, 1 Oct 2026 13:50:13 +0000 Subject: [PATCH 32/86] linting --- src/transformers/distributed/fsdp.py | 12 +++++------- src/transformers/distributed/tensor_parallel.py | 2 +- 2 files changed, 6 insertions(+), 8 deletions(-) diff --git a/src/transformers/distributed/fsdp.py b/src/transformers/distributed/fsdp.py index 7d9a73223101..dbc756111dcb 100644 --- a/src/transformers/distributed/fsdp.py +++ b/src/transformers/distributed/fsdp.py @@ -213,13 +213,11 @@ def apply_fully_sharded_data_parallelism(model: nn.Module, mesh_manager: MeshMan for module_name, module in model.named_modules(): if _get_parameter_plan(module_name, model.ep_plan, is_weight=False) == "ep_dispatch_experts": fully_shard(module, mesh=expert_mesh, reshard_after_forward=True, **fsdp_policy_kwargs) - # Dense parameters on `fsdp` get a per-batch average: gradients summed over `fsdp_size` batches, - # then divided by `fsdp_size`. Experts must match that scale: - # - an EP group holds `ep_size / tp_size` distinct batches (TP ranks share a batch, and each token is - # dispatched exactly once), so an expert's local gradient already sums over that many batches; - # - the `efsdp` reduce then sums `efsdp_size` copies of that expert, one per EP group. - # The expert gradient therefore covers `ep_size / tp_size * efsdp_size = fsdp_size` batches, so divide - # by `fsdp_size` rather than `efsdp_size`. + # Dense parameters on fsdp get a per-batch average: gradients summed over fsdp_size batches, + # then divided by fsdp_size. Experts must match that scale: + # - an EP group holds ep_size / tp_size distinct batches + # - the efsdp reduce then sums efsdp_size copies of that expert. + # The expert gradient therefore covers ep_size / tp_size * efsdp_size = fsdp_size batches module.set_gradient_divide_factor(float(distributed_config.fsdp_size)) if torch.distributed.get_backend(expert_mesh.get_group()) != "nccl": # Non-NCCL backends need to sum first, then apply the division otherwise it runtime error. diff --git a/src/transformers/distributed/tensor_parallel.py b/src/transformers/distributed/tensor_parallel.py index c7d7b12bf134..5690addbc1da 100644 --- a/src/transformers/distributed/tensor_parallel.py +++ b/src/transformers/distributed/tensor_parallel.py @@ -858,7 +858,7 @@ def transform_inputs_pre_forward(self, module, args, kwargs, mesh, *, tp_mesh=No hidden_states = DTensor.from_local(hidden_states, tp_mesh, [Replicate()], run_check=False) top_k_index = DTensor.from_local(top_k_index, tp_mesh, [Replicate()], run_check=False) top_k_weights = DTensor.from_local(top_k_weights, tp_mesh, [Replicate()], run_check=False) - + hidden_states = hidden_states.redistribute(tp_mesh, [Shard(0)]).to_local() top_k_index = top_k_index.redistribute(tp_mesh, [Shard(0)]).to_local() top_k_weights = top_k_weights.redistribute(tp_mesh, [Shard(0)]).to_local() From 925415d49a39edef1af963f26a9a4ebacbc33121 Mon Sep 17 00:00:00 2001 From: 3outeille Date: Wed, 16 Sep 2026 11:05:41 +0000 Subject: [PATCH 33/86] Add dense and expert device mesh views with a MeshManager `initialize_distributed_mesh` now builds two named views of the same ranks: `(pp, fsdp, tp)` for dense layers and `(pp, efsdp, ep)` for experts, both keeping size-one axes so callers select dimensions by name. `MeshManager` routes `ep`/`efsdp` lookups to the expert view and everything else to the dense view. `DistributedConfig` gains `ep_size` (defaults to `tp_size` when `enable_expert_parallel=True`) and `efsdp_size`, with size validation. Model execution is unchanged: expert sharding and FSDP still use the `tp` and `fsdp` axes, and loading rejects `ep_size != tp_size` until the all-to-all dispatcher lands. --- docs/source/en/expert_parallelism.md | 16 +- .../distributed/configuration_utils.py | 27 +- src/transformers/distributed/mixin.py | 26 +- src/transformers/distributed/utils.py | 56 +++-- src/transformers/modeling_utils.py | 5 +- tests/test_distributed_config.py | 232 ++++++++++++++++++ tests/test_fsdp_mixin.py | 2 +- 7 files changed, 325 insertions(+), 39 deletions(-) create mode 100644 tests/test_distributed_config.py diff --git a/docs/source/en/expert_parallelism.md b/docs/source/en/expert_parallelism.md index 59daa55f1bff..d01c1ad94576 100644 --- a/docs/source/en/expert_parallelism.md +++ b/docs/source/en/expert_parallelism.md @@ -66,7 +66,7 @@ distributed_config = DistributedConfig( model = AutoModelForCausalLM.from_pretrained("Qwen/Qwen3-30B-A3B", distributed_config=distributed_config) ``` -The model is loaded on a 2D `(fsdp, tp)` device mesh, and `tp_size * fsdp_size` must equal the number of processes. The expert parallel plan shards the experts across `tp`, then FSDP2 shards every parameter, experts included, across `fsdp` and owns their gradient reduction. Each `fsdp` rank trains on its own part of the batch. +The model is loaded on a `(pp, fsdp, tp)` device mesh with `pp_size=1`, and `tp_size * fsdp_size` must equal the number of processes. The expert parallel plan shards the experts across `tp`, then FSDP2 shards every parameter, experts included, across `fsdp` and owns their gradient reduction. Each `fsdp` rank trains on its own part of the batch. Load the model as usual, then train with [`Trainer`]. It takes the gradient norm across both meshes and gives each mesh its own optimizer param group. [`~Trainer.save_model`] gathers sharded weights into a regular checkpoint. This requires `accelerate>=1.12` so the `Trainer` can mirror `tp_size` and `fsdp_size` into [`~Accelerate.ParallelismConfig`]. @@ -81,6 +81,20 @@ The table below compares EP-only training with 2D EP+FSDP2 on 8xH100 GPUs. The w > [!WARNING] > Resuming from a checkpoint is not supported yet for models sharded at load time, so the [`Trainer`] only accepts `save_only_model=True` or `save_strategy="no"` for them. +## Mesh views + +`DistributedConfig` also accepts an explicit `ep_size`. For the current all-reduce implementation, +set `ep_size=tp_size`; `DistributedConfig(tp_size=4, ep_size=4)` is equivalent to +`DistributedConfig(tp_size=4, enable_expert_parallel=True)`. An explicit `ep_size=1` disables EP. + +Internally, a mesh manager provides two views of the same ranks: `(pp, fsdp, tp)` for dense layers +and `(pp, efsdp, ep)` for experts, where `efsdp_size = fsdp_size * tp_size // ep_size`. +Both retain size-one axes, so callers can select dimensions by name. The mesh builder supports +`ep_size` values that are multiples of `tp_size` and divide `fsdp_size * tp_size`. +Model loading currently rejects enabled EP layouts with `ep_size != tp_size` because all-reduce +requires identical tokens within each expert group. Expert sharding and FSDP continue to use the +`tp` and `fsdp` axes of the dense view. + ## API reference [[autodoc]] DistributedConfig diff --git a/src/transformers/distributed/configuration_utils.py b/src/transformers/distributed/configuration_utils.py index 55e7c85013dc..814c8335b2fd 100644 --- a/src/transformers/distributed/configuration_utils.py +++ b/src/transformers/distributed/configuration_utils.py @@ -33,7 +33,8 @@ class DistributedConfig: enable_sequence_parallel (`bool`, *optional*, defaults to `False`): Reserved for sequence parallelism. Not wired up yet. enable_expert_parallel (`bool`, *optional*, defaults to `False`): - Route MoE models through the expert-parallel path (``base_model_ep_plan``). + Route MoE models through the expert-parallel path (``base_model_ep_plan``). When `ep_size` is + omitted, sets it to `tp_size`. An explicit `ep_size` takes precedence. fsdp_size (`int`, *optional*): Number of devices for FSDP (data parallelism). If `None` and `tp_size` is set, defaults to 1. fsdp_cpu_offload (`bool`, *optional*, defaults to `False`): @@ -42,6 +43,9 @@ class DistributedConfig: Whether to enable mixed precision for FSDP2. pp_size (`int`, *optional*): Number of devices for pipeline parallelism. If `None` and another parallel mode is set, defaults to 1. + ep_size (`int`, *optional*): + Number of devices owning distinct expert shards. Defaults to 1, or to `tp_size` when + `enable_expert_parallel=True`. Model execution currently requires `ep_size=tp_size` when EP is enabled. """ tp_size: int | None = None @@ -52,10 +56,17 @@ class DistributedConfig: fsdp_cpu_offload: bool = False fsdp_mixed_precision: bool = False pp_size: int | None = None + ep_size: int | None = None + + @property + def efsdp_size(self) -> int: + """Size of the expert FSDP axis in the expert mesh view.""" + return self.fsdp_size * self.tp_size // self.ep_size def __post_init__(self): - if self.tp_plan is None and self.tp_size is None and self.fsdp_size is None and self.pp_size is None: - return + for value in (self.tp_size, self.fsdp_size, self.pp_size, self.ep_size): + if value is not None and value < 1: + raise ValueError(f"Parallelism sizes must be >= 1, got {value}.") if self.fsdp_size is None: self.fsdp_size = 1 @@ -73,6 +84,16 @@ def __post_init__(self): elif self.tp_size is None: self.tp_size = 1 + if self.ep_size is None: + self.ep_size = self.tp_size if self.enable_expert_parallel else 1 + self.enable_expert_parallel = self.ep_size > 1 + + if self.ep_size > 1: + if self.ep_size % self.tp_size: + raise ValueError("`ep_size` must be a multiple of `tp_size`.") + if (self.fsdp_size * self.tp_size) % self.ep_size: + raise ValueError("`ep_size` must divide `fsdp_size * tp_size`.") + if self.fsdp_size > 1 and self.pp_size > 1: raise ValueError( "Combining FSDP with pipeline parallelism is not supported yet. " diff --git a/src/transformers/distributed/mixin.py b/src/transformers/distributed/mixin.py index 646b8ac102cc..6f4e96b04d28 100644 --- a/src/transformers/distributed/mixin.py +++ b/src/transformers/distributed/mixin.py @@ -29,6 +29,7 @@ gather_state_dict_for_save, ) from .utils import ( + MeshManager, _distributed_barrier, _get_torch_distributed_rank, _is_torch_distributed_initialized, @@ -49,6 +50,7 @@ class DistributedMixin: """Distributed orchestration and save/load hooks for [`PreTrainedModel`].""" _device_mesh = None + _mesh_manager: MeshManager | None = None _tp_plan: dict[str, str] | None = None _ep_plan: dict[str, str] | None = None _tp_size = None @@ -142,13 +144,18 @@ def prepare_distribute_model( cls, distributed_config: DistributedConfig | dict | None, device_map=None, - ) -> tuple[DistributedConfig | None, object, object]: + ) -> tuple[DistributedConfig | None, object, MeshManager | None]: if distributed_config is None: return None, device_map, None if isinstance(distributed_config, dict): distributed_config = DistributedConfig.from_dict(distributed_config) + if distributed_config.ep_size > 1 and distributed_config.ep_size != distributed_config.tp_size: + raise ValueError( + "All-reduce expert parallelism requires `ep_size=tp_size` and identical tokens per EP group." + ) + if distributed_config.tp_size == 1 and distributed_config.fsdp_size == 1 and distributed_config.pp_size == 1: return distributed_config, device_map, None @@ -157,38 +164,39 @@ def prepare_distribute_model( if distributed_config.fsdp_size > 1 and not is_torch_greater_or_equal("2.7"): raise OSError("FSDP2 requires `torch>=2.7` (distributed checkpoint save/load).") - device_map, device_mesh = initialize_distributed_mesh(distributed_config) + device_map, mesh_manager = initialize_distributed_mesh(distributed_config) - return distributed_config, device_map, device_mesh + return distributed_config, device_map, mesh_manager @classmethod def maybe_distribute_model( cls, model: nn.Module, distributed_config: DistributedConfig | None, - device_mesh, + mesh_manager: MeshManager | None, ): """Apply TP or FSDP2 after model init, before weight loading.""" - if device_mesh is not None: + if mesh_manager is not None: model.config.distributed_config = distributed_config - model._device_mesh = device_mesh + model._mesh_manager = mesh_manager + model._device_mesh = mesh_manager.get_mesh(("pp", "fsdp", "tp")) model._tp_size = distributed_config.tp_size model._fsdp_size = distributed_config.fsdp_size if distributed_config.pp_size > 1: - pp_mesh = device_mesh["pp"] if device_mesh.ndim > 1 else device_mesh + pp_mesh = mesh_manager.get_mesh("pp") model = apply_pipeline_parallelism(model, pp_mesh) # Both may apply: the tensor/expert parallel plan shards across `tp` first, then FSDP2 # shards every parameter (the `tp`-sharded ones included) across `fsdp`. if distributed_config.tp_size > 1: - tp_mesh = device_mesh["tp"] if device_mesh.ndim > 1 else device_mesh + tp_mesh = mesh_manager.get_mesh("tp") if isinstance(distributed_config.tp_plan, dict): model.tp_plan = distributed_config.tp_plan model = apply_tensor_parallelism(model, tp_mesh) if distributed_config.fsdp_size > 1: - fsdp_mesh = device_mesh["fsdp"] if device_mesh.ndim > 1 else device_mesh + fsdp_mesh = mesh_manager.get_mesh("fsdp") model = apply_fully_sharded_data_parallelism(model, fsdp_mesh) return model diff --git a/src/transformers/distributed/utils.py b/src/transformers/distributed/utils.py index 571e4239f087..c843a9139f24 100644 --- a/src/transformers/distributed/utils.py +++ b/src/transformers/distributed/utils.py @@ -25,6 +25,7 @@ if TYPE_CHECKING: + from torch.distributed.device_mesh import DeviceMesh from torch.distributed.tensor import DTensor from .configuration_utils import DistributedConfig @@ -130,6 +131,20 @@ def _distributed_barrier(): torch.distributed.barrier() +class MeshManager: + """Named access to dense and expert parallel axes without exposing their view selection.""" + + def __init__(self, dense_mesh: DeviceMesh, expert_mesh: DeviceMesh): + self._dense_mesh = dense_mesh + self._expert_mesh = expert_mesh + + def get_mesh(self, dims: str | tuple[str, ...]) -> DeviceMesh: + """Select expert axes for `ep`/`efsdp`, otherwise dense axes; DeviceMesh handles slicing.""" + dims = (dims,) if isinstance(dims, str) else dims + mesh = self._expert_mesh if "ep" in dims or "efsdp" in dims else self._dense_mesh + return mesh[dims] + + # Retained for the legacy transformers.integrations.tensor_parallel API. def initialize_tensor_parallelism( tp_plan: str | dict[str, str] | None, tp_size: int | None = None, device_mesh=None, device_map=None @@ -223,22 +238,15 @@ def initialize_fully_sharded_data_parallelism(distributed_config: DistributedCon def initialize_distributed_mesh( distributed_config: DistributedConfig, -): - """Create a device mesh containing every configured parallel dimension.""" - mesh_shape = [] - mesh_dim_names = [] - - if distributed_config.pp_size > 1: - mesh_shape.append(distributed_config.pp_size) - mesh_dim_names.append("pp") - if distributed_config.fsdp_size > 1: - mesh_shape.append(distributed_config.fsdp_size) - mesh_dim_names.append("fsdp") - if distributed_config.tp_size > 1: - mesh_shape.append(distributed_config.tp_size) - mesh_dim_names.append("tp") - - if not mesh_shape: +) -> tuple[torch.device | None, MeshManager | None]: + """Build named dense and expert views, independently of the expert dispatcher. + + Both views include singleton dimensions so callers can always select their axes by name. + Each parameter's FSDP and TP/EP axes come from the same view. Separate roots avoid requiring + the newer `DeviceMesh._unflatten` API; the expert view is unused when EP is disabled. + """ + mesh_shape = (distributed_config.pp_size, distributed_config.fsdp_size, distributed_config.tp_size) + if mesh_shape == (1, 1, 1): return None, None device_type = torch._C._get_accelerator().type @@ -260,15 +268,17 @@ def initialize_distributed_mesh( else: device_map = torch.device(device_type) - device_mesh = torch.distributed.init_device_mesh( + dense_mesh = torch.distributed.init_device_mesh( device_type, - tuple(mesh_shape), - mesh_dim_names=tuple(mesh_dim_names), + mesh_shape, + mesh_dim_names=("pp", "fsdp", "tp"), ) - # A flattened sub-mesh, so an all-reduce over every rank is one collective instead of one per dimension. - if len(mesh_dim_names) > 1: - device_mesh._flatten("_".join(mesh_dim_names)) - return device_map, device_mesh + expert_mesh = torch.distributed.init_device_mesh( + device_type, + (distributed_config.pp_size, distributed_config.efsdp_size, distributed_config.ep_size), + mesh_dim_names=("pp", "efsdp", "ep"), + ) + return device_map, MeshManager(dense_mesh, expert_mesh) def gather_full_state_dict(model) -> dict[str, torch.Tensor]: diff --git a/src/transformers/modeling_utils.py b/src/transformers/modeling_utils.py index d8dd2981e274..07291b2870f6 100644 --- a/src/transformers/modeling_utils.py +++ b/src/transformers/modeling_utils.py @@ -4160,9 +4160,10 @@ def from_pretrained( distributed_config = DistributedConfig(tp_plan=tp_plan, tp_size=tp_size) if distributed_config is not None: - distributed_config, device_map, device_mesh = cls.prepare_distribute_model( + distributed_config, device_map, mesh_manager = cls.prepare_distribute_model( distributed_config, device_map=device_map ) + device_mesh = mesh_manager.get_mesh(("pp", "fsdp", "tp")) if mesh_manager is not None else None if gguf_file is not None and not is_accelerate_available(): raise ValueError("accelerate is required when loading a GGUF file `pip install accelerate`.") @@ -4309,7 +4310,7 @@ def from_pretrained( weight_conversions = get_model_conversion_mapping(model, key_mapping, hf_quantizer) if distributed_config is not None: - model = cls.maybe_distribute_model(model, distributed_config, device_mesh) + model = cls.maybe_distribute_model(model, distributed_config, mesh_manager) # Prepare the full device map if device_map is not None: diff --git a/tests/test_distributed_config.py b/tests/test_distributed_config.py new file mode 100644 index 000000000000..ba9f9c507f66 --- /dev/null +++ b/tests/test_distributed_config.py @@ -0,0 +1,232 @@ +# Copyright 2026 The HuggingFace Team. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +import os +import tempfile +import unittest +from datetime import timedelta +from unittest.mock import patch + +from transformers.distributed import DistributedConfig +from transformers.testing_utils import require_torch, require_torch_greater_or_equal +from transformers.utils import is_torch_available + + +if is_torch_available(): + import torch + import torch.distributed as dist + import torch.multiprocessing as mp + + from transformers.distributed.mixin import DistributedMixin + from transformers.distributed.utils import initialize_distributed_mesh + + +class DistributedConfigTest(unittest.TestCase): + def test_defaults_and_round_trip(self): + for kwargs in ({}, {"tp_size": 4}, {"fsdp_size": 4}, {"pp_size": 4}, {"tp_size": 2, "fsdp_size": 2}): + with self.subTest(kwargs=kwargs): + config = DistributedConfig(**kwargs) + self.assertEqual(config.ep_size, 1) + self.assertFalse(config.enable_expert_parallel) + self.assertEqual(config.efsdp_size, config.fsdp_size * config.tp_size) + self.assertEqual(DistributedConfig.from_dict(config.to_dict()), config) + + def test_legacy_and_explicit_ep_sizes(self): + legacy = DistributedConfig(tp_size=4, fsdp_size=2, enable_expert_parallel=True) + explicit = DistributedConfig(tp_size=4, fsdp_size=2, ep_size=4) + self.assertEqual(legacy, explicit) + self.assertEqual(DistributedConfig.from_dict(explicit.to_dict()), explicit) + for ep_size in (1, 4, 8): + with self.subTest(ep_size=ep_size): + config = DistributedConfig(tp_size=4, fsdp_size=2, ep_size=ep_size, enable_expert_parallel=True) + self.assertEqual(config.ep_size, ep_size) + self.assertEqual(config.enable_expert_parallel, ep_size > 1) + self.assertEqual((config.tp_size, config.fsdp_size), (4, 2)) + + def test_inferred_tp_size(self): + with patch.dict(os.environ, {"WORLD_SIZE": "8"}): + config = DistributedConfig(tp_plan="auto", fsdp_size=2, enable_expert_parallel=True) + self.assertEqual((config.tp_size, config.ep_size, config.efsdp_size), (4, 4, 2)) + + def test_expert_mesh_sizes(self): + for fsdp, tp, ep, efsdp in ((8, 1, 4, 2), (2, 2, 4, 1), (4, 2, 4, 2), (2, 2, 2, 2), (1, 4, 4, 1)): + with self.subTest(fsdp=fsdp, tp=tp, ep=ep): + config = DistributedConfig(fsdp_size=fsdp, tp_size=tp, ep_size=ep) + self.assertEqual(config.efsdp_size, efsdp) + self.assertEqual(DistributedConfig.from_dict(config.to_dict()), config) + + def test_invalid_sizes(self): + for name in ("tp_size", "fsdp_size", "pp_size", "ep_size"): + for value in (0, -1): + with self.subTest(name=name, value=value), self.assertRaisesRegex(ValueError, "must be >= 1"): + DistributedConfig(**{name: value}) + for kwargs, message in ( + ({"tp_size": 4, "ep_size": 2}, "multiple"), + ({"fsdp_size": 4, "ep_size": 3}, "must divide"), + ({"fsdp_size": 2, "pp_size": 2}, "pipeline parallelism"), + ({"ep_size": 2}, "must divide"), + ): + with self.subTest(kwargs=kwargs), self.assertRaisesRegex(ValueError, message): + DistributedConfig(**kwargs) + + +@require_torch +class DistributedMeshValidationTest(unittest.TestCase): + def test_disabled_mesh_does_not_initialize_distributed(self): + with patch("transformers.distributed.utils._ensure_torch_distributed") as initialize: + self.assertEqual(initialize_distributed_mesh(DistributedConfig()), (None, None)) + config, device_map, meshes = DistributedMixin.prepare_distribute_model({}, device_map="cpu") + self.assertEqual(config, DistributedConfig()) + self.assertEqual(device_map, "cpu") + self.assertIsNone(meshes) + initialize.assert_not_called() + + def test_model_loading_rejects_unsupported_ep_layout_before_initialization(self): + with patch("transformers.distributed.mixin.initialize_distributed_mesh") as initialize: + with self.assertRaisesRegex(ValueError, "ep_size=tp_size"): + DistributedMixin.prepare_distribute_model(DistributedConfig(fsdp_size=4, ep_size=2)) + initialize.assert_not_called() + + def test_world_size_mismatch(self): + with ( + patch("transformers.distributed.utils._ensure_torch_distributed"), + patch("torch._C._get_accelerator", return_value=torch.device("cpu")), + patch("torch.distributed.get_world_size", return_value=2), + self.assertRaisesRegex(RuntimeError, "requires 4 processes"), + ): + initialize_distributed_mesh(DistributedConfig(tp_size=4)) + + +def _mesh_worker(rank, rendezvous): + world_size = 4 + dist.init_process_group( + "gloo", + init_method=f"file://{rendezvous}", + rank=rank, + world_size=world_size, + timeout=timedelta(seconds=120), + ) + os.environ["LOCAL_RANK"] = str(rank) + try: + configs = [ + DistributedConfig(fsdp_size=4), + DistributedConfig(tp_size=4), + DistributedConfig(pp_size=4), + DistributedConfig(tp_size=2, fsdp_size=2), + DistributedConfig(tp_size=2, pp_size=2), + ] + configs += [DistributedConfig(fsdp_size=4, ep_size=ep) for ep in (2, 4)] + configs += [DistributedConfig(fsdp_size=2, tp_size=2, ep_size=ep) for ep in (2, 4)] + configs += [DistributedConfig(fsdp_size=1, tp_size=4, ep_size=4)] + for config in configs: + with patch("torch._C._get_accelerator", return_value=torch.device("cpu")): + _, meshes = initialize_distributed_mesh(config) + assert meshes.get_mesh(("pp", "fsdp", "tp")).mesh_dim_names == ("pp", "fsdp", "tp") + for axes in (("pp", "fsdp", "tp"), ("pp", "efsdp", "ep")): + assert meshes.get_mesh(axes).size() == world_size + assert meshes.get_mesh(axes).mesh_dim_names == axes + for name in axes: + assert meshes.get_mesh(name).size() == getattr(config, name + "_size") + assert meshes.get_mesh(("fsdp", "tp")).mesh_dim_names == ("fsdp", "tp") + assert meshes.get_mesh(("efsdp", "ep")).mesh_dim_names == ("efsdp", "ep") + for invalid in ("missing", ("tp", "ep"), ("fsdp", "efsdp")): + try: + meshes.get_mesh(invalid) + except KeyError: + pass + else: + raise AssertionError(f"Accepted invalid mesh dimensions: {invalid}") + stage_size = config.fsdp_size * config.tp_size + stage_start = rank // stage_size * stage_size + expert_rank = (rank - stage_start) % config.ep_size + ep_start = rank // config.ep_size * config.ep_size + assert dist.get_process_group_ranks(meshes.get_mesh("ep").get_group()) == list( + range(ep_start, ep_start + config.ep_size) + ) + assert dist.get_process_group_ranks(meshes.get_mesh("efsdp").get_group()) == list( + range(stage_start + expert_rank, stage_start + stage_size, config.ep_size) + ) + tp_start = rank // config.tp_size * config.tp_size + assert dist.get_process_group_ranks(meshes.get_mesh("tp").get_group()) == list( + range(tp_start, tp_start + config.tp_size) + ) + assert dist.get_process_group_ranks(meshes.get_mesh("fsdp").get_group()) == list( + range(stage_start + rank % config.tp_size, stage_start + stage_size, config.tp_size) + ) + assert dist.get_process_group_ranks(meshes.get_mesh("pp").get_group()) == list( + range(rank % stage_size, world_size, stage_size) + ) + assert meshes.get_mesh("ep").get_group() is meshes.get_mesh("ep").get_group() + finally: + dist.destroy_process_group() + + +def _dense_load_worker(rank, rendezvous): + from transformers import Qwen2Config, Qwen2ForCausalLM + + os.environ.update(RANK=str(rank), LOCAL_RANK=str(rank), WORLD_SIZE="2", LOCAL_WORLD_SIZE="2") + dist.init_process_group( + "gloo", init_method=f"file://{rendezvous}", rank=rank, world_size=2, timeout=timedelta(seconds=120) + ) + try: + torch.manual_seed(42) + config = Qwen2Config( + vocab_size=32, + hidden_size=8, + intermediate_size=8, + num_hidden_layers=2, + num_attention_heads=2, + num_key_value_heads=2, + ) + reference = Qwen2ForCausalLM(config).eval() + source = rendezvous + "_model" + if rank == 0: + reference.save_pretrained(source) + dist.barrier() + inputs = torch.tensor([[1, 2, 3]]) + generation_kwargs = { + "max_new_tokens": 2, + "do_sample": False, + "output_logits": True, + "return_dict_in_generate": True, + } + expected = reference.generate(inputs, **generation_kwargs) + for distributed_config in (DistributedConfig(tp_size=2), DistributedConfig(pp_size=2)): + with patch("torch._C._get_accelerator", return_value=torch.device("cpu")): + model = Qwen2ForCausalLM.from_pretrained(source, distributed_config=distributed_config).eval() + assert model._device_mesh is model._mesh_manager.get_mesh(("pp", "fsdp", "tp")) + actual = model.generate(inputs, **generation_kwargs) + torch.testing.assert_close(actual.sequences, expected.sequences) + torch.testing.assert_close(torch.stack(actual.logits), torch.stack(expected.logits)) + if distributed_config.tp_size > 1: + destination = rendezvous + "_saved" + model.save_pretrained(destination) + dist.barrier() + restored = Qwen2ForCausalLM.from_pretrained(destination).eval() + for name, param in restored.named_parameters(): + torch.testing.assert_close(param, dict(reference.named_parameters())[name], atol=0, rtol=0) + finally: + dist.destroy_process_group() + + +@require_torch +@require_torch_greater_or_equal("2.5") +class DistributedMeshTest(unittest.TestCase): + def test_mesh_groups(self): + with tempfile.TemporaryDirectory() as directory: + mp.spawn(_mesh_worker, args=(os.path.join(directory, "init"),), nprocs=4, join=True) + + def test_dense_load_generate_and_save(self): + with tempfile.TemporaryDirectory() as directory: + mp.spawn(_dense_load_worker, args=(os.path.join(directory, "init"),), nprocs=2, join=True) diff --git a/tests/test_fsdp_mixin.py b/tests/test_fsdp_mixin.py index 18fbcac99ee8..b59108cee436 100644 --- a/tests/test_fsdp_mixin.py +++ b/tests/test_fsdp_mixin.py @@ -539,7 +539,7 @@ def _test_fsdp2_expert_parallel_2d_vs_ddp_impl(rank, config_class, config_dict, distributed_config=DistributedConfig(tp_size=2, fsdp_size=dp, enable_expert_parallel=True), ) assert model.tp_size == 2 and model.fsdp_size == dp - assert model._device_mesh.mesh_dim_names == ("fsdp", "tp") + assert model._device_mesh.mesh_dim_names == ("pp", "fsdp", "tp") model.train() optimizer = torch.optim.Adam(model.parameters(), lr=LR, foreach=False) dp_rank = model._device_mesh["fsdp"].get_local_rank() From 0c227f1d303cf4755a8bace1dd01e7386cf2583b Mon Sep 17 00:00:00 2001 From: 3outeille Date: Wed, 16 Sep 2026 11:05:42 +0000 Subject: [PATCH 34/86] cleaning --- tests/test_distributed_config.py | 232 ------------------------------- 1 file changed, 232 deletions(-) delete mode 100644 tests/test_distributed_config.py diff --git a/tests/test_distributed_config.py b/tests/test_distributed_config.py deleted file mode 100644 index ba9f9c507f66..000000000000 --- a/tests/test_distributed_config.py +++ /dev/null @@ -1,232 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import os -import tempfile -import unittest -from datetime import timedelta -from unittest.mock import patch - -from transformers.distributed import DistributedConfig -from transformers.testing_utils import require_torch, require_torch_greater_or_equal -from transformers.utils import is_torch_available - - -if is_torch_available(): - import torch - import torch.distributed as dist - import torch.multiprocessing as mp - - from transformers.distributed.mixin import DistributedMixin - from transformers.distributed.utils import initialize_distributed_mesh - - -class DistributedConfigTest(unittest.TestCase): - def test_defaults_and_round_trip(self): - for kwargs in ({}, {"tp_size": 4}, {"fsdp_size": 4}, {"pp_size": 4}, {"tp_size": 2, "fsdp_size": 2}): - with self.subTest(kwargs=kwargs): - config = DistributedConfig(**kwargs) - self.assertEqual(config.ep_size, 1) - self.assertFalse(config.enable_expert_parallel) - self.assertEqual(config.efsdp_size, config.fsdp_size * config.tp_size) - self.assertEqual(DistributedConfig.from_dict(config.to_dict()), config) - - def test_legacy_and_explicit_ep_sizes(self): - legacy = DistributedConfig(tp_size=4, fsdp_size=2, enable_expert_parallel=True) - explicit = DistributedConfig(tp_size=4, fsdp_size=2, ep_size=4) - self.assertEqual(legacy, explicit) - self.assertEqual(DistributedConfig.from_dict(explicit.to_dict()), explicit) - for ep_size in (1, 4, 8): - with self.subTest(ep_size=ep_size): - config = DistributedConfig(tp_size=4, fsdp_size=2, ep_size=ep_size, enable_expert_parallel=True) - self.assertEqual(config.ep_size, ep_size) - self.assertEqual(config.enable_expert_parallel, ep_size > 1) - self.assertEqual((config.tp_size, config.fsdp_size), (4, 2)) - - def test_inferred_tp_size(self): - with patch.dict(os.environ, {"WORLD_SIZE": "8"}): - config = DistributedConfig(tp_plan="auto", fsdp_size=2, enable_expert_parallel=True) - self.assertEqual((config.tp_size, config.ep_size, config.efsdp_size), (4, 4, 2)) - - def test_expert_mesh_sizes(self): - for fsdp, tp, ep, efsdp in ((8, 1, 4, 2), (2, 2, 4, 1), (4, 2, 4, 2), (2, 2, 2, 2), (1, 4, 4, 1)): - with self.subTest(fsdp=fsdp, tp=tp, ep=ep): - config = DistributedConfig(fsdp_size=fsdp, tp_size=tp, ep_size=ep) - self.assertEqual(config.efsdp_size, efsdp) - self.assertEqual(DistributedConfig.from_dict(config.to_dict()), config) - - def test_invalid_sizes(self): - for name in ("tp_size", "fsdp_size", "pp_size", "ep_size"): - for value in (0, -1): - with self.subTest(name=name, value=value), self.assertRaisesRegex(ValueError, "must be >= 1"): - DistributedConfig(**{name: value}) - for kwargs, message in ( - ({"tp_size": 4, "ep_size": 2}, "multiple"), - ({"fsdp_size": 4, "ep_size": 3}, "must divide"), - ({"fsdp_size": 2, "pp_size": 2}, "pipeline parallelism"), - ({"ep_size": 2}, "must divide"), - ): - with self.subTest(kwargs=kwargs), self.assertRaisesRegex(ValueError, message): - DistributedConfig(**kwargs) - - -@require_torch -class DistributedMeshValidationTest(unittest.TestCase): - def test_disabled_mesh_does_not_initialize_distributed(self): - with patch("transformers.distributed.utils._ensure_torch_distributed") as initialize: - self.assertEqual(initialize_distributed_mesh(DistributedConfig()), (None, None)) - config, device_map, meshes = DistributedMixin.prepare_distribute_model({}, device_map="cpu") - self.assertEqual(config, DistributedConfig()) - self.assertEqual(device_map, "cpu") - self.assertIsNone(meshes) - initialize.assert_not_called() - - def test_model_loading_rejects_unsupported_ep_layout_before_initialization(self): - with patch("transformers.distributed.mixin.initialize_distributed_mesh") as initialize: - with self.assertRaisesRegex(ValueError, "ep_size=tp_size"): - DistributedMixin.prepare_distribute_model(DistributedConfig(fsdp_size=4, ep_size=2)) - initialize.assert_not_called() - - def test_world_size_mismatch(self): - with ( - patch("transformers.distributed.utils._ensure_torch_distributed"), - patch("torch._C._get_accelerator", return_value=torch.device("cpu")), - patch("torch.distributed.get_world_size", return_value=2), - self.assertRaisesRegex(RuntimeError, "requires 4 processes"), - ): - initialize_distributed_mesh(DistributedConfig(tp_size=4)) - - -def _mesh_worker(rank, rendezvous): - world_size = 4 - dist.init_process_group( - "gloo", - init_method=f"file://{rendezvous}", - rank=rank, - world_size=world_size, - timeout=timedelta(seconds=120), - ) - os.environ["LOCAL_RANK"] = str(rank) - try: - configs = [ - DistributedConfig(fsdp_size=4), - DistributedConfig(tp_size=4), - DistributedConfig(pp_size=4), - DistributedConfig(tp_size=2, fsdp_size=2), - DistributedConfig(tp_size=2, pp_size=2), - ] - configs += [DistributedConfig(fsdp_size=4, ep_size=ep) for ep in (2, 4)] - configs += [DistributedConfig(fsdp_size=2, tp_size=2, ep_size=ep) for ep in (2, 4)] - configs += [DistributedConfig(fsdp_size=1, tp_size=4, ep_size=4)] - for config in configs: - with patch("torch._C._get_accelerator", return_value=torch.device("cpu")): - _, meshes = initialize_distributed_mesh(config) - assert meshes.get_mesh(("pp", "fsdp", "tp")).mesh_dim_names == ("pp", "fsdp", "tp") - for axes in (("pp", "fsdp", "tp"), ("pp", "efsdp", "ep")): - assert meshes.get_mesh(axes).size() == world_size - assert meshes.get_mesh(axes).mesh_dim_names == axes - for name in axes: - assert meshes.get_mesh(name).size() == getattr(config, name + "_size") - assert meshes.get_mesh(("fsdp", "tp")).mesh_dim_names == ("fsdp", "tp") - assert meshes.get_mesh(("efsdp", "ep")).mesh_dim_names == ("efsdp", "ep") - for invalid in ("missing", ("tp", "ep"), ("fsdp", "efsdp")): - try: - meshes.get_mesh(invalid) - except KeyError: - pass - else: - raise AssertionError(f"Accepted invalid mesh dimensions: {invalid}") - stage_size = config.fsdp_size * config.tp_size - stage_start = rank // stage_size * stage_size - expert_rank = (rank - stage_start) % config.ep_size - ep_start = rank // config.ep_size * config.ep_size - assert dist.get_process_group_ranks(meshes.get_mesh("ep").get_group()) == list( - range(ep_start, ep_start + config.ep_size) - ) - assert dist.get_process_group_ranks(meshes.get_mesh("efsdp").get_group()) == list( - range(stage_start + expert_rank, stage_start + stage_size, config.ep_size) - ) - tp_start = rank // config.tp_size * config.tp_size - assert dist.get_process_group_ranks(meshes.get_mesh("tp").get_group()) == list( - range(tp_start, tp_start + config.tp_size) - ) - assert dist.get_process_group_ranks(meshes.get_mesh("fsdp").get_group()) == list( - range(stage_start + rank % config.tp_size, stage_start + stage_size, config.tp_size) - ) - assert dist.get_process_group_ranks(meshes.get_mesh("pp").get_group()) == list( - range(rank % stage_size, world_size, stage_size) - ) - assert meshes.get_mesh("ep").get_group() is meshes.get_mesh("ep").get_group() - finally: - dist.destroy_process_group() - - -def _dense_load_worker(rank, rendezvous): - from transformers import Qwen2Config, Qwen2ForCausalLM - - os.environ.update(RANK=str(rank), LOCAL_RANK=str(rank), WORLD_SIZE="2", LOCAL_WORLD_SIZE="2") - dist.init_process_group( - "gloo", init_method=f"file://{rendezvous}", rank=rank, world_size=2, timeout=timedelta(seconds=120) - ) - try: - torch.manual_seed(42) - config = Qwen2Config( - vocab_size=32, - hidden_size=8, - intermediate_size=8, - num_hidden_layers=2, - num_attention_heads=2, - num_key_value_heads=2, - ) - reference = Qwen2ForCausalLM(config).eval() - source = rendezvous + "_model" - if rank == 0: - reference.save_pretrained(source) - dist.barrier() - inputs = torch.tensor([[1, 2, 3]]) - generation_kwargs = { - "max_new_tokens": 2, - "do_sample": False, - "output_logits": True, - "return_dict_in_generate": True, - } - expected = reference.generate(inputs, **generation_kwargs) - for distributed_config in (DistributedConfig(tp_size=2), DistributedConfig(pp_size=2)): - with patch("torch._C._get_accelerator", return_value=torch.device("cpu")): - model = Qwen2ForCausalLM.from_pretrained(source, distributed_config=distributed_config).eval() - assert model._device_mesh is model._mesh_manager.get_mesh(("pp", "fsdp", "tp")) - actual = model.generate(inputs, **generation_kwargs) - torch.testing.assert_close(actual.sequences, expected.sequences) - torch.testing.assert_close(torch.stack(actual.logits), torch.stack(expected.logits)) - if distributed_config.tp_size > 1: - destination = rendezvous + "_saved" - model.save_pretrained(destination) - dist.barrier() - restored = Qwen2ForCausalLM.from_pretrained(destination).eval() - for name, param in restored.named_parameters(): - torch.testing.assert_close(param, dict(reference.named_parameters())[name], atol=0, rtol=0) - finally: - dist.destroy_process_group() - - -@require_torch -@require_torch_greater_or_equal("2.5") -class DistributedMeshTest(unittest.TestCase): - def test_mesh_groups(self): - with tempfile.TemporaryDirectory() as directory: - mp.spawn(_mesh_worker, args=(os.path.join(directory, "init"),), nprocs=4, join=True) - - def test_dense_load_generate_and_save(self): - with tempfile.TemporaryDirectory() as directory: - mp.spawn(_dense_load_worker, args=(os.path.join(directory, "init"),), nprocs=2, join=True) From a5958d4e8b32023be9e6109048be9204d00f3748 Mon Sep 17 00:00:00 2001 From: 3outeille Date: Wed, 16 Sep 2026 11:05:42 +0000 Subject: [PATCH 35/86] clean --- src/transformers/distributed/utils.py | 7 +------ 1 file changed, 1 insertion(+), 6 deletions(-) diff --git a/src/transformers/distributed/utils.py b/src/transformers/distributed/utils.py index c843a9139f24..c06182d86b2e 100644 --- a/src/transformers/distributed/utils.py +++ b/src/transformers/distributed/utils.py @@ -239,12 +239,7 @@ def initialize_fully_sharded_data_parallelism(distributed_config: DistributedCon def initialize_distributed_mesh( distributed_config: DistributedConfig, ) -> tuple[torch.device | None, MeshManager | None]: - """Build named dense and expert views, independently of the expert dispatcher. - - Both views include singleton dimensions so callers can always select their axes by name. - Each parameter's FSDP and TP/EP axes come from the same view. Separate roots avoid requiring - the newer `DeviceMesh._unflatten` API; the expert view is unused when EP is disabled. - """ + """Create a device mesh containing every configured parallel dimension.""" mesh_shape = (distributed_config.pp_size, distributed_config.fsdp_size, distributed_config.tp_size) if mesh_shape == (1, 1, 1): return None, None From 657191640c68b2edae1d8f6695c95bbb5b523aef Mon Sep 17 00:00:00 2001 From: 3outeille Date: Wed, 16 Sep 2026 12:15:08 +0000 Subject: [PATCH 36/86] update --- docs/source/en/expert_parallelism.md | 18 ++++++------ docs/source/en/fsdp.md | 2 +- docs/source/en/model_doc/hy_v4.md | 2 +- docs/source/en/model_doc/minimax_m3_vl.md | 2 +- .../distributed/configuration_utils.py | 28 +++++++++++++++---- src/transformers/distributed/mixin.py | 2 +- src/transformers/distributed/utils.py | 4 +++ .../deepseek_v4/test_modeling_deepseek_v4.py | 4 +-- tests/test_fsdp_mixin.py | 2 +- tests/test_tensor_parallel_mixin.py | 2 +- 10 files changed, 44 insertions(+), 22 deletions(-) diff --git a/docs/source/en/expert_parallelism.md b/docs/source/en/expert_parallelism.md index d01c1ad94576..8d23ca6c9feb 100644 --- a/docs/source/en/expert_parallelism.md +++ b/docs/source/en/expert_parallelism.md @@ -19,7 +19,7 @@ rendered properly in your Markdown viewer. ## DistributedConfig -Enable expert parallelism with the [`DistributedConfig`] class and the `enable_expert_parallel` argument. +Enable expert parallelism with the [`DistributedConfig`] class and the `ep_size` argument. The current all-reduce implementation requires `ep_size=tp_size`, so every rank in an expert group receives the same tokens. ```py import os @@ -30,7 +30,7 @@ from transformers.distributed.configuration_utils import DistributedConfig distributed_config = DistributedConfig( tp_size=int(os.environ["WORLD_SIZE"]), - enable_expert_parallel=True, + ep_size=int(os.environ["WORLD_SIZE"]), ) model = AutoModelForCausalLM.from_pretrained( @@ -42,7 +42,7 @@ model = AutoModelForCausalLM.from_pretrained( > [!TIP] > Expert parallelism automatically enables [tensor parallelism](./perf_infer_gpu_multi) for attention layers. -This argument switches to the `ep_plan` (expert parallel plan) defined in each MoE model's config file. The [`GroupedGemmParallel`] class splits expert weights so each device loads only its local experts. The `ep_router` routes tokens to experts and an all-reduce operation combines their outputs. +Setting `ep_size > 1` switches to the `ep_plan` (expert parallel plan) defined in each MoE model's config file. The [`GroupedGemmParallel`] class splits expert weights so each device loads only its local experts. The `ep_router` routes tokens to experts and an all-reduce operation combines their outputs. Launch your inference script with [torchrun](https://pytorch.org/docs/stable/elastic/run.html) and specify how many devices to use. The number of devices must evenly divide the total number of experts. @@ -52,16 +52,16 @@ torchrun --nproc-per-node 8 your_script.py ## Combining with FSDP2 -Expert parallelism only shards the experts. Everything else (attention, embeddings, norms) and its optimizer state is replicated on every expert-parallel rank, which limits how large a model you can train. Add [FSDP2](./fsdp) on a second mesh dimension with `fsdp_size`, and keep using `tp_size` for the expert parallel width (`tp_size` is the EP size). +Expert parallelism only shards the experts. Everything else (attention, embeddings, norms) and its optimizer state is replicated on every expert-parallel rank, which limits how large a model you can train. Add [FSDP2](./fsdp) on a second mesh dimension with `fsdp_size`, and keep `ep_size=tp_size` for the expert parallel width. ```py from transformers import AutoModelForCausalLM from transformers.distributed import DistributedConfig distributed_config = DistributedConfig( - tp_size=4, # expert parallel size + tp_size=4, + ep_size=4, # expert parallel size, must match tp_size fsdp_size=2, # data parallel shards - enable_expert_parallel=True, ) model = AutoModelForCausalLM.from_pretrained("Qwen/Qwen3-30B-A3B", distributed_config=distributed_config) ``` @@ -83,9 +83,9 @@ The table below compares EP-only training with 2D EP+FSDP2 on 8xH100 GPUs. The w ## Mesh views -`DistributedConfig` also accepts an explicit `ep_size`. For the current all-reduce implementation, -set `ep_size=tp_size`; `DistributedConfig(tp_size=4, ep_size=4)` is equivalent to -`DistributedConfig(tp_size=4, enable_expert_parallel=True)`. An explicit `ep_size=1` disables EP. +The legacy `enable_expert_parallel=True` flag is a deprecated alias for `ep_size=tp_size` when `ep_size` +is omitted, and will be removed in v5.20. It emits a `FutureWarning` and leaves `tp_size` and `fsdp_size` +unchanged. An explicit `ep_size` takes precedence over the flag, and `ep_size=1` disables EP. Internally, a mesh manager provides two views of the same ranks: `(pp, fsdp, tp)` for dense layers and `(pp, efsdp, ep)` for experts, where `efsdp_size = fsdp_size * tp_size // ep_size`. diff --git a/docs/source/en/fsdp.md b/docs/source/en/fsdp.md index dbbd629f22d1..5c12d7b55df3 100644 --- a/docs/source/en/fsdp.md +++ b/docs/source/en/fsdp.md @@ -122,7 +122,7 @@ TrainingArguments( > [!TIP] -> For mixture-of-experts models, `fsdp_size` can be combined with `tp_size` and `enable_expert_parallel=True` to shard the experts across one mesh dimension and everything else across the other. See [expert parallelism](./expert_parallelism#combining-with-fsdp2). +> For mixture-of-experts models, `fsdp_size` can be combined with `tp_size` and `ep_size` to shard the experts across one mesh dimension and everything else across the other. See [expert parallelism](./expert_parallelism#combining-with-fsdp2). ## Next steps diff --git a/docs/source/en/model_doc/hy_v4.md b/docs/source/en/model_doc/hy_v4.md index c6b3b31cf969..dd8cd789c622 100644 --- a/docs/source/en/model_doc/hy_v4.md +++ b/docs/source/en/model_doc/hy_v4.md @@ -83,7 +83,7 @@ model = AutoModelForCausalLM.from_pretrained( model = AutoModelForCausalLM.from_pretrained( model_id, dtype=torch.bfloat16, - distributed_config=DistributedConfig(tp_size=16, enable_expert_parallel=True), + distributed_config=DistributedConfig(tp_size=16, ep_size=16), ) ``` diff --git a/docs/source/en/model_doc/minimax_m3_vl.md b/docs/source/en/model_doc/minimax_m3_vl.md index 671d59922a44..2688f8c31ea2 100644 --- a/docs/source/en/model_doc/minimax_m3_vl.md +++ b/docs/source/en/model_doc/minimax_m3_vl.md @@ -167,7 +167,7 @@ model = AutoModelForImageTextToText.from_pretrained( quantization_config=FineGrainedFP8Config(dequantize=True), distributed_config=DistributedConfig( tp_size=int(os.environ["WORLD_SIZE"]), - enable_expert_parallel=True, + ep_size=int(os.environ["WORLD_SIZE"]), ), attn_implementation="kernels-staging/msa@v0", # MSA block-sparse attention kernel ) diff --git a/src/transformers/distributed/configuration_utils.py b/src/transformers/distributed/configuration_utils.py index 814c8335b2fd..ad33f96d39b7 100644 --- a/src/transformers/distributed/configuration_utils.py +++ b/src/transformers/distributed/configuration_utils.py @@ -14,6 +14,7 @@ import json import os +import warnings from dataclasses import asdict, dataclass from typing import Literal @@ -33,8 +34,8 @@ class DistributedConfig: enable_sequence_parallel (`bool`, *optional*, defaults to `False`): Reserved for sequence parallelism. Not wired up yet. enable_expert_parallel (`bool`, *optional*, defaults to `False`): - Route MoE models through the expert-parallel path (``base_model_ep_plan``). When `ep_size` is - omitted, sets it to `tp_size`. An explicit `ep_size` takes precedence. + Deprecated alias for `ep_size=tp_size` when `ep_size` is omitted, removed in v5.20. An explicit + `ep_size` takes precedence. This flag does not change `tp_size` or `fsdp_size`. fsdp_size (`int`, *optional*): Number of devices for FSDP (data parallelism). If `None` and `tp_size` is set, defaults to 1. fsdp_cpu_offload (`bool`, *optional*, defaults to `False`): @@ -44,8 +45,8 @@ class DistributedConfig: pp_size (`int`, *optional*): Number of devices for pipeline parallelism. If `None` and another parallel mode is set, defaults to 1. ep_size (`int`, *optional*): - Number of devices owning distinct expert shards. Defaults to 1, or to `tp_size` when - `enable_expert_parallel=True`. Model execution currently requires `ep_size=tp_size` when EP is enabled. + Number of devices owning distinct expert shards. Defaults to 1. Set it explicitly to enable EP. + Model execution currently requires `ep_size=tp_size` when EP is enabled. """ tp_size: int | None = None @@ -64,6 +65,11 @@ def efsdp_size(self) -> int: return self.fsdp_size * self.tp_size // self.ep_size def __post_init__(self): + self._resolve_parallelism() + self._validate_mesh_config() + + def _resolve_parallelism(self): + """Resolve parallel sizes and legacy EP settings.""" for value in (self.tp_size, self.fsdp_size, self.pp_size, self.ep_size): if value is not None and value < 1: raise ValueError(f"Parallelism sizes must be >= 1, got {value}.") @@ -84,10 +90,22 @@ def __post_init__(self): elif self.tp_size is None: self.tp_size = 1 + if self.enable_expert_parallel and self.ep_size is None: + self.ep_size = self.tp_size + warnings.warn( + f"`enable_expert_parallel` without `ep_size` is deprecated and will be removed in v5.20. " + f"Use ep_size={self.ep_size} instead.", + FutureWarning, + stacklevel=4, + ) + if self.ep_size is None: - self.ep_size = self.tp_size if self.enable_expert_parallel else 1 + self.ep_size = 1 + # Retain the legacy attribute for callers; internal EP decisions use ep_size. self.enable_expert_parallel = self.ep_size > 1 + def _validate_mesh_config(self): + """Validate mesh sizes before the model's expert plan is available.""" if self.ep_size > 1: if self.ep_size % self.tp_size: raise ValueError("`ep_size` must be a multiple of `tp_size`.") diff --git a/src/transformers/distributed/mixin.py b/src/transformers/distributed/mixin.py index 6f4e96b04d28..821eb8f2a733 100644 --- a/src/transformers/distributed/mixin.py +++ b/src/transformers/distributed/mixin.py @@ -88,7 +88,7 @@ def tp_plan(self) -> dict[str, str]: if hasattr(self.config, "distributed_config") and self.config.distributed_config.enable_expert_parallel: if not self._ep_plan: raise ValueError( - f"Expert parallelism was requested (`enable_expert_parallel=True`), but " + f"Expert parallelism was requested (`ep_size > 1`), but " f"`{self.__class__.__name__}` does not define an expert-parallel plan. Add a " f"`base_model_ep_plan` to its config, or disable expert parallelism." ) diff --git a/src/transformers/distributed/utils.py b/src/transformers/distributed/utils.py index c06182d86b2e..aa80eae9c3f5 100644 --- a/src/transformers/distributed/utils.py +++ b/src/transformers/distributed/utils.py @@ -135,6 +135,10 @@ class MeshManager: """Named access to dense and expert parallel axes without exposing their view selection.""" def __init__(self, dense_mesh: DeviceMesh, expert_mesh: DeviceMesh): + """ + dense_mesh: (pp, fsdp, tp) -> attention, dense MLPs, embeddings, lm_heads + expert_mesh: (pp, efsdp, ep) -> experts + """ self._dense_mesh = dense_mesh self._expert_mesh = expert_mesh diff --git a/tests/models/deepseek_v4/test_modeling_deepseek_v4.py b/tests/models/deepseek_v4/test_modeling_deepseek_v4.py index aecaa6584ac1..64f6e89e7d61 100644 --- a/tests/models/deepseek_v4/test_modeling_deepseek_v4.py +++ b/tests/models/deepseek_v4/test_modeling_deepseek_v4.py @@ -454,7 +454,7 @@ def main() -> int: attn_implementation="eager", experts_implementation=LOADTIME_DISPATCH, distributed_config=DistributedConfig( - tp_size=int(os.environ["WORLD_SIZE"]), enable_expert_parallel=True + tp_size=int(os.environ["WORLD_SIZE"]), ep_size=int(os.environ["WORLD_SIZE"]) ), ) model.eval() @@ -533,7 +533,7 @@ def main() -> int: attn_implementation="eager", experts_implementation=LOADTIME_DISPATCH, distributed_config=DistributedConfig( - tp_size=int(os.environ["WORLD_SIZE"]), enable_expert_parallel=True + tp_size=int(os.environ["WORLD_SIZE"]), ep_size=int(os.environ["WORLD_SIZE"]) ), ) model.eval() diff --git a/tests/test_fsdp_mixin.py b/tests/test_fsdp_mixin.py index b59108cee436..efcaed91a5f1 100644 --- a/tests/test_fsdp_mixin.py +++ b/tests/test_fsdp_mixin.py @@ -536,7 +536,7 @@ def _test_fsdp2_expert_parallel_2d_vs_ddp_impl(rank, config_class, config_dict, model = AutoModelForCausalLM.from_pretrained( init_model_dir, torch_dtype=dtype, - distributed_config=DistributedConfig(tp_size=2, fsdp_size=dp, enable_expert_parallel=True), + distributed_config=DistributedConfig(tp_size=2, fsdp_size=dp, ep_size=2), ) assert model.tp_size == 2 and model.fsdp_size == dp assert model._device_mesh.mesh_dim_names == ("pp", "fsdp", "tp") diff --git a/tests/test_tensor_parallel_mixin.py b/tests/test_tensor_parallel_mixin.py index 22d92bdd0948..38f78a0163f3 100644 --- a/tests/test_tensor_parallel_mixin.py +++ b/tests/test_tensor_parallel_mixin.py @@ -394,7 +394,7 @@ def _load_ep_and_reference_models(model_path, model_class): """Load EP model and non-EP reference model for comparison.""" model_ep = model_class.from_pretrained( model_path, - distributed_config=DistributedConfig(tp_size=dist.get_world_size(), enable_expert_parallel=True), + distributed_config=DistributedConfig(tp_size=dist.get_world_size(), ep_size=dist.get_world_size()), ) dist.barrier() From 90fbd09509071f1fbfe7fea9b8affdabb047c0b6 Mon Sep 17 00:00:00 2001 From: 3outeille Date: Wed, 16 Sep 2026 12:19:49 +0000 Subject: [PATCH 37/86] clean doc --- docs/source/en/expert_parallelism.md | 16 ---------------- 1 file changed, 16 deletions(-) diff --git a/docs/source/en/expert_parallelism.md b/docs/source/en/expert_parallelism.md index 8d23ca6c9feb..74687635dd94 100644 --- a/docs/source/en/expert_parallelism.md +++ b/docs/source/en/expert_parallelism.md @@ -66,8 +66,6 @@ distributed_config = DistributedConfig( model = AutoModelForCausalLM.from_pretrained("Qwen/Qwen3-30B-A3B", distributed_config=distributed_config) ``` -The model is loaded on a `(pp, fsdp, tp)` device mesh with `pp_size=1`, and `tp_size * fsdp_size` must equal the number of processes. The expert parallel plan shards the experts across `tp`, then FSDP2 shards every parameter, experts included, across `fsdp` and owns their gradient reduction. Each `fsdp` rank trains on its own part of the batch. - Load the model as usual, then train with [`Trainer`]. It takes the gradient norm across both meshes and gives each mesh its own optimizer param group. [`~Trainer.save_model`] gathers sharded weights into a regular checkpoint. This requires `accelerate>=1.12` so the `Trainer` can mirror `tp_size` and `fsdp_size` into [`~Accelerate.ParallelismConfig`]. The table below compares EP-only training with 2D EP+FSDP2 on 8xH100 GPUs. The workload is full fine-tuning of Qwen3-30B-A3B in bf16 at sequence length 2048. More FSDP shards cut peak memory, and tokens/s drop some because FSDP2 all-gathers and reduce-scatters the experts across `fsdp`. @@ -81,20 +79,6 @@ The table below compares EP-only training with 2D EP+FSDP2 on 8xH100 GPUs. The w > [!WARNING] > Resuming from a checkpoint is not supported yet for models sharded at load time, so the [`Trainer`] only accepts `save_only_model=True` or `save_strategy="no"` for them. -## Mesh views - -The legacy `enable_expert_parallel=True` flag is a deprecated alias for `ep_size=tp_size` when `ep_size` -is omitted, and will be removed in v5.20. It emits a `FutureWarning` and leaves `tp_size` and `fsdp_size` -unchanged. An explicit `ep_size` takes precedence over the flag, and `ep_size=1` disables EP. - -Internally, a mesh manager provides two views of the same ranks: `(pp, fsdp, tp)` for dense layers -and `(pp, efsdp, ep)` for experts, where `efsdp_size = fsdp_size * tp_size // ep_size`. -Both retain size-one axes, so callers can select dimensions by name. The mesh builder supports -`ep_size` values that are multiples of `tp_size` and divide `fsdp_size * tp_size`. -Model loading currently rejects enabled EP layouts with `ep_size != tp_size` because all-reduce -requires identical tokens within each expert group. Expert sharding and FSDP continue to use the -`tp` and `fsdp` axes of the dense view. - ## API reference [[autodoc]] DistributedConfig From e4561f09561a946ced76dfbe6ea949218df6f2a3 Mon Sep 17 00:00:00 2001 From: 3outeille Date: Wed, 16 Sep 2026 13:30:18 +0000 Subject: [PATCH 38/86] Use ep_size in the Kimi K2.5 expert parallel example --- docs/source/en/model_doc/kimi_k25.md | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/docs/source/en/model_doc/kimi_k25.md b/docs/source/en/model_doc/kimi_k25.md index 8fc5d85b15b0..20e22ab0cf92 100644 --- a/docs/source/en/model_doc/kimi_k25.md +++ b/docs/source/en/model_doc/kimi_k25.md @@ -49,7 +49,7 @@ import torch from transformers import AutoProcessor, AutoTokenizer, AutoModelForImageTextToText from transformers.distributed.configuration_utils import DistributedConfig -distributed_config = DistributedConfig(enable_expert_parallel=True) +distributed_config = DistributedConfig(tp_size=int(os.environ["WORLD_SIZE"]), ep_size=int(os.environ["WORLD_SIZE"])) processor = AutoProcessor.from_pretrained('moonshotai/Kimi-K2.6') model = AutoModelForImageTextToText.from_pretrained( From 7f4951f97733dce12a5865bdaf470039888c85bc Mon Sep 17 00:00:00 2001 From: 3outeille Date: Wed, 16 Sep 2026 13:32:37 +0000 Subject: [PATCH 39/86] Enable expert parallelism in the Mega MoE example --- docs/source/en/experts_interface.md | 7 +++++-- 1 file changed, 5 insertions(+), 2 deletions(-) diff --git a/docs/source/en/experts_interface.md b/docs/source/en/experts_interface.md index baada16c77d7..65c45f532772 100644 --- a/docs/source/en/experts_interface.md +++ b/docs/source/en/experts_interface.md @@ -153,14 +153,17 @@ This backend requires: - A Blackwell GPU (compute capability ≥ 10.0) with a CUDA toolkit (`nvcc`) 12.9 or later. - FP4-packed expert weights paired with UE8M0 weight scales (the pre-quantized checkpoint typically declares `expert_dtype="fp4"` and `scale_fmt="ue8m0"` in its config). -- A `torch.distributed` process group for the expert-parallel group, which the tensor-parallel wrapping supplies automatically. +- A `torch.distributed` process group for the expert-parallel group, which the expert-parallel wrapping supplies automatically when `ep_size > 1`. ```py import os from transformers import AutoModelForCausalLM, DistributedConfig -distributed_config = DistributedConfig(tp_size=int(os.environ["WORLD_SIZE"])) +distributed_config = DistributedConfig( + tp_size=int(os.environ["WORLD_SIZE"]), + ep_size=int(os.environ["WORLD_SIZE"]), +) model = AutoModelForCausalLM.from_pretrained( "deepseek-ai/DeepSeek-V4", experts_implementation="deepgemm_megamoe", From e65e3eafff2cc3378887e63d8a39e5fdfe2d9501 Mon Sep 17 00:00:00 2001 From: 3outeille Date: Fri, 18 Sep 2026 15:17:34 +0000 Subject: [PATCH 40/86] warn once --- .../distributed/configuration_utils.py | 14 ++++++++------ 1 file changed, 8 insertions(+), 6 deletions(-) diff --git a/src/transformers/distributed/configuration_utils.py b/src/transformers/distributed/configuration_utils.py index ad33f96d39b7..05a0e405264c 100644 --- a/src/transformers/distributed/configuration_utils.py +++ b/src/transformers/distributed/configuration_utils.py @@ -18,6 +18,7 @@ from dataclasses import asdict, dataclass from typing import Literal +from .utils import _get_torch_distributed_rank @dataclass class DistributedConfig: @@ -92,12 +93,13 @@ def _resolve_parallelism(self): if self.enable_expert_parallel and self.ep_size is None: self.ep_size = self.tp_size - warnings.warn( - f"`enable_expert_parallel` without `ep_size` is deprecated and will be removed in v5.20. " - f"Use ep_size={self.ep_size} instead.", - FutureWarning, - stacklevel=4, - ) + if _get_torch_distributed_rank() == 0: + warnings.warn( + f"`enable_expert_parallel` without `ep_size` is deprecated and will be removed in v5.20. " + f"Use ep_size={self.ep_size} instead.", + FutureWarning, + stacklevel=4, + ) if self.ep_size is None: self.ep_size = 1 From 03af9e8b46f179fd8b17f779a6c25ce5b2291785 Mon Sep 17 00:00:00 2001 From: 3outeille Date: Wed, 16 Sep 2026 11:05:41 +0000 Subject: [PATCH 41/86] Add dense and expert device mesh views with a MeshManager `initialize_distributed_mesh` now builds two named views of the same ranks: `(pp, fsdp, tp)` for dense layers and `(pp, efsdp, ep)` for experts, both keeping size-one axes so callers select dimensions by name. `MeshManager` routes `ep`/`efsdp` lookups to the expert view and everything else to the dense view. `DistributedConfig` gains `ep_size` (defaults to `tp_size` when `enable_expert_parallel=True`) and `efsdp_size`, with size validation. Model execution is unchanged: expert sharding and FSDP still use the `tp` and `fsdp` axes, and loading rejects `ep_size != tp_size` until the all-to-all dispatcher lands. --- docs/source/en/expert_parallelism.md | 14 ++ tests/test_distributed_config.py | 232 +++++++++++++++++++++++++++ 2 files changed, 246 insertions(+) create mode 100644 tests/test_distributed_config.py diff --git a/docs/source/en/expert_parallelism.md b/docs/source/en/expert_parallelism.md index 74687635dd94..62a34b648322 100644 --- a/docs/source/en/expert_parallelism.md +++ b/docs/source/en/expert_parallelism.md @@ -79,6 +79,20 @@ The table below compares EP-only training with 2D EP+FSDP2 on 8xH100 GPUs. The w > [!WARNING] > Resuming from a checkpoint is not supported yet for models sharded at load time, so the [`Trainer`] only accepts `save_only_model=True` or `save_strategy="no"` for them. +## Mesh views + +`DistributedConfig` also accepts an explicit `ep_size`. For the current all-reduce implementation, +set `ep_size=tp_size`; `DistributedConfig(tp_size=4, ep_size=4)` is equivalent to +`DistributedConfig(tp_size=4, enable_expert_parallel=True)`. An explicit `ep_size=1` disables EP. + +Internally, a mesh manager provides two views of the same ranks: `(pp, fsdp, tp)` for dense layers +and `(pp, efsdp, ep)` for experts, where `efsdp_size = fsdp_size * tp_size // ep_size`. +Both retain size-one axes, so callers can select dimensions by name. The mesh builder supports +`ep_size` values that are multiples of `tp_size` and divide `fsdp_size * tp_size`. +Model loading currently rejects enabled EP layouts with `ep_size != tp_size` because all-reduce +requires identical tokens within each expert group. Expert sharding and FSDP continue to use the +`tp` and `fsdp` axes of the dense view. + ## API reference [[autodoc]] DistributedConfig diff --git a/tests/test_distributed_config.py b/tests/test_distributed_config.py new file mode 100644 index 000000000000..ba9f9c507f66 --- /dev/null +++ b/tests/test_distributed_config.py @@ -0,0 +1,232 @@ +# Copyright 2026 The HuggingFace Team. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +import os +import tempfile +import unittest +from datetime import timedelta +from unittest.mock import patch + +from transformers.distributed import DistributedConfig +from transformers.testing_utils import require_torch, require_torch_greater_or_equal +from transformers.utils import is_torch_available + + +if is_torch_available(): + import torch + import torch.distributed as dist + import torch.multiprocessing as mp + + from transformers.distributed.mixin import DistributedMixin + from transformers.distributed.utils import initialize_distributed_mesh + + +class DistributedConfigTest(unittest.TestCase): + def test_defaults_and_round_trip(self): + for kwargs in ({}, {"tp_size": 4}, {"fsdp_size": 4}, {"pp_size": 4}, {"tp_size": 2, "fsdp_size": 2}): + with self.subTest(kwargs=kwargs): + config = DistributedConfig(**kwargs) + self.assertEqual(config.ep_size, 1) + self.assertFalse(config.enable_expert_parallel) + self.assertEqual(config.efsdp_size, config.fsdp_size * config.tp_size) + self.assertEqual(DistributedConfig.from_dict(config.to_dict()), config) + + def test_legacy_and_explicit_ep_sizes(self): + legacy = DistributedConfig(tp_size=4, fsdp_size=2, enable_expert_parallel=True) + explicit = DistributedConfig(tp_size=4, fsdp_size=2, ep_size=4) + self.assertEqual(legacy, explicit) + self.assertEqual(DistributedConfig.from_dict(explicit.to_dict()), explicit) + for ep_size in (1, 4, 8): + with self.subTest(ep_size=ep_size): + config = DistributedConfig(tp_size=4, fsdp_size=2, ep_size=ep_size, enable_expert_parallel=True) + self.assertEqual(config.ep_size, ep_size) + self.assertEqual(config.enable_expert_parallel, ep_size > 1) + self.assertEqual((config.tp_size, config.fsdp_size), (4, 2)) + + def test_inferred_tp_size(self): + with patch.dict(os.environ, {"WORLD_SIZE": "8"}): + config = DistributedConfig(tp_plan="auto", fsdp_size=2, enable_expert_parallel=True) + self.assertEqual((config.tp_size, config.ep_size, config.efsdp_size), (4, 4, 2)) + + def test_expert_mesh_sizes(self): + for fsdp, tp, ep, efsdp in ((8, 1, 4, 2), (2, 2, 4, 1), (4, 2, 4, 2), (2, 2, 2, 2), (1, 4, 4, 1)): + with self.subTest(fsdp=fsdp, tp=tp, ep=ep): + config = DistributedConfig(fsdp_size=fsdp, tp_size=tp, ep_size=ep) + self.assertEqual(config.efsdp_size, efsdp) + self.assertEqual(DistributedConfig.from_dict(config.to_dict()), config) + + def test_invalid_sizes(self): + for name in ("tp_size", "fsdp_size", "pp_size", "ep_size"): + for value in (0, -1): + with self.subTest(name=name, value=value), self.assertRaisesRegex(ValueError, "must be >= 1"): + DistributedConfig(**{name: value}) + for kwargs, message in ( + ({"tp_size": 4, "ep_size": 2}, "multiple"), + ({"fsdp_size": 4, "ep_size": 3}, "must divide"), + ({"fsdp_size": 2, "pp_size": 2}, "pipeline parallelism"), + ({"ep_size": 2}, "must divide"), + ): + with self.subTest(kwargs=kwargs), self.assertRaisesRegex(ValueError, message): + DistributedConfig(**kwargs) + + +@require_torch +class DistributedMeshValidationTest(unittest.TestCase): + def test_disabled_mesh_does_not_initialize_distributed(self): + with patch("transformers.distributed.utils._ensure_torch_distributed") as initialize: + self.assertEqual(initialize_distributed_mesh(DistributedConfig()), (None, None)) + config, device_map, meshes = DistributedMixin.prepare_distribute_model({}, device_map="cpu") + self.assertEqual(config, DistributedConfig()) + self.assertEqual(device_map, "cpu") + self.assertIsNone(meshes) + initialize.assert_not_called() + + def test_model_loading_rejects_unsupported_ep_layout_before_initialization(self): + with patch("transformers.distributed.mixin.initialize_distributed_mesh") as initialize: + with self.assertRaisesRegex(ValueError, "ep_size=tp_size"): + DistributedMixin.prepare_distribute_model(DistributedConfig(fsdp_size=4, ep_size=2)) + initialize.assert_not_called() + + def test_world_size_mismatch(self): + with ( + patch("transformers.distributed.utils._ensure_torch_distributed"), + patch("torch._C._get_accelerator", return_value=torch.device("cpu")), + patch("torch.distributed.get_world_size", return_value=2), + self.assertRaisesRegex(RuntimeError, "requires 4 processes"), + ): + initialize_distributed_mesh(DistributedConfig(tp_size=4)) + + +def _mesh_worker(rank, rendezvous): + world_size = 4 + dist.init_process_group( + "gloo", + init_method=f"file://{rendezvous}", + rank=rank, + world_size=world_size, + timeout=timedelta(seconds=120), + ) + os.environ["LOCAL_RANK"] = str(rank) + try: + configs = [ + DistributedConfig(fsdp_size=4), + DistributedConfig(tp_size=4), + DistributedConfig(pp_size=4), + DistributedConfig(tp_size=2, fsdp_size=2), + DistributedConfig(tp_size=2, pp_size=2), + ] + configs += [DistributedConfig(fsdp_size=4, ep_size=ep) for ep in (2, 4)] + configs += [DistributedConfig(fsdp_size=2, tp_size=2, ep_size=ep) for ep in (2, 4)] + configs += [DistributedConfig(fsdp_size=1, tp_size=4, ep_size=4)] + for config in configs: + with patch("torch._C._get_accelerator", return_value=torch.device("cpu")): + _, meshes = initialize_distributed_mesh(config) + assert meshes.get_mesh(("pp", "fsdp", "tp")).mesh_dim_names == ("pp", "fsdp", "tp") + for axes in (("pp", "fsdp", "tp"), ("pp", "efsdp", "ep")): + assert meshes.get_mesh(axes).size() == world_size + assert meshes.get_mesh(axes).mesh_dim_names == axes + for name in axes: + assert meshes.get_mesh(name).size() == getattr(config, name + "_size") + assert meshes.get_mesh(("fsdp", "tp")).mesh_dim_names == ("fsdp", "tp") + assert meshes.get_mesh(("efsdp", "ep")).mesh_dim_names == ("efsdp", "ep") + for invalid in ("missing", ("tp", "ep"), ("fsdp", "efsdp")): + try: + meshes.get_mesh(invalid) + except KeyError: + pass + else: + raise AssertionError(f"Accepted invalid mesh dimensions: {invalid}") + stage_size = config.fsdp_size * config.tp_size + stage_start = rank // stage_size * stage_size + expert_rank = (rank - stage_start) % config.ep_size + ep_start = rank // config.ep_size * config.ep_size + assert dist.get_process_group_ranks(meshes.get_mesh("ep").get_group()) == list( + range(ep_start, ep_start + config.ep_size) + ) + assert dist.get_process_group_ranks(meshes.get_mesh("efsdp").get_group()) == list( + range(stage_start + expert_rank, stage_start + stage_size, config.ep_size) + ) + tp_start = rank // config.tp_size * config.tp_size + assert dist.get_process_group_ranks(meshes.get_mesh("tp").get_group()) == list( + range(tp_start, tp_start + config.tp_size) + ) + assert dist.get_process_group_ranks(meshes.get_mesh("fsdp").get_group()) == list( + range(stage_start + rank % config.tp_size, stage_start + stage_size, config.tp_size) + ) + assert dist.get_process_group_ranks(meshes.get_mesh("pp").get_group()) == list( + range(rank % stage_size, world_size, stage_size) + ) + assert meshes.get_mesh("ep").get_group() is meshes.get_mesh("ep").get_group() + finally: + dist.destroy_process_group() + + +def _dense_load_worker(rank, rendezvous): + from transformers import Qwen2Config, Qwen2ForCausalLM + + os.environ.update(RANK=str(rank), LOCAL_RANK=str(rank), WORLD_SIZE="2", LOCAL_WORLD_SIZE="2") + dist.init_process_group( + "gloo", init_method=f"file://{rendezvous}", rank=rank, world_size=2, timeout=timedelta(seconds=120) + ) + try: + torch.manual_seed(42) + config = Qwen2Config( + vocab_size=32, + hidden_size=8, + intermediate_size=8, + num_hidden_layers=2, + num_attention_heads=2, + num_key_value_heads=2, + ) + reference = Qwen2ForCausalLM(config).eval() + source = rendezvous + "_model" + if rank == 0: + reference.save_pretrained(source) + dist.barrier() + inputs = torch.tensor([[1, 2, 3]]) + generation_kwargs = { + "max_new_tokens": 2, + "do_sample": False, + "output_logits": True, + "return_dict_in_generate": True, + } + expected = reference.generate(inputs, **generation_kwargs) + for distributed_config in (DistributedConfig(tp_size=2), DistributedConfig(pp_size=2)): + with patch("torch._C._get_accelerator", return_value=torch.device("cpu")): + model = Qwen2ForCausalLM.from_pretrained(source, distributed_config=distributed_config).eval() + assert model._device_mesh is model._mesh_manager.get_mesh(("pp", "fsdp", "tp")) + actual = model.generate(inputs, **generation_kwargs) + torch.testing.assert_close(actual.sequences, expected.sequences) + torch.testing.assert_close(torch.stack(actual.logits), torch.stack(expected.logits)) + if distributed_config.tp_size > 1: + destination = rendezvous + "_saved" + model.save_pretrained(destination) + dist.barrier() + restored = Qwen2ForCausalLM.from_pretrained(destination).eval() + for name, param in restored.named_parameters(): + torch.testing.assert_close(param, dict(reference.named_parameters())[name], atol=0, rtol=0) + finally: + dist.destroy_process_group() + + +@require_torch +@require_torch_greater_or_equal("2.5") +class DistributedMeshTest(unittest.TestCase): + def test_mesh_groups(self): + with tempfile.TemporaryDirectory() as directory: + mp.spawn(_mesh_worker, args=(os.path.join(directory, "init"),), nprocs=4, join=True) + + def test_dense_load_generate_and_save(self): + with tempfile.TemporaryDirectory() as directory: + mp.spawn(_dense_load_worker, args=(os.path.join(directory, "init"),), nprocs=2, join=True) From 69a145b21146bbf58da9ff56b1452fca317d22dd Mon Sep 17 00:00:00 2001 From: 3outeille Date: Wed, 16 Sep 2026 11:05:42 +0000 Subject: [PATCH 42/86] cleaning --- tests/test_distributed_config.py | 232 ------------------------------- 1 file changed, 232 deletions(-) delete mode 100644 tests/test_distributed_config.py diff --git a/tests/test_distributed_config.py b/tests/test_distributed_config.py deleted file mode 100644 index ba9f9c507f66..000000000000 --- a/tests/test_distributed_config.py +++ /dev/null @@ -1,232 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import os -import tempfile -import unittest -from datetime import timedelta -from unittest.mock import patch - -from transformers.distributed import DistributedConfig -from transformers.testing_utils import require_torch, require_torch_greater_or_equal -from transformers.utils import is_torch_available - - -if is_torch_available(): - import torch - import torch.distributed as dist - import torch.multiprocessing as mp - - from transformers.distributed.mixin import DistributedMixin - from transformers.distributed.utils import initialize_distributed_mesh - - -class DistributedConfigTest(unittest.TestCase): - def test_defaults_and_round_trip(self): - for kwargs in ({}, {"tp_size": 4}, {"fsdp_size": 4}, {"pp_size": 4}, {"tp_size": 2, "fsdp_size": 2}): - with self.subTest(kwargs=kwargs): - config = DistributedConfig(**kwargs) - self.assertEqual(config.ep_size, 1) - self.assertFalse(config.enable_expert_parallel) - self.assertEqual(config.efsdp_size, config.fsdp_size * config.tp_size) - self.assertEqual(DistributedConfig.from_dict(config.to_dict()), config) - - def test_legacy_and_explicit_ep_sizes(self): - legacy = DistributedConfig(tp_size=4, fsdp_size=2, enable_expert_parallel=True) - explicit = DistributedConfig(tp_size=4, fsdp_size=2, ep_size=4) - self.assertEqual(legacy, explicit) - self.assertEqual(DistributedConfig.from_dict(explicit.to_dict()), explicit) - for ep_size in (1, 4, 8): - with self.subTest(ep_size=ep_size): - config = DistributedConfig(tp_size=4, fsdp_size=2, ep_size=ep_size, enable_expert_parallel=True) - self.assertEqual(config.ep_size, ep_size) - self.assertEqual(config.enable_expert_parallel, ep_size > 1) - self.assertEqual((config.tp_size, config.fsdp_size), (4, 2)) - - def test_inferred_tp_size(self): - with patch.dict(os.environ, {"WORLD_SIZE": "8"}): - config = DistributedConfig(tp_plan="auto", fsdp_size=2, enable_expert_parallel=True) - self.assertEqual((config.tp_size, config.ep_size, config.efsdp_size), (4, 4, 2)) - - def test_expert_mesh_sizes(self): - for fsdp, tp, ep, efsdp in ((8, 1, 4, 2), (2, 2, 4, 1), (4, 2, 4, 2), (2, 2, 2, 2), (1, 4, 4, 1)): - with self.subTest(fsdp=fsdp, tp=tp, ep=ep): - config = DistributedConfig(fsdp_size=fsdp, tp_size=tp, ep_size=ep) - self.assertEqual(config.efsdp_size, efsdp) - self.assertEqual(DistributedConfig.from_dict(config.to_dict()), config) - - def test_invalid_sizes(self): - for name in ("tp_size", "fsdp_size", "pp_size", "ep_size"): - for value in (0, -1): - with self.subTest(name=name, value=value), self.assertRaisesRegex(ValueError, "must be >= 1"): - DistributedConfig(**{name: value}) - for kwargs, message in ( - ({"tp_size": 4, "ep_size": 2}, "multiple"), - ({"fsdp_size": 4, "ep_size": 3}, "must divide"), - ({"fsdp_size": 2, "pp_size": 2}, "pipeline parallelism"), - ({"ep_size": 2}, "must divide"), - ): - with self.subTest(kwargs=kwargs), self.assertRaisesRegex(ValueError, message): - DistributedConfig(**kwargs) - - -@require_torch -class DistributedMeshValidationTest(unittest.TestCase): - def test_disabled_mesh_does_not_initialize_distributed(self): - with patch("transformers.distributed.utils._ensure_torch_distributed") as initialize: - self.assertEqual(initialize_distributed_mesh(DistributedConfig()), (None, None)) - config, device_map, meshes = DistributedMixin.prepare_distribute_model({}, device_map="cpu") - self.assertEqual(config, DistributedConfig()) - self.assertEqual(device_map, "cpu") - self.assertIsNone(meshes) - initialize.assert_not_called() - - def test_model_loading_rejects_unsupported_ep_layout_before_initialization(self): - with patch("transformers.distributed.mixin.initialize_distributed_mesh") as initialize: - with self.assertRaisesRegex(ValueError, "ep_size=tp_size"): - DistributedMixin.prepare_distribute_model(DistributedConfig(fsdp_size=4, ep_size=2)) - initialize.assert_not_called() - - def test_world_size_mismatch(self): - with ( - patch("transformers.distributed.utils._ensure_torch_distributed"), - patch("torch._C._get_accelerator", return_value=torch.device("cpu")), - patch("torch.distributed.get_world_size", return_value=2), - self.assertRaisesRegex(RuntimeError, "requires 4 processes"), - ): - initialize_distributed_mesh(DistributedConfig(tp_size=4)) - - -def _mesh_worker(rank, rendezvous): - world_size = 4 - dist.init_process_group( - "gloo", - init_method=f"file://{rendezvous}", - rank=rank, - world_size=world_size, - timeout=timedelta(seconds=120), - ) - os.environ["LOCAL_RANK"] = str(rank) - try: - configs = [ - DistributedConfig(fsdp_size=4), - DistributedConfig(tp_size=4), - DistributedConfig(pp_size=4), - DistributedConfig(tp_size=2, fsdp_size=2), - DistributedConfig(tp_size=2, pp_size=2), - ] - configs += [DistributedConfig(fsdp_size=4, ep_size=ep) for ep in (2, 4)] - configs += [DistributedConfig(fsdp_size=2, tp_size=2, ep_size=ep) for ep in (2, 4)] - configs += [DistributedConfig(fsdp_size=1, tp_size=4, ep_size=4)] - for config in configs: - with patch("torch._C._get_accelerator", return_value=torch.device("cpu")): - _, meshes = initialize_distributed_mesh(config) - assert meshes.get_mesh(("pp", "fsdp", "tp")).mesh_dim_names == ("pp", "fsdp", "tp") - for axes in (("pp", "fsdp", "tp"), ("pp", "efsdp", "ep")): - assert meshes.get_mesh(axes).size() == world_size - assert meshes.get_mesh(axes).mesh_dim_names == axes - for name in axes: - assert meshes.get_mesh(name).size() == getattr(config, name + "_size") - assert meshes.get_mesh(("fsdp", "tp")).mesh_dim_names == ("fsdp", "tp") - assert meshes.get_mesh(("efsdp", "ep")).mesh_dim_names == ("efsdp", "ep") - for invalid in ("missing", ("tp", "ep"), ("fsdp", "efsdp")): - try: - meshes.get_mesh(invalid) - except KeyError: - pass - else: - raise AssertionError(f"Accepted invalid mesh dimensions: {invalid}") - stage_size = config.fsdp_size * config.tp_size - stage_start = rank // stage_size * stage_size - expert_rank = (rank - stage_start) % config.ep_size - ep_start = rank // config.ep_size * config.ep_size - assert dist.get_process_group_ranks(meshes.get_mesh("ep").get_group()) == list( - range(ep_start, ep_start + config.ep_size) - ) - assert dist.get_process_group_ranks(meshes.get_mesh("efsdp").get_group()) == list( - range(stage_start + expert_rank, stage_start + stage_size, config.ep_size) - ) - tp_start = rank // config.tp_size * config.tp_size - assert dist.get_process_group_ranks(meshes.get_mesh("tp").get_group()) == list( - range(tp_start, tp_start + config.tp_size) - ) - assert dist.get_process_group_ranks(meshes.get_mesh("fsdp").get_group()) == list( - range(stage_start + rank % config.tp_size, stage_start + stage_size, config.tp_size) - ) - assert dist.get_process_group_ranks(meshes.get_mesh("pp").get_group()) == list( - range(rank % stage_size, world_size, stage_size) - ) - assert meshes.get_mesh("ep").get_group() is meshes.get_mesh("ep").get_group() - finally: - dist.destroy_process_group() - - -def _dense_load_worker(rank, rendezvous): - from transformers import Qwen2Config, Qwen2ForCausalLM - - os.environ.update(RANK=str(rank), LOCAL_RANK=str(rank), WORLD_SIZE="2", LOCAL_WORLD_SIZE="2") - dist.init_process_group( - "gloo", init_method=f"file://{rendezvous}", rank=rank, world_size=2, timeout=timedelta(seconds=120) - ) - try: - torch.manual_seed(42) - config = Qwen2Config( - vocab_size=32, - hidden_size=8, - intermediate_size=8, - num_hidden_layers=2, - num_attention_heads=2, - num_key_value_heads=2, - ) - reference = Qwen2ForCausalLM(config).eval() - source = rendezvous + "_model" - if rank == 0: - reference.save_pretrained(source) - dist.barrier() - inputs = torch.tensor([[1, 2, 3]]) - generation_kwargs = { - "max_new_tokens": 2, - "do_sample": False, - "output_logits": True, - "return_dict_in_generate": True, - } - expected = reference.generate(inputs, **generation_kwargs) - for distributed_config in (DistributedConfig(tp_size=2), DistributedConfig(pp_size=2)): - with patch("torch._C._get_accelerator", return_value=torch.device("cpu")): - model = Qwen2ForCausalLM.from_pretrained(source, distributed_config=distributed_config).eval() - assert model._device_mesh is model._mesh_manager.get_mesh(("pp", "fsdp", "tp")) - actual = model.generate(inputs, **generation_kwargs) - torch.testing.assert_close(actual.sequences, expected.sequences) - torch.testing.assert_close(torch.stack(actual.logits), torch.stack(expected.logits)) - if distributed_config.tp_size > 1: - destination = rendezvous + "_saved" - model.save_pretrained(destination) - dist.barrier() - restored = Qwen2ForCausalLM.from_pretrained(destination).eval() - for name, param in restored.named_parameters(): - torch.testing.assert_close(param, dict(reference.named_parameters())[name], atol=0, rtol=0) - finally: - dist.destroy_process_group() - - -@require_torch -@require_torch_greater_or_equal("2.5") -class DistributedMeshTest(unittest.TestCase): - def test_mesh_groups(self): - with tempfile.TemporaryDirectory() as directory: - mp.spawn(_mesh_worker, args=(os.path.join(directory, "init"),), nprocs=4, join=True) - - def test_dense_load_generate_and_save(self): - with tempfile.TemporaryDirectory() as directory: - mp.spawn(_dense_load_worker, args=(os.path.join(directory, "init"),), nprocs=2, join=True) From caa7640946fab01428e39e5920ff2926eda25e25 Mon Sep 17 00:00:00 2001 From: 3outeille Date: Wed, 16 Sep 2026 12:15:08 +0000 Subject: [PATCH 43/86] update --- docs/source/en/expert_parallelism.md | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/docs/source/en/expert_parallelism.md b/docs/source/en/expert_parallelism.md index 62a34b648322..64c67a661799 100644 --- a/docs/source/en/expert_parallelism.md +++ b/docs/source/en/expert_parallelism.md @@ -81,9 +81,9 @@ The table below compares EP-only training with 2D EP+FSDP2 on 8xH100 GPUs. The w ## Mesh views -`DistributedConfig` also accepts an explicit `ep_size`. For the current all-reduce implementation, -set `ep_size=tp_size`; `DistributedConfig(tp_size=4, ep_size=4)` is equivalent to -`DistributedConfig(tp_size=4, enable_expert_parallel=True)`. An explicit `ep_size=1` disables EP. +The legacy `enable_expert_parallel=True` flag is a deprecated alias for `ep_size=tp_size` when `ep_size` +is omitted, and will be removed in v5.20. It emits a `FutureWarning` and leaves `tp_size` and `fsdp_size` +unchanged. An explicit `ep_size` takes precedence over the flag, and `ep_size=1` disables EP. Internally, a mesh manager provides two views of the same ranks: `(pp, fsdp, tp)` for dense layers and `(pp, efsdp, ep)` for experts, where `efsdp_size = fsdp_size * tp_size // ep_size`. From a3de4dc26902f8424eb946027e863a77c4e9da7b Mon Sep 17 00:00:00 2001 From: 3outeille Date: Wed, 16 Sep 2026 12:19:49 +0000 Subject: [PATCH 44/86] clean doc --- docs/source/en/expert_parallelism.md | 14 -------------- 1 file changed, 14 deletions(-) diff --git a/docs/source/en/expert_parallelism.md b/docs/source/en/expert_parallelism.md index 64c67a661799..74687635dd94 100644 --- a/docs/source/en/expert_parallelism.md +++ b/docs/source/en/expert_parallelism.md @@ -79,20 +79,6 @@ The table below compares EP-only training with 2D EP+FSDP2 on 8xH100 GPUs. The w > [!WARNING] > Resuming from a checkpoint is not supported yet for models sharded at load time, so the [`Trainer`] only accepts `save_only_model=True` or `save_strategy="no"` for them. -## Mesh views - -The legacy `enable_expert_parallel=True` flag is a deprecated alias for `ep_size=tp_size` when `ep_size` -is omitted, and will be removed in v5.20. It emits a `FutureWarning` and leaves `tp_size` and `fsdp_size` -unchanged. An explicit `ep_size` takes precedence over the flag, and `ep_size=1` disables EP. - -Internally, a mesh manager provides two views of the same ranks: `(pp, fsdp, tp)` for dense layers -and `(pp, efsdp, ep)` for experts, where `efsdp_size = fsdp_size * tp_size // ep_size`. -Both retain size-one axes, so callers can select dimensions by name. The mesh builder supports -`ep_size` values that are multiples of `tp_size` and divide `fsdp_size * tp_size`. -Model loading currently rejects enabled EP layouts with `ep_size != tp_size` because all-reduce -requires identical tokens within each expert group. Expert sharding and FSDP continue to use the -`tp` and `fsdp` axes of the dense view. - ## API reference [[autodoc]] DistributedConfig From 84cb4beecb3532a05daf3cf762976fd23c3e3d54 Mon Sep 17 00:00:00 2001 From: Ferdinand Mom <47445085+3outeille@users.noreply.github.com> Date: Sat, 19 Sep 2026 00:07:32 +0900 Subject: [PATCH 45/86] Update src/transformers/distributed/mixin.py Co-authored-by: Arthur <48595927+ArthurZucker@users.noreply.github.com> --- src/transformers/distributed/mixin.py | 1 + 1 file changed, 1 insertion(+) diff --git a/src/transformers/distributed/mixin.py b/src/transformers/distributed/mixin.py index 821eb8f2a733..224dddb63eeb 100644 --- a/src/transformers/distributed/mixin.py +++ b/src/transformers/distributed/mixin.py @@ -154,6 +154,7 @@ def prepare_distribute_model( if distributed_config.ep_size > 1 and distributed_config.ep_size != distributed_config.tp_size: raise ValueError( "All-reduce expert parallelism requires `ep_size=tp_size` and identical tokens per EP group." + "The token dispatch version allowing `ep_size != tp_size` is coming soon!" ) if distributed_config.tp_size == 1 and distributed_config.fsdp_size == 1 and distributed_config.pp_size == 1: From b55219529272a55a24c690d1538ca89c1d2ec3c4 Mon Sep 17 00:00:00 2001 From: 3outeille Date: Fri, 18 Sep 2026 15:42:27 +0000 Subject: [PATCH 46/86] style: ruff import block formatting in distributed config --- src/transformers/distributed/configuration_utils.py | 1 + 1 file changed, 1 insertion(+) diff --git a/src/transformers/distributed/configuration_utils.py b/src/transformers/distributed/configuration_utils.py index 05a0e405264c..9012271ec545 100644 --- a/src/transformers/distributed/configuration_utils.py +++ b/src/transformers/distributed/configuration_utils.py @@ -20,6 +20,7 @@ from .utils import _get_torch_distributed_rank + @dataclass class DistributedConfig: """ From d1cf2079d56bf34a4e74b0afa406a726e80108ef Mon Sep 17 00:00:00 2001 From: 3outeille Date: Wed, 30 Sep 2026 14:43:29 +0000 Subject: [PATCH 47/86] add more comments --- src/transformers/distributed/mixin.py | 8 +++--- src/transformers/distributed/utils.py | 37 +++++++++++++++++++++------ 2 files changed, 33 insertions(+), 12 deletions(-) diff --git a/src/transformers/distributed/mixin.py b/src/transformers/distributed/mixin.py index 224dddb63eeb..75b799f57e15 100644 --- a/src/transformers/distributed/mixin.py +++ b/src/transformers/distributed/mixin.py @@ -29,7 +29,7 @@ gather_state_dict_for_save, ) from .utils import ( - MeshManager, + TransformersDeviceMesh, _distributed_barrier, _get_torch_distributed_rank, _is_torch_distributed_initialized, @@ -50,7 +50,7 @@ class DistributedMixin: """Distributed orchestration and save/load hooks for [`PreTrainedModel`].""" _device_mesh = None - _mesh_manager: MeshManager | None = None + _mesh_manager: TransformersDeviceMesh | None = None _tp_plan: dict[str, str] | None = None _ep_plan: dict[str, str] | None = None _tp_size = None @@ -144,7 +144,7 @@ def prepare_distribute_model( cls, distributed_config: DistributedConfig | dict | None, device_map=None, - ) -> tuple[DistributedConfig | None, object, MeshManager | None]: + ) -> tuple[DistributedConfig | None, object, TransformersDeviceMesh | None]: if distributed_config is None: return None, device_map, None @@ -174,7 +174,7 @@ def maybe_distribute_model( cls, model: nn.Module, distributed_config: DistributedConfig | None, - mesh_manager: MeshManager | None, + mesh_manager: TransformersDeviceMesh | None, ): """Apply TP or FSDP2 after model init, before weight loading.""" if mesh_manager is not None: diff --git a/src/transformers/distributed/utils.py b/src/transformers/distributed/utils.py index aa80eae9c3f5..39cd8256bd8f 100644 --- a/src/transformers/distributed/utils.py +++ b/src/transformers/distributed/utils.py @@ -131,14 +131,35 @@ def _distributed_barrier(): torch.distributed.barrier() -class MeshManager: - """Named access to dense and expert parallel axes without exposing their view selection.""" +class TransformersDeviceMesh: + """ + Holds the device meshes used by a model. + + dense layers and experts are sharded differently, so they need different views of the same ranks. + + dense : (pp, fsdp, tp) attention, dense MLPs, embeddings, lm_heads + expert : (pp, efsdp, ep) experts + + Both views cover the same world, so pp * fsdp * tp == pp * efsdp * ep. + efsdp is not something you pick, it is whatever is left once ep is fixed: + efsdp = fsdp * tp / ep. It is the FSDP axis for expert weights same role `fsdp` plays for the dense params. + + There is no etp (expert tensor parallel) axis yet meaning experts are never tensor-sharded here. + If one were ever added, the identity would become pp * efsdp * ep * etp == pp * fsdp * tp and efsdp would shrink by etp + (efsdp = fsdp * tp / (ep * etp)) + + Why not reuse fsdp mesh ? When ep_size == tp_size, efsdp == fsdp, both in size and in which + ranks are grouped together, so the fsdp axis of the dense mesh would work for experts too. + As soon as ep_size != tp_size the two group different ranks and you need a separate axis. + + Regarding ep value, We decide to default it to node width (8 on most machines) so all-to-all never leaves the node. + - On a single node, ep == fsdp * tp thus efsdp = 1, the axis does nothing. + - On several nodes, we still keep ep at node width, since all-to-all across nodes is expensive. + However, each node then holds a full copy of the expert group and efsdp is the number of copies, which is where FSDP happens for the experts + i.e: 2 nodes x 8 GPUs -> efsdp = 16 / 8 = 2, one EP group per node, two copies, sharded over efsdp. + """ def __init__(self, dense_mesh: DeviceMesh, expert_mesh: DeviceMesh): - """ - dense_mesh: (pp, fsdp, tp) -> attention, dense MLPs, embeddings, lm_heads - expert_mesh: (pp, efsdp, ep) -> experts - """ self._dense_mesh = dense_mesh self._expert_mesh = expert_mesh @@ -242,7 +263,7 @@ def initialize_fully_sharded_data_parallelism(distributed_config: DistributedCon def initialize_distributed_mesh( distributed_config: DistributedConfig, -) -> tuple[torch.device | None, MeshManager | None]: +) -> tuple[torch.device | None, TransformersDeviceMesh | None]: """Create a device mesh containing every configured parallel dimension.""" mesh_shape = (distributed_config.pp_size, distributed_config.fsdp_size, distributed_config.tp_size) if mesh_shape == (1, 1, 1): @@ -277,7 +298,7 @@ def initialize_distributed_mesh( (distributed_config.pp_size, distributed_config.efsdp_size, distributed_config.ep_size), mesh_dim_names=("pp", "efsdp", "ep"), ) - return device_map, MeshManager(dense_mesh, expert_mesh) + return device_map, TransformersDeviceMesh(dense_mesh, expert_mesh) def gather_full_state_dict(model) -> dict[str, torch.Tensor]: From 4b655fbfd45ca5b6ce6943bf31129c898d263239 Mon Sep 17 00:00:00 2001 From: 3outeille Date: Fri, 18 Sep 2026 17:22:06 +0000 Subject: [PATCH 48/86] Decouple tp_plan and ep_plan for expert parallelism. Keep TP and EP plans separate on the model, resolve overrides via resolve_parallel_plans, and apply both through tensor parallel sharding on the tp mesh when ep_size matches tp_size. --- docs/source/en/expert_parallelism.md | 25 +- .../distributed/configuration_utils.py | 26 ++- src/transformers/distributed/mixin.py | 83 ++++--- .../distributed/tensor_parallel.py | 88 +++++-- src/transformers/modeling_utils.py | 4 +- tests/tensor_parallel/test_tensor_parallel.py | 214 +++++++++++++++++- tests/test_modeling_common.py | 7 +- tests/test_tensor_parallel_mixin.py | 7 +- 8 files changed, 376 insertions(+), 78 deletions(-) diff --git a/docs/source/en/expert_parallelism.md b/docs/source/en/expert_parallelism.md index 74687635dd94..785c7f7be2f2 100644 --- a/docs/source/en/expert_parallelism.md +++ b/docs/source/en/expert_parallelism.md @@ -39,10 +39,12 @@ model = AutoModelForCausalLM.from_pretrained( ) ``` -> [!TIP] -> Expert parallelism automatically enables [tensor parallelism](./perf_infer_gpu_multi) for attention layers. +Each MoE model defines two plans in its config: `base_model_tp_plan` for the dense modules and `base_model_ep_plan` for the experts. They are exposed on the loaded model as `model.tp_plan` and `model.ep_plan`. With `tp_size > 1` and `ep_size > 1`, both apply: the [tensor parallel](./perf_infer_gpu_multi) plan shards attention and the dense MLPs, and the expert parallel plan shards the experts. EP rules take precedence over TP rules for the same modules, so expert weights are sharded once, by the EP plan. In the EP plan, the [`GroupedGemmParallel`] style splits the expert weights along the expert dimension so each rank loads only its local experts, and `ep_router` masks the experts that live on other ranks before an all-reduce combines the expert outputs. + +`tp_plan` is applied only when `tp_size > 1`, and `ep_plan` only when `ep_size > 1`. With TP enabled and EP disabled, the full TP plan applies, expert rules included. -Setting `ep_size > 1` switches to the `ep_plan` (expert parallel plan) defined in each MoE model's config file. The [`GroupedGemmParallel`] class splits expert weights so each device loads only its local experts. The `ep_router` routes tokens to experts and an all-reduce operation combines their outputs. +> [!TIP] +> `enable_expert_parallel=True` is a deprecated alias for `ep_size=tp_size`, used only when `ep_size` is omitted, and emits a `FutureWarning`. Launch your inference script with [torchrun](https://pytorch.org/docs/stable/elastic/run.html) and specify how many devices to use. The number of devices must evenly divide the total number of experts. @@ -50,9 +52,24 @@ Launch your inference script with [torchrun](https://pytorch.org/docs/stable/ela torchrun --nproc-per-node 8 your_script.py ``` +### Overriding the plans + +Pass `tp_plan={...}` or `ep_plan={...}` to [`DistributedConfig`] to override individual rules of the predefined plans. Unspecified rules are kept, and the merged plans are stored on the model. Each key must match a module, a parameter, or an existing plan entry; otherwise loading raises a `ValueError` before anything is sharded. Use the full path as seen from the loaded model, so `model.layers.*` for a causal LM and `layers.*` for its base model. + +```py +distributed_config = DistributedConfig( + tp_size=4, + ep_size=4, + tp_plan={"model.layers.*.self_attn.q_proj": "colwise_rep"}, + ep_plan={"model.layers.*.mlp.experts.down_proj": "grouped_gemm"}, +) +``` + +Providing a plan does not infer parallel sizes: set `tp_size` and `ep_size` explicitly. + ## Combining with FSDP2 -Expert parallelism only shards the experts. Everything else (attention, embeddings, norms) and its optimizer state is replicated on every expert-parallel rank, which limits how large a model you can train. Add [FSDP2](./fsdp) on a second mesh dimension with `fsdp_size`, and keep `ep_size=tp_size` for the expert parallel width. +Tensor and expert parallelism shard the weights across `tp`, but the optimizer state and the modules without a rule are still replicated on every rank of the group, which limits how large a model you can train. Add [FSDP2](./fsdp) on a second mesh dimension with `fsdp_size`, and keep `ep_size=tp_size` for the expert parallel width. ```py from transformers import AutoModelForCausalLM diff --git a/src/transformers/distributed/configuration_utils.py b/src/transformers/distributed/configuration_utils.py index 9012271ec545..91926709897b 100644 --- a/src/transformers/distributed/configuration_utils.py +++ b/src/transformers/distributed/configuration_utils.py @@ -18,8 +18,6 @@ from dataclasses import asdict, dataclass from typing import Literal -from .utils import _get_torch_distributed_rank - @dataclass class DistributedConfig: @@ -32,7 +30,8 @@ class DistributedConfig: `WORLD_SIZE // (other_parallel_size)`. If `None` and no `tp_plan` is set, defaults to 1. tp_plan (`dict[str, str]` or `"auto"`, *optional*): Tensor parallel sharding plan. Pass `"auto"`, or leave as `None` when `tp_size` is set, to use the - model's predefined `base_model_tp_plan`. Pass a dictionary to override the predefined plan. + model's predefined `base_model_tp_plan`. Pass a dictionary to override individual rules of that plan; + unspecified rules are kept. enable_sequence_parallel (`bool`, *optional*, defaults to `False`): Reserved for sequence parallelism. Not wired up yet. enable_expert_parallel (`bool`, *optional*, defaults to `False`): @@ -49,6 +48,10 @@ class DistributedConfig: ep_size (`int`, *optional*): Number of devices owning distinct expert shards. Defaults to 1. Set it explicitly to enable EP. Model execution currently requires `ep_size=tp_size` when EP is enabled. + ep_plan (`dict[str, str]`, *optional*): + Expert parallel sharding plan. Leave as `None` to use the model's predefined `base_model_ep_plan`. Pass a + dictionary to override individual rules of that plan; unspecified rules are kept. Applied only when + `ep_size > 1`, and its rules take precedence over `tp_plan` rules for the same modules. """ tp_size: int | None = None @@ -60,6 +63,7 @@ class DistributedConfig: fsdp_mixed_precision: bool = False pp_size: int | None = None ep_size: int | None = None + ep_plan: dict[str, str] | None = None @property def efsdp_size(self) -> int: @@ -94,13 +98,12 @@ def _resolve_parallelism(self): if self.enable_expert_parallel and self.ep_size is None: self.ep_size = self.tp_size - if _get_torch_distributed_rank() == 0: - warnings.warn( - f"`enable_expert_parallel` without `ep_size` is deprecated and will be removed in v5.20. " - f"Use ep_size={self.ep_size} instead.", - FutureWarning, - stacklevel=4, - ) + warnings.warn( + f"`enable_expert_parallel` without `ep_size` is deprecated and will be removed in v5.20. " + f"Use ep_size={self.ep_size} instead.", + FutureWarning, + stacklevel=4, + ) if self.ep_size is None: self.ep_size = 1 @@ -109,6 +112,9 @@ def _resolve_parallelism(self): def _validate_mesh_config(self): """Validate mesh sizes before the model's expert plan is available.""" + if self.ep_plan is not None and not isinstance(self.ep_plan, dict): + raise ValueError("`ep_plan` must be a dictionary or None.") + if self.ep_size > 1: if self.ep_size % self.tp_size: raise ValueError("`ep_size` must be a multiple of `tp_size`.") diff --git a/src/transformers/distributed/mixin.py b/src/transformers/distributed/mixin.py index 75b799f57e15..4ae276547185 100644 --- a/src/transformers/distributed/mixin.py +++ b/src/transformers/distributed/mixin.py @@ -24,9 +24,10 @@ from .fsdp import apply_fully_sharded_data_parallelism, is_fsdp_managed_module from .pipeline_parallel import apply_pipeline_parallelism from .tensor_parallel import ( - _validate_tp_plan_styles, + _validate_parallel_plan_styles, apply_tensor_parallelism, gather_state_dict_for_save, + resolve_parallel_plans, ) from .utils import ( TransformersDeviceMesh, @@ -84,16 +85,13 @@ def init_parallel_plans(self) -> None: @property def tp_plan(self) -> dict[str, str]: - """The full tp plan for the model's modules.""" - if hasattr(self.config, "distributed_config") and self.config.distributed_config.enable_expert_parallel: - if not self._ep_plan: - raise ValueError( - f"Expert parallelism was requested (`ep_size > 1`), but " - f"`{self.__class__.__name__}` does not define an expert-parallel plan. Add a " - f"`base_model_ep_plan` to its config, or disable expert parallelism." - ) - return self._ep_plan - return self._tp_plan + """The full tensor parallel plan for the model's modules.""" + return self._tp_plan if self._tp_plan is not None else {} + + @property + def ep_plan(self) -> dict[str, str]: + """The full expert parallel plan for the model's modules, kept separate from `tp_plan`.""" + return self._ep_plan if self._ep_plan is not None else {} @property def fsdp_plan(self) -> dict[str, str]: @@ -111,7 +109,7 @@ def tp_plan(self, plan: dict[str, str] | None): if not isinstance(plan, dict): raise ValueError("Can only set a dictionary as `tp_plan`") - _validate_tp_plan_styles(plan) + _validate_parallel_plan_styles(plan) model_param_names = [name for name, _ in self.named_parameters()] for layer_pattern in plan.keys(): @@ -129,6 +127,17 @@ def tp_plan(self, plan: dict[str, str] | None): self._tp_plan = plan + @ep_plan.setter + def ep_plan(self, plan: dict[str, str] | None): + if plan is None: + self._ep_plan = {} + return + if not isinstance(plan, dict): + raise ValueError("Can only set a dictionary as `ep_plan`") + + _validate_parallel_plan_styles(plan) + self._ep_plan = plan + @pp_plan.setter def pp_plan(self, plan: dict[str, tuple[str, str]] | None): if plan is None: @@ -154,7 +163,6 @@ def prepare_distribute_model( if distributed_config.ep_size > 1 and distributed_config.ep_size != distributed_config.tp_size: raise ValueError( "All-reduce expert parallelism requires `ep_size=tp_size` and identical tokens per EP group." - "The token dispatch version allowing `ep_size != tp_size` is coming soon!" ) if distributed_config.tp_size == 1 and distributed_config.fsdp_size == 1 and distributed_config.pp_size == 1: @@ -176,29 +184,32 @@ def maybe_distribute_model( distributed_config: DistributedConfig | None, mesh_manager: TransformersDeviceMesh | None, ): - """Apply TP or FSDP2 after model init, before weight loading.""" - if mesh_manager is not None: - model.config.distributed_config = distributed_config - model._mesh_manager = mesh_manager - model._device_mesh = mesh_manager.get_mesh(("pp", "fsdp", "tp")) - model._tp_size = distributed_config.tp_size - model._fsdp_size = distributed_config.fsdp_size - - if distributed_config.pp_size > 1: - pp_mesh = mesh_manager.get_mesh("pp") - model = apply_pipeline_parallelism(model, pp_mesh) - - # Both may apply: the tensor/expert parallel plan shards across `tp` first, then FSDP2 - # shards every parameter (the `tp`-sharded ones included) across `fsdp`. - if distributed_config.tp_size > 1: - tp_mesh = mesh_manager.get_mesh("tp") - if isinstance(distributed_config.tp_plan, dict): - model.tp_plan = distributed_config.tp_plan - model = apply_tensor_parallelism(model, tp_mesh) - - if distributed_config.fsdp_size > 1: - fsdp_mesh = mesh_manager.get_mesh("fsdp") - model = apply_fully_sharded_data_parallelism(model, fsdp_mesh) + """Apply pipeline, tensor and expert parallelism, then FSDP2, after model init and before weight loading.""" + if mesh_manager is None: + return model + + model.config.distributed_config = distributed_config + model._mesh_manager = mesh_manager + model._device_mesh = mesh_manager.get_mesh(("pp", "fsdp", "tp")) + model._tp_size = distributed_config.tp_size + model._fsdp_size = distributed_config.fsdp_size + + # Resolve both plans before sharding anything: overrides are merged into `model.tp_plan` / `model.ep_plan`, + # and the experts named by the EP plan are removed from the TP plan so they are sharded once. + tp_plan, ep_plan = resolve_parallel_plans(model, distributed_config) + + if distributed_config.pp_size > 1: + model = apply_pipeline_parallelism(model, mesh_manager.get_mesh("pp")) + + tp_mesh = mesh_manager.get_mesh("tp") + if tp_plan: + model = apply_tensor_parallelism(model, tp_mesh, tp_plan) + if ep_plan: + # Legacy masked EP: the EP group is the TP group, every rank keeps every token. + model = apply_tensor_parallelism(model, tp_mesh, ep_plan) + + if distributed_config.fsdp_size > 1: + model = apply_fully_sharded_data_parallelism(model, mesh_manager.get_mesh("fsdp")) return model def should_save_on_this_rank(self, is_main_process: bool) -> bool: diff --git a/src/transformers/distributed/tensor_parallel.py b/src/transformers/distributed/tensor_parallel.py index 3151f6264a12..99973c5abf09 100644 --- a/src/transformers/distributed/tensor_parallel.py +++ b/src/transformers/distributed/tensor_parallel.py @@ -15,12 +15,21 @@ import contextlib import re +from fnmatch import fnmatchcase +from typing import TYPE_CHECKING from ..utils import logging from ..utils.generic import GeneralInterface from ..utils.import_utils import is_torch_available, is_torch_distributed_available +if TYPE_CHECKING: + from torch import nn + from torch.distributed.device_mesh import DeviceMesh + + from .configuration_utils import DistributedConfig + + logger = logging.get_logger(__name__) if is_torch_available(): @@ -71,21 +80,21 @@ def verify_tp_plan(expected_keys: list[str], tp_plan: dict[str, str] | None): logger.warning(f"The following layers were not sharded: {', '.join(unsharded_layers)}") -def _get_parameter_tp_plan(parameter_name: str, tp_plan: dict[str, str], is_weight=True) -> str | None: +def _get_parameter_plan(parameter_name: str, plan: dict[str, str], is_weight=True) -> str | None: """ - Get the TP style for a parameter from the TP plan. + Get the parallel style for a parameter or module from a TP or EP plan. - The TP plan is a dictionary that maps parameter names to TP styles. + The plan is a dictionary that maps parameter or module names to parallel styles. The parameter name can be a generic name with wildcards (e.g. "*.weight") or a specific name (e.g. "layer_1.weight"). The `is_weight` is important because for weights, we want to support `.weights` and `.bias` cases seamlessly! but not parent classes for `post_init` calls """ generic_param_name = replace_layer_number_by_wildcard(parameter_name) - if generic_param_name in tp_plan: - return tp_plan[generic_param_name] - elif is_weight and "." in generic_param_name and (module_name := generic_param_name.rsplit(".", 1)[0]) in tp_plan: - return tp_plan[module_name] + if generic_param_name in plan: + return plan[generic_param_name] + elif is_weight and "." in generic_param_name and (module_name := generic_param_name.rsplit(".", 1)[0]) in plan: + return plan[module_name] return None @@ -792,32 +801,77 @@ class ParallelInterface(GeneralInterface): ALL_PARALLEL_STYLES: ParallelInterface = ParallelInterface() -def _validate_tp_plan_styles(tp_plan: dict[str, str] | None) -> None: - unsupported_styles = {style for style in (tp_plan or {}).values() if style not in ALL_PARALLEL_STYLES} +def _validate_parallel_plan_styles(plan: dict[str, str] | None) -> None: + unsupported_styles = {style for style in (plan or {}).values() if style not in ALL_PARALLEL_STYLES} if unsupported_styles: raise ValueError( - f"Unsupported tensor parallel styles: {unsupported_styles}. " - f"Supported styles are {list(ALL_PARALLEL_STYLES.keys())}" + f"Unsupported parallel styles: {unsupported_styles}. Supported styles are {list(ALL_PARALLEL_STYLES.keys())}" + ) + + +def resolve_parallel_plans( + model: nn.Module, distributed_config: DistributedConfig +) -> tuple[dict[str, str], dict[str, str]]: + """Merge the `DistributedConfig` overrides into the model's plans and split them between TP and EP. + + Returns the TP plan to apply to the dense modules and the EP plan to apply to the experts. Each plan is empty + when its parallel size is 1. EP owns every module it names, so TP rules for those modules and their children + are dropped: expert weights are sharded once, by the EP plan. + """ + # Reject invalid paths before merging, e.g. "layers.*" when the model uses "model.layers.*". + layer_names = {name for name, _ in model.named_modules()} | {name for name, _ in model.named_parameters()} + layer_names |= {replace_layer_number_by_wildcard(name) for name in layer_names} + for plan_name in ("tp_plan", "ep_plan"): + override = getattr(distributed_config, plan_name) + if isinstance(override, dict): + valid_names = layer_names | set(getattr(model, plan_name)) + for pattern in override: + if pattern not in valid_names: + raise ValueError( + f"The `{plan_name}` pattern {pattern!r} does not match any module, parameter, " + f"or existing plan entry in {type(model).__name__}. " + "Check the full path, including any 'model.' prefix." + ) + + if isinstance(distributed_config.tp_plan, dict): + model._tp_plan = model.tp_plan | distributed_config.tp_plan + if isinstance(distributed_config.ep_plan, dict): + model._ep_plan = model.ep_plan | distributed_config.ep_plan + + tp_plan = dict(model.tp_plan) if distributed_config.tp_size > 1 else {} + ep_plan = dict(model.ep_plan) if distributed_config.ep_size > 1 else {} + if distributed_config.ep_size > 1 and not ep_plan: + raise ValueError( + f"Expert parallelism was requested (`ep_size={distributed_config.ep_size}`), but `{type(model).__name__}` " + "does not define an expert-parallel plan. Pass `ep_plan` in `DistributedConfig`, add a " + "`base_model_ep_plan` to the model's config, or disable expert parallelism." ) + def is_expert_path(name: str) -> bool: + return any(fnmatchcase(name, path) or fnmatchcase(name, path + ".*") for path in ep_plan) + + tp_plan = {name: style for name, style in tp_plan.items() if not is_expert_path(name)} + _validate_parallel_plan_styles(tp_plan) + _validate_parallel_plan_styles(ep_plan) + return tp_plan, ep_plan -def apply_tensor_parallelism(model, tp_mesh): - """DTensor backend: shard params as placeholders and install TP forward hooks.""" - _validate_tp_plan_styles(model.tp_plan) +def apply_tensor_parallelism(model: nn.Module, tp_mesh: DeviceMesh, plan: dict[str, str] | None = None): + plan = model.tp_plan if plan is None else plan + _validate_parallel_plan_styles(plan) for name, module in model.named_modules(): # Create DTensor placeholders so the loader knows which shard belongs to this rank. for p_name, _ in list(module.named_parameters(recurse=False)): full = f"{name}.{p_name}" if name else p_name - style_name = _get_parameter_tp_plan(parameter_name=full, tp_plan=model.tp_plan, is_weight=True) + style_name = _get_parameter_plan(parameter_name=full, plan=plan, is_weight=True) if style_name is not None and style_name in ALL_PARALLEL_STYLES: style = ALL_PARALLEL_STYLES[style_name] style.validate_param(module, p_name, tp_mesh, parameter_name=full) style.shard_param(module, p_name, tp_mesh) - # Install the input/output transforms required by this module's TP style. - style_name = _get_parameter_tp_plan(parameter_name=name, tp_plan=model.tp_plan, is_weight=False) + # Install the input/output transforms required by this module's style. + style_name = _get_parameter_plan(parameter_name=name, plan=plan, is_weight=False) if style_name is not None and style_name in ALL_PARALLEL_STYLES: if style_name == "mla_kv_a_proj": # MLA needs to know the qk_rope_head_dim to split the projection output into KV and RoPE parts. diff --git a/src/transformers/modeling_utils.py b/src/transformers/modeling_utils.py index 07291b2870f6..ccb9eb334ad4 100644 --- a/src/transformers/modeling_utils.py +++ b/src/transformers/modeling_utils.py @@ -57,7 +57,7 @@ from .distributed import DistributedConfig from .distributed.mixin import DistributedMixin from .distributed.sharding_utils import _dtensor_from_local_like -from .distributed.tensor_parallel import _get_parameter_tp_plan, verify_tp_plan +from .distributed.tensor_parallel import _get_parameter_plan, verify_tp_plan from .distributed.utils import ( _get_torch_distributed_world_size, _is_torch_distributed_initialized, @@ -5013,7 +5013,7 @@ def get_total_byte_count( param_byte_count = param.numel() * dtype_size if len(tp_plan) > 0: - is_part_of_plan = _get_parameter_tp_plan(param_name, tp_plan, is_weight=True) is not None + is_part_of_plan = _get_parameter_plan(param_name, tp_plan, is_weight=True) is not None param_byte_count //= _get_torch_distributed_world_size() if is_part_of_plan else 1 total_byte_count[device] += param_byte_count diff --git a/tests/tensor_parallel/test_tensor_parallel.py b/tests/tensor_parallel/test_tensor_parallel.py index c83741361fd5..c39930ef2307 100644 --- a/tests/tensor_parallel/test_tensor_parallel.py +++ b/tests/tensor_parallel/test_tensor_parallel.py @@ -16,8 +16,9 @@ import torch -from transformers import AutoModelForCausalLM +from transformers import AutoModelForCausalLM, Qwen3MoeConfig, Qwen3MoeForCausalLM, Qwen3MoeModel from transformers.distributed import tensor_parallel +from transformers.distributed.configuration_utils import DistributedConfig from transformers.distributed.sharding_utils import DtensorShardOperation from transformers.distributed.tensor_parallel import ( ALL_PARALLEL_STYLES, @@ -26,7 +27,216 @@ PackedRowwiseParallel, RowwiseParallel, ) -from transformers.testing_utils import TestCasePlus, is_tensor_parallel_test +from transformers.testing_utils import TestCasePlus, is_tensor_parallel_test, require_torch + + +# Qwen3 MoE's predefined plans, as resolved on `Qwen3MoeModel` (no `model.` prefix). +DENSE_TP_PLAN = { + "layers.*.self_attn.q_proj": "colwise", + "layers.*.self_attn.k_proj": "colwise", + "layers.*.self_attn.v_proj": "colwise", + "layers.*.self_attn.q_norm": "replicated_with_grad_allreduce", + "layers.*.self_attn.k_norm": "replicated_with_grad_allreduce", + "layers.*.self_attn.o_proj": "rowwise", + "layers.*.mlp.gate_proj": "colwise", + "layers.*.mlp.up_proj": "colwise", + "layers.*.mlp.down_proj": "rowwise", +} +EXPERT_TP_PLAN = { + "layers.*.mlp.experts.gate_up_proj": "packed_colwise", + "layers.*.mlp.experts.down_proj": "rowwise", + "layers.*.mlp.experts": "moe_tp_experts", +} +EP_PLAN = { + "layers.*.mlp.gate": "ep_router", + "layers.*.mlp.experts.gate_up_proj": "grouped_gemm", + "layers.*.mlp.experts.down_proj": "grouped_gemm", + "layers.*.mlp.experts": "moe_tp_experts", +} + + +@require_torch +class TestParallelPlanResolution(TestCasePlus): + def setUp(self): + super().setUp() + self.config = Qwen3MoeConfig( + vocab_size=32, + hidden_size=16, + intermediate_size=32, + moe_intermediate_size=8, + num_hidden_layers=1, + num_attention_heads=4, + num_key_value_heads=4, + head_dim=4, + num_experts=4, + num_experts_per_tok=2, + ) + with torch.device("meta"): + self.model = Qwen3MoeModel(self.config) + + def test_ep_plan_setter(self): + self.model.ep_plan = None + self.assertEqual(self.model.ep_plan, {}) + with self.assertRaisesRegex(ValueError, "Can only set a dictionary"): + self.model.ep_plan = "auto" + with self.assertRaisesRegex(ValueError, "Unsupported parallel styles"): + self.model.ep_plan = {"layers.*.mlp.experts": "invalid_style"} + self.model.ep_plan = EP_PLAN + self.assertEqual(self.model.ep_plan, EP_PLAN) + + def test_disabled_parallelism_has_no_plans(self): + for config in (DistributedConfig(), DistributedConfig(fsdp_size=8), DistributedConfig(pp_size=2)): + with self.subTest(config=config): + self.assertEqual(tensor_parallel.resolve_parallel_plans(self.model, config), ({}, {})) + + def test_tp_only_keeps_experts_in_tp_plan(self): + tp_plan, ep_plan = tensor_parallel.resolve_parallel_plans(self.model, DistributedConfig(tp_size=4)) + self.assertEqual(tp_plan, DENSE_TP_PLAN | EXPERT_TP_PLAN) + self.assertEqual(ep_plan, {}) + + def test_ep_takes_experts_and_router_out_of_tp_plan(self): + for config in ( + DistributedConfig(tp_size=4, ep_size=4), + DistributedConfig(tp_size=2, fsdp_size=2, ep_size=2), + ): + with self.subTest(config=config): + tp_plan, ep_plan = tensor_parallel.resolve_parallel_plans(self.model, config) + self.assertEqual(tp_plan, DENSE_TP_PLAN) + self.assertEqual(ep_plan, EP_PLAN) + + def test_legacy_flag_is_an_alias_for_ep_size(self): + with warnings.catch_warnings(record=True) as caught: + warnings.simplefilter("always") + legacy = DistributedConfig(tp_size=4, enable_expert_parallel=True) + self.assertEqual([w.category for w in caught], [FutureWarning]) + self.assertIn("Use ep_size=4 instead", str(caught[0].message)) + with warnings.catch_warnings(record=True) as caught: + warnings.simplefilter("always") + explicit = DistributedConfig(tp_size=4, ep_size=4) + disabled = DistributedConfig(tp_size=4, ep_size=1, enable_expert_parallel=True) + self.assertEqual(caught, []) + self.assertEqual(legacy, explicit) + self.assertFalse(disabled.enable_expert_parallel) + self.assertEqual( + tensor_parallel.resolve_parallel_plans(self.model, legacy), + tensor_parallel.resolve_parallel_plans(self.model, explicit), + ) + + def test_ep_plan_is_a_dict_and_round_trips(self): + with self.assertRaisesRegex(ValueError, "`ep_plan` must be a dictionary or None"): + DistributedConfig(tp_size=4, ep_size=4, ep_plan="auto") + config = DistributedConfig(tp_size=4, ep_size=4, ep_plan={"layers.*.mlp.gate": "ep_router"}) + self.assertEqual(config.to_dict()["ep_plan"], {"layers.*.mlp.gate": "ep_router"}) + self.assertEqual(DistributedConfig.from_dict(config.to_dict()), config) + + def test_overrides_merge_into_the_predefined_plans(self): + config = DistributedConfig( + tp_size=4, + ep_size=4, + tp_plan={"layers.*.self_attn.q_proj": "colwise_rep"}, + ep_plan={"layers.*.mlp.experts.down_proj": "rowwise"}, + ) + tp_plan, ep_plan = tensor_parallel.resolve_parallel_plans(self.model, config) + self.assertEqual(tp_plan, DENSE_TP_PLAN | {"layers.*.self_attn.q_proj": "colwise_rep"}) + self.assertEqual(ep_plan, EP_PLAN | {"layers.*.mlp.experts.down_proj": "rowwise"}) + # The merged plans are stored on the model, the config defaults are untouched. + self.assertEqual(self.model.tp_plan["layers.*.self_attn.q_proj"], "colwise_rep") + self.assertEqual(self.model.ep_plan["layers.*.mlp.experts.down_proj"], "rowwise") + self.assertEqual(self.model.config.base_model_tp_plan["layers.*.self_attn.q_proj"], "colwise") + self.assertEqual(self.model.config.base_model_ep_plan["layers.*.mlp.experts.down_proj"], "grouped_gemm") + # The overrides are not rewritten with the merged plans. + self.assertEqual(config.tp_plan, {"layers.*.self_attn.q_proj": "colwise_rep"}) + self.assertEqual(config.ep_plan, {"layers.*.mlp.experts.down_proj": "rowwise"}) + # The merged EP plan stays on the model but is not applied while EP is disabled. + tp_plan, ep_plan = tensor_parallel.resolve_parallel_plans(self.model, DistributedConfig(tp_size=4)) + self.assertEqual(tp_plan, DENSE_TP_PLAN | EXPERT_TP_PLAN | {"layers.*.self_attn.q_proj": "colwise_rep"}) + self.assertEqual(ep_plan, {}) + self.assertEqual(self.model.ep_plan["layers.*.mlp.experts.down_proj"], "rowwise") + + def test_ep_rules_take_precedence_over_tp_rules_for_the_same_modules(self): + config = DistributedConfig( + tp_size=4, + ep_size=4, + tp_plan={"layers.*.mlp.experts.gate_up_proj": "packed_rowwise", "layers.*.mlp.gate": "colwise"}, + ) + tp_plan, ep_plan = tensor_parallel.resolve_parallel_plans(self.model, config) + self.assertEqual(tp_plan, DENSE_TP_PLAN) + self.assertEqual(ep_plan, EP_PLAN) + # The custom TP rules are kept on the model and apply as soon as EP is disabled. + tp_plan, ep_plan = tensor_parallel.resolve_parallel_plans(self.model, DistributedConfig(tp_size=4)) + self.assertEqual(tp_plan["layers.*.mlp.experts.gate_up_proj"], "packed_rowwise") + self.assertEqual(tp_plan["layers.*.mlp.gate"], "colwise") + self.assertEqual(ep_plan, {}) + + def test_ep_requires_an_expert_plan(self): + self.model.ep_plan = None + with self.assertRaisesRegex(ValueError, "does not define an expert-parallel plan"): + tensor_parallel.resolve_parallel_plans(self.model, DistributedConfig(tp_size=4, ep_size=4)) + config = DistributedConfig(tp_size=4, ep_size=4, ep_plan=EP_PLAN) + self.assertEqual(tensor_parallel.resolve_parallel_plans(self.model, config), (DENSE_TP_PLAN, EP_PLAN)) + + def test_unmatched_override_keys_raise_without_changing_plans(self): + original_tp_plan, original_ep_plan = self.model.tp_plan.copy(), self.model.ep_plan.copy() + for plan_name in ("tp_plan", "ep_plan"): + for key in ("layers.*.mlp.experst", "layers.*.mlp.experts.missing_weight", "model.layers.*.mlp.experts"): + with self.subTest(plan_name=plan_name, key=key): + config = DistributedConfig(tp_size=4, ep_size=4, **{plan_name: {key: "grouped_gemm"}}) + with self.assertRaisesRegex(ValueError, f"The `{plan_name}` pattern .* does not match") as error: + tensor_parallel.resolve_parallel_plans(self.model, config) + self.assertIn(key, str(error.exception)) + self.assertIn("Qwen3MoeModel", str(error.exception)) + self.assertEqual(self.model.tp_plan, original_tp_plan) + self.assertEqual(self.model.ep_plan, original_ep_plan) + + def test_override_keys_can_match_modules_parameters_or_existing_plan_keys(self): + for plan_name in ("tp_plan", "ep_plan"): + # `gate_proj` is in the predefined TP plan even though this MoE model has no such module. + for key in ("layers.*.mlp", "layers.0.self_attn.q_proj.weight", "layers.*.mlp.gate_proj"): + with self.subTest(plan_name=plan_name, key=key): + original = getattr(self.model, plan_name).copy() + if plan_name == "ep_plan": + self.model.ep_plan = original | {"layers.*.mlp.gate_proj": "colwise"} + config = DistributedConfig(tp_size=4, **{plan_name: {key: "colwise_rep"}}) + tensor_parallel.resolve_parallel_plans(self.model, config) + self.assertEqual(getattr(self.model, plan_name)[key], "colwise_rep") + setattr(self.model, plan_name, original) + + def test_head_model_overrides_need_the_model_prefix(self): + with torch.device("meta"): + model = Qwen3MoeForCausalLM(self.config) + config = DistributedConfig(tp_size=4, ep_size=4, ep_plan={"layers.*.mlp.gate": "ep_router"}) + with self.assertRaisesRegex(ValueError, "including any 'model.' prefix"): + tensor_parallel.resolve_parallel_plans(model, config) + + config = DistributedConfig( + tp_size=4, + ep_size=4, + tp_plan={"model.layers.*.self_attn.q_proj": "colwise_rep"}, + ep_plan={"model.layers.*.mlp.gate": "ep_router"}, + ) + tp_plan, ep_plan = tensor_parallel.resolve_parallel_plans(model, config) + expected_tp_plan = {f"model.{k}": v for k, v in DENSE_TP_PLAN.items()} | {"lm_head": "colwise_gather_output"} + self.assertEqual(tp_plan, expected_tp_plan | config.tp_plan) + self.assertEqual(ep_plan, {f"model.{k}": v for k, v in EP_PLAN.items()}) + + def test_masked_ep_shards_and_installs_hooks_on_the_tp_mesh(self): + tp_mesh = object() + _, ep_plan = tensor_parallel.resolve_parallel_plans(self.model, DistributedConfig(tp_size=4, ep_size=4)) + experts, router = self.model.layers[0].mlp.experts, self.model.layers[0].mlp.gate + with ( + patch.object(ALL_PARALLEL_STYLES["grouped_gemm"], "validate_param") as validate, + patch.object(ALL_PARALLEL_STYLES["grouped_gemm"], "shard_param") as shard, + patch.object(ALL_PARALLEL_STYLES["moe_tp_experts"], "install_forward") as install_experts, + patch.object(ALL_PARALLEL_STYLES["ep_router"], "install_forward") as install_router, + ): + result = tensor_parallel.apply_tensor_parallelism(self.model, tp_mesh, ep_plan) + self.assertIs(result, self.model) + self.assertEqual(shard.call_count, 2) + for name in ("gate_up_proj", "down_proj"): + validate.assert_any_call(experts, name, tp_mesh, parameter_name=f"layers.0.mlp.experts.{name}") + shard.assert_any_call(experts, name, tp_mesh) + install_experts.assert_called_once_with(experts, tp_mesh) + install_router.assert_called_once_with(router, tp_mesh) @is_tensor_parallel_test diff --git a/tests/test_modeling_common.py b/tests/test_modeling_common.py index 99e9456487ae..38df5eca748d 100644 --- a/tests/test_modeling_common.py +++ b/tests/test_modeling_common.py @@ -136,7 +136,7 @@ from torch import nn from transformers import MODEL_MAPPING - from transformers.distributed.tensor_parallel import _get_parameter_tp_plan + from transformers.distributed.tensor_parallel import _get_parameter_plan from transformers.integrations.accelerate import compute_module_sizes from transformers.modeling_utils import load_state_dict from transformers.pytorch_utils import id_tensor_storage @@ -4874,10 +4874,9 @@ def test_tp_plan_matches_params(self): for pattern in tp_plan: # Check if this given pattern matches any param or module (the value attributed to the pattern does not matter) pattern_usage[pattern] = any( - _get_parameter_tp_plan(param, {pattern: ""}, is_weight=True) is not None for param in param_names + _get_parameter_plan(param, {pattern: ""}, is_weight=True) is not None for param in param_names ) or any( - _get_parameter_tp_plan(module, {pattern: ""}, is_weight=False) is not None - for module in module_names + _get_parameter_plan(module, {pattern: ""}, is_weight=False) is not None for module in module_names ) unused_entries = {k for k, v in pattern_usage.items() if not v} diff --git a/tests/test_tensor_parallel_mixin.py b/tests/test_tensor_parallel_mixin.py index 38f78a0163f3..c82e1de4671a 100644 --- a/tests/test_tensor_parallel_mixin.py +++ b/tests/test_tensor_parallel_mixin.py @@ -20,7 +20,7 @@ from transformers import TorchAoConfig, set_seed from transformers.distributed.configuration_utils import DistributedConfig -from transformers.distributed.tensor_parallel import _get_parameter_tp_plan +from transformers.distributed.tensor_parallel import _get_parameter_plan from transformers.testing_utils import ( is_tensor_parallel_test, is_torch_available, @@ -182,7 +182,7 @@ def _verify_tp_sharding(rank, model_tp, model_ref): # Verify sharding is correct for dim in range(param.ndim): if param.size(dim) != param_full.size(dim): - param_plan = _get_parameter_tp_plan(name, model_tp._tp_plan, is_weight=True) + param_plan = _get_parameter_plan(name, model_tp._tp_plan, is_weight=True) if param_plan in ("packed_colwise", "packed_rowwise"): expected_size = param_full.size(dim) // world_size assert param.size(dim) == expected_size, ( @@ -268,7 +268,7 @@ def _test_tp_backward_impl(rank, model_path, model_class, atol, rtol): if grad.shape != grad_tp.shape: for dim in range(grad.ndim): if grad.size(dim) != grad_tp.size(dim): - param_plan = _get_parameter_tp_plan(name, model_tp._tp_plan, is_weight=True) + param_plan = _get_parameter_plan(name, model_tp._tp_plan, is_weight=True) if param_plan in ("packed_colwise", "packed_rowwise"): # interleaved slicing grad = get_packed_grad_shard(grad, world_size, rank, dim) @@ -392,6 +392,7 @@ def _test_tp_generation_quantized_impl(_rank, model_path, model_class, max_new_t def _load_ep_and_reference_models(model_path, model_class): """Load EP model and non-EP reference model for comparison.""" + # All-reduce EP: every rank sees the same tokens, so TP and EP span the same ranks. model_ep = model_class.from_pretrained( model_path, distributed_config=DistributedConfig(tp_size=dist.get_world_size(), ep_size=dist.get_world_size()), From a4f11e3a69ffa1a11c13576c82719003f18d7a40 Mon Sep 17 00:00:00 2001 From: 3outeille Date: Tue, 22 Sep 2026 09:22:27 +0000 Subject: [PATCH 49/86] Match expert paths with a regex and keep only expert rules under token dispatch Replace the fnmatch check in resolve_parallel_plans with a plan-pattern regex so plan keys are matched literally except for `*`. When the EP plan uses `ep_dispatch_experts`, drop the router masking rules from the EP plan and only take the expert modules and their parameters out of the TP plan. --- .../distributed/tensor_parallel.py | 24 +++++++++++++++---- 1 file changed, 19 insertions(+), 5 deletions(-) diff --git a/src/transformers/distributed/tensor_parallel.py b/src/transformers/distributed/tensor_parallel.py index 99973c5abf09..e5296f72d126 100644 --- a/src/transformers/distributed/tensor_parallel.py +++ b/src/transformers/distributed/tensor_parallel.py @@ -15,7 +15,6 @@ import contextlib import re -from fnmatch import fnmatchcase from typing import TYPE_CHECKING from ..utils import logging @@ -51,6 +50,14 @@ def replace_layer_number_by_wildcard(name: str) -> str: return re.sub(r"\.\d+(\.|$)", lambda m: ".*" + m.group(1), name) +def _plan_pattern_to_regex(pattern: str) -> str: + """ + Translate a plan key into a regex, where `*` stands for any run of characters (typically a layer index, e.g. + `"model.layers.*.mlp.experts"`). Every other character is matched literally. + """ + return ".*".join(re.escape(part) for part in pattern.split("*")) + + def verify_tp_plan(expected_keys: list[str], tp_plan: dict[str, str] | None): """ Verify the TP plan of the model, log a warning if the layers that were not sharded and the rules that were not applied. @@ -847,10 +854,17 @@ def resolve_parallel_plans( "`base_model_ep_plan` to the model's config, or disable expert parallelism." ) - def is_expert_path(name: str) -> bool: - return any(fnmatchcase(name, path) or fnmatchcase(name, path + ".*") for path in ep_plan) - - tp_plan = {name: style for name, style in tp_plan.items() if not is_expert_path(name)} + def is_expert_path(name: str, paths: list[str]) -> bool: + # An EP path also owns its children, e.g. `...experts.gate_up_proj` under `...experts`. + return any(re.fullmatch(rf"{_plan_pattern_to_regex(path)}(\..*)?", name) for path in paths) + + expert_paths = list(ep_plan) + if "ep_dispatch_experts" in ep_plan.values(): + # Dispatch finds each expert's owner from the global expert ids, so the router masking hooks + # (`ep_router`) must not run: keep only the expert modules and their parameter rules. + expert_paths = [name for name, style in ep_plan.items() if style in ("moe_tp_experts", "ep_dispatch_experts")] + ep_plan = {name: style for name, style in ep_plan.items() if is_expert_path(name, expert_paths)} + tp_plan = {name: style for name, style in tp_plan.items() if not is_expert_path(name, expert_paths)} _validate_parallel_plan_styles(tp_plan) _validate_parallel_plan_styles(ep_plan) return tp_plan, ep_plan From a5cd182bd3e64377a2af76f6f8483a7725a4224f Mon Sep 17 00:00:00 2001 From: 3outeille Date: Wed, 30 Sep 2026 17:12:44 +0000 Subject: [PATCH 50/86] revert change on warning --- .../distributed/configuration_utils.py | 15 +++++++++------ 1 file changed, 9 insertions(+), 6 deletions(-) diff --git a/src/transformers/distributed/configuration_utils.py b/src/transformers/distributed/configuration_utils.py index 91926709897b..77800f12a64d 100644 --- a/src/transformers/distributed/configuration_utils.py +++ b/src/transformers/distributed/configuration_utils.py @@ -18,6 +18,8 @@ from dataclasses import asdict, dataclass from typing import Literal +from .utils import _get_torch_distributed_rank + @dataclass class DistributedConfig: @@ -98,12 +100,13 @@ def _resolve_parallelism(self): if self.enable_expert_parallel and self.ep_size is None: self.ep_size = self.tp_size - warnings.warn( - f"`enable_expert_parallel` without `ep_size` is deprecated and will be removed in v5.20. " - f"Use ep_size={self.ep_size} instead.", - FutureWarning, - stacklevel=4, - ) + if _get_torch_distributed_rank() == 0: + warnings.warn( + f"`enable_expert_parallel` without `ep_size` is deprecated and will be removed in v5.20. " + f"Use ep_size={self.ep_size} instead.", + FutureWarning, + stacklevel=4, + ) if self.ep_size is None: self.ep_size = 1 From 3d45692b242d576f78f87c7d7fd8f164642a3b03 Mon Sep 17 00:00:00 2001 From: Ferdinand Mom <47445085+3outeille@users.noreply.github.com> Date: Thu, 1 Oct 2026 05:34:19 +0900 Subject: [PATCH 51/86] Apply suggestion from @ArthurZucker Co-authored-by: Arthur <48595927+ArthurZucker@users.noreply.github.com> --- .../distributed/tensor_parallel.py | 20 +++++-------------- 1 file changed, 5 insertions(+), 15 deletions(-) diff --git a/src/transformers/distributed/tensor_parallel.py b/src/transformers/distributed/tensor_parallel.py index e5296f72d126..2a57300287c6 100644 --- a/src/transformers/distributed/tensor_parallel.py +++ b/src/transformers/distributed/tensor_parallel.py @@ -826,24 +826,14 @@ def resolve_parallel_plans( are dropped: expert weights are sharded once, by the EP plan. """ # Reject invalid paths before merging, e.g. "layers.*" when the model uses "model.layers.*". - layer_names = {name for name, _ in model.named_modules()} | {name for name, _ in model.named_parameters()} - layer_names |= {replace_layer_number_by_wildcard(name) for name in layer_names} + names = {replace_layer_number_by_wildcard(n) for n, _ in chain(model.named_modules(), model.named_parameters())} for plan_name in ("tp_plan", "ep_plan"): override = getattr(distributed_config, plan_name) if isinstance(override, dict): - valid_names = layer_names | set(getattr(model, plan_name)) - for pattern in override: - if pattern not in valid_names: - raise ValueError( - f"The `{plan_name}` pattern {pattern!r} does not match any module, parameter, " - f"or existing plan entry in {type(model).__name__}. " - "Check the full path, including any 'model.' prefix." - ) - - if isinstance(distributed_config.tp_plan, dict): - model._tp_plan = model.tp_plan | distributed_config.tp_plan - if isinstance(distributed_config.ep_plan, dict): - model._ep_plan = model.ep_plan | distributed_config.ep_plan + plan = getattr(model, plan_name) + if unknown := override.keys() - names - plan.keys(): + raise ValueError(f"`{plan_name}` keys {sorted(unknown)} match nothing in {type(model).__name__}.") + setattr(model, f"_{plan_name}", plan | override) tp_plan = dict(model.tp_plan) if distributed_config.tp_size > 1 else {} ep_plan = dict(model.ep_plan) if distributed_config.ep_size > 1 else {} From 04752eb4cffbef4159d0753510913ae89177eb72 Mon Sep 17 00:00:00 2001 From: 3outeille Date: Wed, 30 Sep 2026 21:21:54 +0000 Subject: [PATCH 52/86] cleaning --- .../distributed/tensor_parallel.py | 27 ++++++------------- tests/tensor_parallel/test_tensor_parallel.py | 6 ++--- 2 files changed, 11 insertions(+), 22 deletions(-) diff --git a/src/transformers/distributed/tensor_parallel.py b/src/transformers/distributed/tensor_parallel.py index 2a57300287c6..5864520fac99 100644 --- a/src/transformers/distributed/tensor_parallel.py +++ b/src/transformers/distributed/tensor_parallel.py @@ -15,6 +15,7 @@ import contextlib import re +from itertools import chain from typing import TYPE_CHECKING from ..utils import logging @@ -50,14 +51,6 @@ def replace_layer_number_by_wildcard(name: str) -> str: return re.sub(r"\.\d+(\.|$)", lambda m: ".*" + m.group(1), name) -def _plan_pattern_to_regex(pattern: str) -> str: - """ - Translate a plan key into a regex, where `*` stands for any run of characters (typically a layer index, e.g. - `"model.layers.*.mlp.experts"`). Every other character is matched literally. - """ - return ".*".join(re.escape(part) for part in pattern.split("*")) - - def verify_tp_plan(expected_keys: list[str], tp_plan: dict[str, str] | None): """ Verify the TP plan of the model, log a warning if the layers that were not sharded and the rules that were not applied. @@ -844,17 +837,13 @@ def resolve_parallel_plans( "`base_model_ep_plan` to the model's config, or disable expert parallelism." ) - def is_expert_path(name: str, paths: list[str]) -> bool: - # An EP path also owns its children, e.g. `...experts.gate_up_proj` under `...experts`. - return any(re.fullmatch(rf"{_plan_pattern_to_regex(path)}(\..*)?", name) for path in paths) - - expert_paths = list(ep_plan) - if "ep_dispatch_experts" in ep_plan.values(): - # Dispatch finds each expert's owner from the global expert ids, so the router masking hooks - # (`ep_router`) must not run: keep only the expert modules and their parameter rules. - expert_paths = [name for name, style in ep_plan.items() if style in ("moe_tp_experts", "ep_dispatch_experts")] - ep_plan = {name: style for name, style in ep_plan.items() if is_expert_path(name, expert_paths)} - tp_plan = {name: style for name, style in tp_plan.items() if not is_expert_path(name, expert_paths)} + if "ep_dispatch_experts" in ep_plan.values() and "ep_router" in ep_plan.values(): + raise ValueError("`ep_dispatch_experts` routes tokens itself; remove the `ep_router` rules from `ep_plan`.") + + # EP rules take precedence: drop TP rules on EP modules and their children. + is_expert = re.compile(rf"(?:{'|'.join(map(re.escape, ep_plan))})(?:\..+)?").fullmatch + tp_plan = {name: style for name, style in tp_plan.items() if not is_expert(name)} + _validate_parallel_plan_styles(tp_plan) _validate_parallel_plan_styles(ep_plan) return tp_plan, ep_plan diff --git a/tests/tensor_parallel/test_tensor_parallel.py b/tests/tensor_parallel/test_tensor_parallel.py index c39930ef2307..a265cbd4087f 100644 --- a/tests/tensor_parallel/test_tensor_parallel.py +++ b/tests/tensor_parallel/test_tensor_parallel.py @@ -181,7 +181,7 @@ def test_unmatched_override_keys_raise_without_changing_plans(self): for key in ("layers.*.mlp.experst", "layers.*.mlp.experts.missing_weight", "model.layers.*.mlp.experts"): with self.subTest(plan_name=plan_name, key=key): config = DistributedConfig(tp_size=4, ep_size=4, **{plan_name: {key: "grouped_gemm"}}) - with self.assertRaisesRegex(ValueError, f"The `{plan_name}` pattern .* does not match") as error: + with self.assertRaisesRegex(ValueError, f"`{plan_name}` keys .* match nothing in") as error: tensor_parallel.resolve_parallel_plans(self.model, config) self.assertIn(key, str(error.exception)) self.assertIn("Qwen3MoeModel", str(error.exception)) @@ -191,7 +191,7 @@ def test_unmatched_override_keys_raise_without_changing_plans(self): def test_override_keys_can_match_modules_parameters_or_existing_plan_keys(self): for plan_name in ("tp_plan", "ep_plan"): # `gate_proj` is in the predefined TP plan even though this MoE model has no such module. - for key in ("layers.*.mlp", "layers.0.self_attn.q_proj.weight", "layers.*.mlp.gate_proj"): + for key in ("layers.*.mlp", "layers.*.self_attn.q_proj.weight", "layers.*.mlp.gate_proj"): with self.subTest(plan_name=plan_name, key=key): original = getattr(self.model, plan_name).copy() if plan_name == "ep_plan": @@ -205,7 +205,7 @@ def test_head_model_overrides_need_the_model_prefix(self): with torch.device("meta"): model = Qwen3MoeForCausalLM(self.config) config = DistributedConfig(tp_size=4, ep_size=4, ep_plan={"layers.*.mlp.gate": "ep_router"}) - with self.assertRaisesRegex(ValueError, "including any 'model.' prefix"): + with self.assertRaisesRegex(ValueError, "match nothing in Qwen3MoeForCausalLM"): tensor_parallel.resolve_parallel_plans(model, config) config = DistributedConfig( From c41435b0186a2605e1e23c5d5deeb2e2f4f062cf Mon Sep 17 00:00:00 2001 From: 3outeille Date: Fri, 18 Sep 2026 16:59:48 +0000 Subject: [PATCH 53/86] Fix tied embeddings for models with an EP-only base plan --- src/transformers/configuration_utils.py | 6 ++++-- 1 file changed, 4 insertions(+), 2 deletions(-) diff --git a/src/transformers/configuration_utils.py b/src/transformers/configuration_utils.py index 69c89e0422b3..ef53a3117477 100755 --- a/src/transformers/configuration_utils.py +++ b/src/transformers/configuration_utils.py @@ -388,9 +388,11 @@ def __post_init__(self, **kwargs): self.per_layer_config = per_layer_config # TODO: to support models whose input embedding module is not named `embed_tokens` (e.g. GPT-NeoX's `embed_in`). - if getattr(self, "tie_word_embeddings", False) and self.base_model_tp_plan is not None: + if getattr(self, "tie_word_embeddings", False) and ( + self.base_model_tp_plan is not None or self.base_model_ep_plan is not None + ): self.base_model_tp_plan = { - **self.base_model_tp_plan, + **(self.base_model_tp_plan or {}), "embed_tokens": "embedding_rowwise", } From 8dc8333b252871a02c30e3463bba9c582f31ad29 Mon Sep 17 00:00:00 2001 From: 3outeille Date: Fri, 18 Sep 2026 16:59:48 +0000 Subject: [PATCH 54/86] [distributed] Add expert-parallel token dispatch, default for Qwen3 MoE --- docs/source/en/expert_parallelism.md | 108 +++++++++- .../distributed/configuration_utils.py | 24 ++- src/transformers/distributed/fsdp.py | 23 +- src/transformers/distributed/mixin.py | 33 ++- .../distributed/tensor_parallel.py | 199 +++++++++++++++++- .../qwen3_moe/configuration_qwen3_moe.py | 10 +- tests/tensor_parallel/test_tensor_parallel.py | 43 ++-- 7 files changed, 390 insertions(+), 50 deletions(-) diff --git a/docs/source/en/expert_parallelism.md b/docs/source/en/expert_parallelism.md index 785c7f7be2f2..61d646c56f23 100644 --- a/docs/source/en/expert_parallelism.md +++ b/docs/source/en/expert_parallelism.md @@ -19,7 +19,7 @@ rendered properly in your Markdown viewer. ## DistributedConfig -Enable expert parallelism with the [`DistributedConfig`] class and the `ep_size` argument. The current all-reduce implementation requires `ep_size=tp_size`, so every rank in an expert group receives the same tokens. +Enable expert parallelism with the [`DistributedConfig`] class and the `ep_size` argument. Most models route with masking and all-reduce, which requires `ep_size=tp_size` so every rank in an expert group receives the same tokens. Models whose plan uses [token dispatch](#token-dispatch), such as Qwen3 MoE, can set `ep_size` independently of `tp_size`. ```py import os @@ -43,10 +43,17 @@ Each MoE model defines two plans in its config: `base_model_tp_plan` for the den `tp_plan` is applied only when `tp_size > 1`, and `ep_plan` only when `ep_size > 1`. With TP enabled and EP disabled, the full TP plan applies, expert rules included. +The expert forward rule in `ep_plan` selects how tokens reach the experts: + +| rule | mechanism | layout | +| :--- | :--- | :--- | +| `"moe_tp_experts"` with `"ep_router"` on the router | masking and all-reduce: every rank runs its local experts on the whole batch, the router masks the others, and an all-reduce combines the outputs | `ep_size=tp_size` | +| `"ep_dispatch_experts"` | [token dispatch](#token-dispatch): each rank keeps its own tokens and only exchanges the routed (token, expert) pairs with two all-to-all collectives | `ep_size` a multiple of `tp_size` that divides `fsdp_size * tp_size` | + > [!TIP] > `enable_expert_parallel=True` is a deprecated alias for `ep_size=tp_size`, used only when `ep_size` is omitted, and emits a `FutureWarning`. -Launch your inference script with [torchrun](https://pytorch.org/docs/stable/elastic/run.html) and specify how many devices to use. The number of devices must evenly divide the total number of experts. +Launch your inference script with [torchrun](https://pytorch.org/docs/stable/elastic/run.html). The number of processes must equal `tp_size * fsdp_size * pp_size`, and `ep_size` must evenly divide the number of experts. ```zsh torchrun --nproc-per-node 8 your_script.py @@ -67,9 +74,98 @@ distributed_config = DistributedConfig( Providing a plan does not infer parallel sizes: set `tp_size` and `ep_size` explicitly. +Qwen3 MoE defaults to `"ep_dispatch_experts"`. To use masking and all-reduce instead, set `ep_size=tp_size` and override both the router and the expert forward rules: + +```py +distributed_config = DistributedConfig( + tp_size=4, + ep_size=4, + ep_plan={ + "model.layers.*.mlp.gate": "ep_router", + "model.layers.*.mlp.experts": "moe_tp_experts", + }, +) +``` + +Conversely, override the expert forward rule of a model whose plan uses masking with `"ep_dispatch_experts"` to use token dispatch. The router rule is then ignored, since dispatch needs the global expert ids to find each expert's owner. + +## Token dispatch + +With token dispatch, each rank trains on its own part of the batch. At every MoE layer, a rank routes its tokens, sends each (token, expert) pair to the rank that owns the expert with an all-to-all, runs its local experts on what it receives, gets the results back with a second all-to-all and combines them with the routing weights. Only the routed activations and expert outputs travel, and no rank computes experts for tokens it does not own. + +```py +from transformers import AutoModelForCausalLM +from transformers.distributed import DistributedConfig + +distributed_config = DistributedConfig( + tp_size=1, + fsdp_size=8, + ep_size=4, +) +model = AutoModelForCausalLM.from_pretrained("Qwen/Qwen3-30B-A3B", distributed_config=distributed_config) +``` + +With `tp_size=1`, `ep_size` must divide `fsdp_size` and the number of experts, and attention is not bound by `num_key_value_heads`. For the rest of the model: + +- The parameters outside the experts are sharded with [FSDP2](./fsdp) across `fsdp`, which reduces their gradients. +- The experts are sharded across `ep` and, when `efsdp_size = fsdp_size * tp_size // ep_size` is larger than one, additionally FSDP-sharded across `efsdp`. They are always FSDP-wrapped, so `fsdp_mixed_precision` and `fsdp_cpu_offload` apply to them too and [`~PreTrainedModel.save_pretrained`] gathers them like any other parameter. +- An expert parallel group holds `ep_size / tp_size` batches, so an expert's gradient is a sum over that many batches. The `efsdp` reduction divides by `fsdp_size` instead of its group size, which gives the same per-batch average FSDP2 takes for the dense modules. +- Every local expert also processes one zero pad row per layer. A rank whose experts received no tokens still joins the reverse all-to-all and the expert gradient reduction. + +### How the sizes combine + +- `tp_size * fsdp_size` is the number of processes. `ep_size` adds none: it regroups the same ranks for the expert weights only. +- `ep_size` cuts the expert list into `ep_size` blocks. Each rank computes `num_experts / ep_size` experts, and `ep_size` consecutive ranks hold one complete set. That set of ranks is the group the all-to-all runs in. +- `efsdp_size = fsdp_size * tp_size / ep_size` is how many complete copies of the expert set exist. Ranks at the same position in different copies shard those experts for memory and average their gradients, like FSDP does for the dense modules. +- The batch a rank holds depends on `fsdp` only. Consecutive ranks form a TP group and get the same batch; the `fsdp_size` groups get different batches. + +Eight processes, `tp_size=2, fsdp_size=4`, eight experts: + +```text +rank 0 1 2 3 4 5 6 7 +batch [====B0====] [====B1====] [====B2====] [====B3====] one batch per TP pair +tp 0 1 0 1 0 1 0 1 + +ep_size=2 E0-3 E4-7 E0-3 E4-7 E0-3 E4-7 E0-3 E4-7 group = a TP pair, efsdp_size=4 +ep_size=4 E0E1 E2E3 E4E5 E6E7 E0E1 E2E3 E4E5 E6E7 group = two pairs, efsdp_size=2 +ep_size=8 E0 E1 E2 E3 E4 E5 E6 E7 group = all ranks, efsdp_size=1 +``` + +Two numbers follow from the picture: + +- Inside an EP group, each token exists `tp_size` times, once per rank of the pair that holds its batch. This does not depend on `ep_size`. +- An EP group holds `ep_size / tp_size` different batches. This is the count an expert's gradient sums over. + +### With tensor parallelism + +Set `tp_size > 1` to shard the dense modules with the TP plan while the experts use dispatch. On eight processes: + +```py +distributed_config = DistributedConfig( + tp_size=2, + fsdp_size=4, + ep_size=4, +) +``` + +Each pair of TP ranks receives the same batch, because tensor parallelism replicates the activations inside the pair. If both ranks dispatched all of their tokens, the rank owning an expert would receive every token twice, compute it twice, and its weight gradient would double. Experts are whole on one rank, so the duplicate cannot be split by weights, and the owner is usually another rank, so it cannot be resolved by ownership as masking does. The pair therefore splits the rows: each TP rank dispatches a disjoint `1 / tp_size` of the tokens, results come back to the rank that sent them, and an all-reduce over the pair of the zero-padded halves restores the replicated output the next layer expects. The split is by `tp_size`, not `ep_size`, since only the ranks that hold a batch can send it. Expert groups span four ranks and each expert is FSDP-sharded across `efsdp_size = 2` ranks, while the trunk's FSDP group spans four ranks. The model's usual TP constraints, such as attention-head divisibility, still apply to the dense modules. Token slices may be uneven or empty, including during single-token decoding. + +Token dispatch cannot be combined yet with pipeline parallelism yet; use `pp_size=1` (not tested yet) +These configurations each use eight GPUs: + +| Configuration | Result | +| :--- | :--- | +| `DistributedConfig(tp_size=4, fsdp_size=2, ep_size=4)` | Dispatch with TP groups of four, each slicing its batch in four (Qwen3 MoE default plan); unless you specify a ep_plan to use the legacy masked EP | +| `DistributedConfig(tp_size=1, fsdp_size=8, ep_size=4)` | Dispatch with an independent batch on each rank, no slicing, and experts FSDP-sharded across pairs of ranks. | +| `DistributedConfig(tp_size=2, fsdp_size=4, ep_size=4)` | Dispatch with a TP pair per batch, each pair slicing its batch in two; two batches per expert group. | +| `DistributedConfig(tp_size=8, ep_size=8)` | Dispatch with every rank sharing one batch, sliced in eight, or masking and all-reduce for a masked plan. | + +> [!WARNING] +> The [`Trainer`] does not account for token dispatch yet: batch and token counting assume the all-reduce layout, where the ranks of a TP group share a batch and `fsdp_size` data-parallel shards exist. Trainer support for dispatch comes in a follow-up. + ## Combining with FSDP2 -Tensor and expert parallelism shard the weights across `tp`, but the optimizer state and the modules without a rule are still replicated on every rank of the group, which limits how large a model you can train. Add [FSDP2](./fsdp) on a second mesh dimension with `fsdp_size`, and keep `ep_size=tp_size` for the expert parallel width. +Tensor and expert parallelism shard the weights across `tp`, but the optimizer state and the modules without a rule are still replicated on every rank of the group, which limits how large a model you can train. Add [FSDP2](./fsdp) on a second mesh dimension with `fsdp_size`. With masking and all-reduce, keep `ep_size=tp_size` and pass `ep_plan={"layers.*.mlp.gate": "ep_router", "layers.*.mlp.experts": "moe_tp_experts"}``. ```py from transformers import AutoModelForCausalLM @@ -77,8 +173,12 @@ from transformers.distributed import DistributedConfig distributed_config = DistributedConfig( tp_size=4, - ep_size=4, # expert parallel size, must match tp_size + ep_size=4, # expert parallel size, must match tp_size with masking and all-reduce fsdp_size=2, # data parallel shards + ep_plan={ + "model.layers.*.mlp.gate": "ep_router", + "model.layers.*.mlp.experts": "moe_tp_experts", + }, ) model = AutoModelForCausalLM.from_pretrained("Qwen/Qwen3-30B-A3B", distributed_config=distributed_config) ``` diff --git a/src/transformers/distributed/configuration_utils.py b/src/transformers/distributed/configuration_utils.py index 77800f12a64d..a2b3c14248ac 100644 --- a/src/transformers/distributed/configuration_utils.py +++ b/src/transformers/distributed/configuration_utils.py @@ -18,6 +18,7 @@ from dataclasses import asdict, dataclass from typing import Literal +from ..utils import is_torch_greater_or_equal from .utils import _get_torch_distributed_rank @@ -48,12 +49,14 @@ class DistributedConfig: pp_size (`int`, *optional*): Number of devices for pipeline parallelism. If `None` and another parallel mode is set, defaults to 1. ep_size (`int`, *optional*): - Number of devices owning distinct expert shards. Defaults to 1. Set it explicitly to enable EP. - Model execution currently requires `ep_size=tp_size` when EP is enabled. + Number of devices owning distinct expert shards. Defaults to 1. Set it explicitly to enable EP. Must be + a multiple of `tp_size` and divide `fsdp_size * tp_size`. All-reduce expert plans require + `ep_size=tp_size`; token dispatch (`"ep_dispatch_experts"`) also allows `ep_size > tp_size`. ep_plan (`dict[str, str]`, *optional*): Expert parallel sharding plan. Leave as `None` to use the model's predefined `base_model_ep_plan`. Pass a dictionary to override individual rules of that plan; unspecified rules are kept. Applied only when - `ep_size > 1`, and its rules take precedence over `tp_plan` rules for the same modules. + `ep_size > 1`, and its rules take precedence over `tp_plan` rules for the same modules. An + `"ep_dispatch_experts"` rule selects all-to-all token dispatch instead of router masking and all-reduce. """ tp_size: int | None = None @@ -130,6 +133,21 @@ def _validate_mesh_config(self): "Use DistributedConfig(tp_size=N, fsdp_size=M), or combine TP and PP." ) + def _validate_resolved_ep_plan(self, ep_plan: dict[str, str]): + """Validate the layout against the resolved EP plan, once the model's defaults and overrides are merged.""" + if self.ep_size <= 1 or not ep_plan: + return + + if "ep_dispatch_experts" in ep_plan.values(): + if self.pp_size > 1: + raise ValueError("Combining token dispatch with pipeline parallelism is not supported/tested yet.") + if not is_torch_greater_or_equal("2.7"): + raise OSError("Expert-parallel token dispatch requires `torch>=2.7`.") + elif {"ep_router", "moe_tp_experts"}.issubset(ep_plan.values()) and self.ep_size != self.tp_size: + raise ValueError( + "All-reduce expert parallelism requires `ep_size=tp_size`, so every rank of an expert group sees the same tokens" + ) + @classmethod def from_dict(cls, config_dict: dict, **kwargs) -> "DistributedConfig": merged = {**config_dict, **kwargs} diff --git a/src/transformers/distributed/fsdp.py b/src/transformers/distributed/fsdp.py index aef3a5979a8c..d651ae61f398 100644 --- a/src/transformers/distributed/fsdp.py +++ b/src/transformers/distributed/fsdp.py @@ -27,6 +27,7 @@ import torch.nn as nn from .configuration_utils import DistributedConfig + from .utils import MeshManager if is_torch_available(): import torch @@ -184,15 +185,15 @@ def verify_fsdp_plan(module_names: list[str], fsdp_plan: dict[str, str] | None) logger.warning(f"The following FSDP rules were not applied to any module: {unused_rules}") -def apply_fully_sharded_data_parallelism( - model: nn.Module, fsdp_mesh: torch.distributed.device_mesh.DeviceMesh -) -> nn.Module: +def apply_fully_sharded_data_parallelism(model: nn.Module, mesh_manager: MeshManager) -> nn.Module: """ - Apply FSDP2 (fully_shard) to a model. + Apply FSDP2 (fully_shard) to a model: dispatched experts on `efsdp` and the rest of the modules on `fsdp`. Torch availability, distributed initialization and the version requirement are asserted upstream by `initialize_distributed_mesh`. """ + distributed_config = model.config.distributed_config + fsdp_mesh = mesh_manager.get_mesh("fsdp") fsdp_plan = dict(getattr(model, "_fsdp_plan", None) or {}) if not fsdp_plan: raise ValueError( @@ -206,6 +207,20 @@ def apply_fully_sharded_data_parallelism( adapted_fsdp_plan = _resolve_tied_embed_lm_head_plan(fsdp_plan, model) reshard_targets, no_reshard_targets = expand_fsdp_plan(model, adapted_fsdp_plan) + fsdp_policy_kwargs = _get_fsdp_policy_kwargs(distributed_config) + if distributed_config.ep_size > 1 and "ep_dispatch_experts" in model.ep_plan.values(): + expert_mesh = mesh_manager.get_mesh("efsdp") + for module in model.modules(): + if getattr(module, "_is_expert_parallel", False): + fully_shard(module, mesh=expert_mesh, reshard_after_forward=True, **fsdp_policy_kwargs) + # An expert group spans several data-parallel batches, so an expert's gradient sums over + # all of them. FSDP2 would divide by the efsdp group size; dividing by fsdp_size instead + # gives the same per-batch average the dense modules get on the fsdp mesh, even when efsdp has a single rank. + module.set_gradient_divide_factor(float(distributed_config.fsdp_size)) + if torch.distributed.get_backend(expert_mesh.get_group()) != "nccl": + # Non-NCCL backends need to sum first, then apply the division otherwise it runtime error. + module.set_force_sum_reduction_for_comms(True) + for module_name, module in reshard_targets: fully_shard(module, mesh=fsdp_mesh, reshard_after_forward=True, **fsdp_policy_kwargs) logger.debug(f"Applied fully_shard to {module_name} (reshard=True)") diff --git a/src/transformers/distributed/mixin.py b/src/transformers/distributed/mixin.py index 4ae276547185..584e2e5a6fb4 100644 --- a/src/transformers/distributed/mixin.py +++ b/src/transformers/distributed/mixin.py @@ -25,6 +25,7 @@ from .pipeline_parallel import apply_pipeline_parallelism from .tensor_parallel import ( _validate_parallel_plan_styles, + apply_expert_parallelism, apply_tensor_parallelism, gather_state_dict_for_save, resolve_parallel_plans, @@ -197,19 +198,29 @@ def maybe_distribute_model( # Resolve both plans before sharding anything: overrides are merged into `model.tp_plan` / `model.ep_plan`, # and the experts named by the EP plan are removed from the TP plan so they are sharded once. tp_plan, ep_plan = resolve_parallel_plans(model, distributed_config) + distributed_config._validate_resolved_ep_plan(ep_plan) if distributed_config.pp_size > 1: model = apply_pipeline_parallelism(model, mesh_manager.get_mesh("pp")) - tp_mesh = mesh_manager.get_mesh("tp") if tp_plan: - model = apply_tensor_parallelism(model, tp_mesh, tp_plan) + model = apply_tensor_parallelism(model, mesh_manager.get_mesh("tp"), tp_plan) + if ep_plan: - # Legacy masked EP: the EP group is the TP group, every rank keeps every token. - model = apply_tensor_parallelism(model, tp_mesh, ep_plan) + tp_mesh = mesh_manager.get_mesh("tp") + ep_mesh = mesh_manager.get_mesh("ep") + + if {"ep_router", "moe_tp_experts"}.issubset(ep_plan.values()): + # Legacy masked EP: the EP group is the TP group, every rank keeps every token. + model = apply_tensor_parallelism(model, tp_mesh, ep_plan) + elif "ep_dispatch_experts" in ep_plan.values(): + # EP + DP with tp_size >= 1: the ranks of a TP group share the same batch. If we want a specific token, + # we will have to slice the batch here in order to avoid computing tp_size times the same batch. + model = apply_expert_parallelism(model, ep_mesh, tp_mesh, ep_plan) + + if distributed_config.fsdp_size > 1 or "ep_dispatch_experts" in ep_plan.values(): + model = apply_fully_sharded_data_parallelism(model, mesh_manager) - if distributed_config.fsdp_size > 1: - model = apply_fully_sharded_data_parallelism(model, mesh_manager.get_mesh("fsdp")) return model def should_save_on_this_rank(self, is_main_process: bool) -> bool: @@ -270,9 +281,11 @@ def gather_sharded_state_dict_for_save( if distributed_config is None: return state_dict - if distributed_config.fsdp_size > 1: - # Also covers the 2-D (fsdp, tp) mesh: every parameter is FSDP-managed, and the full - # state dict is only materialized on rank 0. + if distributed_config.fsdp_size > 1 or ( + distributed_config.ep_size > 1 and "ep_dispatch_experts" in self.ep_plan.values() + ): + # Also covers the 2-D (fsdp, tp) mesh and token dispatch: every parameter is FSDP-managed, and the + # full state dict is only materialized on rank 0. if not _is_torch_distributed_initialized(): raise ValueError( "Saving an FSDP-wrapped model requires torch.distributed to be initialized. " @@ -294,5 +307,5 @@ def barrier_after_gathered_checkpoint_save(self, distributed_config: Distributed """Barrier so non-writer ranks wait for rank 0 to finish gathered checkpoint writes.""" if distributed_config is None: return - if distributed_config.tp_size > 1 or distributed_config.fsdp_size > 1: + if distributed_config.tp_size > 1 or distributed_config.fsdp_size > 1 or distributed_config.ep_size > 1: _distributed_barrier() diff --git a/src/transformers/distributed/tensor_parallel.py b/src/transformers/distributed/tensor_parallel.py index 5864520fac99..d6c14907efe6 100644 --- a/src/transformers/distributed/tensor_parallel.py +++ b/src/transformers/distributed/tensor_parallel.py @@ -15,6 +15,7 @@ import contextlib import re +from collections.abc import Callable from itertools import chain from typing import TYPE_CHECKING @@ -37,6 +38,7 @@ if is_torch_distributed_available(): import torch.distributed as dist + from torch.distributed.nn.functional import all_to_all_single from torch.distributed.tensor import DTensor, Partial, Replicate, Shard, distribute_tensor from torch.distributed.tensor.placement_types import _StridedShard @@ -748,6 +750,165 @@ def transform_output_post_forward(self, module, output, mesh): return output +class EpDispatchExpertsParallel(MoeExpertsParallel): + """ + Dispatch disjoint TP token slices to the experts' owners, then replicate the combined output on TP. + + Example: + Let's say we have 8 experts [E0, E7] with DistributedConfig(tp_size=2, fsdp_size=4, ep_size=4). That imply: + - Since fsdp_size=4, we have 4 batches B denoted [B0, B3] + - Because we have tp_size=2, that means *both ranks share the same batch* + - efsdp = (fsdp_size * tp_size) / ep_size = 4 * 2 / 4 = 2 + + GPU 0 1 2 3 4 5 6 7 + | | | | | | | | + dense view --------------------------------------------------------- + batch [====B0====] [====B1====] [====B2====] [====B3====] + tp_size [___________ 0 ___________] [___________ 1 ___________] + fsdp_size 0 1 2 3 + + expert view --------------------------------------------------------- + + experts E0E1 E2E3 E4E5 E6E7 E0E1 E2E3 E4E5 E6E7 + ep_size 0 1 2 3 0 1 2 3 + efsdp_size [___________ 0 ___________] [___________ 1 ___________] + + Assume token 1 in B0 chose E4, which lives on another rank, so it must travel by all-to-all. B0 sits on 2 ranks (cf diagram), so if both sent it, E4 would compute it twice. + We need to make sure that token 1 is in rank 0 range, so rank 0 sends it and rank 1 does not have it in its slice. + """ + + def _dispatch_tokens( + self, + hidden_states: torch.Tensor, + top_k_index: torch.Tensor, + num_local_experts: int, + ep_group, + ep_size: int, + ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, list[int], list[int]]: + """Send each selected (token, expert) pair to the rank that owns the expert. + + Also returns the sort order and the per-rank split sizes that `_combine_tokens` needs to reverse the exchange. + """ + hidden_dim = hidden_states.size(-1) + num_top_k = top_k_index.size(-1) + + # Sorting the selected pairs by expert groups them by owner rank, since each rank owns a contiguous range of + # experts, and the per-expert counts tell every receiver which expert each token it gets is for. The split + # sizes are the one host sync of the layer. + expert_ids = top_k_index.reshape(-1) + order = torch.argsort(expert_ids) + send_tokens = hidden_states[order // num_top_k] + send_counts = torch.zeros(num_local_experts * ep_size, dtype=torch.long, device=hidden_states.device) + send_counts = send_counts.scatter_add_(0, expert_ids, torch.ones_like(expert_ids)).view( + ep_size, num_local_experts + ) + recv_counts = torch.empty_like(send_counts) + torch.distributed.all_to_all_single(recv_counts, send_counts, group=ep_group) + send_sizes, recv_sizes = torch.stack([send_counts.sum(dim=1), recv_counts.sum(dim=1)]).tolist() + recv_tokens = all_to_all_single( + send_tokens.new_empty(sum(recv_sizes), hidden_dim), + send_tokens, + output_split_sizes=recv_sizes, + input_split_sizes=send_sizes, + group=ep_group, + ) + recv_expert_ids = torch.arange(num_local_experts, device=hidden_states.device).repeat(ep_size) + recv_expert_ids = recv_expert_ids.repeat_interleave(recv_counts.reshape(-1), output_size=sum(recv_sizes)) + return recv_tokens, recv_expert_ids, order, send_sizes, recv_sizes + + def _run_local_experts( + self, + experts_forward: Callable, + tokens: torch.Tensor, + expert_ids: torch.Tensor, + num_local_experts: int, + ) -> torch.Tensor: + """Run local experts with top-1 routing and unit weights; apply routing weights after combine.""" + # One zero row per expert keeps tokens and all expert weights connected to backward. Without it, + # empty eager experts can skip the reverse all-to-all and FSDP reduction, leaving other ranks waiting. + num_tokens, hidden_dim = tokens.shape + local_expert_ids = torch.arange(num_local_experts, device=tokens.device) + tokens = torch.cat([tokens, tokens.new_zeros(num_local_experts, hidden_dim)]) + expert_ids = torch.cat([expert_ids, local_expert_ids]).unsqueeze(-1) + weights = torch.ones_like(expert_ids, dtype=tokens.dtype) + return experts_forward(tokens, expert_ids, weights)[:num_tokens] + + def _combine_tokens( + self, + expert_output: torch.Tensor, + top_k_weights: torch.Tensor, + order: torch.Tensor, + send_sizes: list[int], + recv_sizes: list[int], + ep_group, + ) -> torch.Tensor: + """Return expert outputs to the token owners and combine them with routing weights.""" + num_tokens, num_top_k = top_k_weights.shape + hidden_dim = expert_output.size(-1) + recv_out = all_to_all_single( + expert_output.new_empty(order.numel(), hidden_dim), + expert_output, + output_split_sizes=send_sizes, + input_split_sizes=recv_sizes, + group=ep_group, + ) + # Restore the original (token, top-k slot) order, then apply routing weights. + token_outputs = torch.empty_like(recv_out) + token_outputs[order] = recv_out + token_outputs = token_outputs.view(num_tokens, num_top_k, hidden_dim) + return (token_outputs * top_k_weights.unsqueeze(-1)).sum(dim=1) + + def transform_inputs_pre_forward(self, module, args, kwargs, mesh, *, tp_mesh=None): + hidden_states, top_k_index, top_k_weights = args + if isinstance(hidden_states, DTensor): + hidden_states = hidden_states.to_local() + if isinstance(top_k_weights, DTensor): + top_k_weights = top_k_weights.to_local() + if tp_mesh is None or tp_mesh.size() == 1: + return (hidden_states, top_k_index, top_k_weights), kwargs + + tp_group = tp_mesh.get_group() + hidden_states = _AllReduceBackward.apply(hidden_states, tp_group) + top_k_weights = _AllReduceBackward.apply(top_k_weights, tp_group) + # TP ranks share the same batch. Slice tokens here so each is dispatched only once, then restore the full output in the post hook + num_tokens, tp_rank, tp_size = hidden_states.size(0), tp_mesh.get_local_rank(), tp_mesh.size() + rows = slice(num_tokens * tp_rank // tp_size, num_tokens * (tp_rank + 1) // tp_size) + return (hidden_states[rows], top_k_index[rows], top_k_weights[rows]), kwargs + + def transform_output_post_forward(self, module, output, mesh, *, tp_mesh=None, num_tokens=None): + if tp_mesh is None or tp_mesh.size() == 1: + return output + tp_rank, tp_size = tp_mesh.get_local_rank(), tp_mesh.size() + rows = slice(num_tokens * tp_rank // tp_size, num_tokens * (tp_rank + 1) // tp_size) + full_output = output.new_zeros(num_tokens, output.size(-1)) + full_output[rows] = output + return _AllReduceForward.apply(full_output, tp_mesh.get_group()) + + def install_forward(self, module, ep_mesh, *, tp_mesh=None): + experts_forward = module.forward + ep_group, ep_size = ep_mesh.get_group(), ep_mesh.size() + + def tp_forward(hidden_states, top_k_index, top_k_weights): + # Read the full token count before the pre hook slices the inputs on TP. + num_tokens = hidden_states.size(0) + (hidden_states, top_k_index, top_k_weights), _ = self.transform_inputs_pre_forward( + module, (hidden_states, top_k_index, top_k_weights), {}, ep_mesh, tp_mesh=tp_mesh + ) + with self.context_around_forward(module, ep_mesh): + tokens, expert_ids, order, send_sizes, recv_sizes = self._dispatch_tokens( + hidden_states, top_k_index, module.num_experts, ep_group, ep_size + ) + expert_output = self._run_local_experts(experts_forward, tokens, expert_ids, module.num_experts) + output = self._combine_tokens( + expert_output, top_k_weights, order, send_sizes, recv_sizes, ep_group + ).to(hidden_states.dtype) + + return self.transform_output_post_forward(module, output, ep_mesh, tp_mesh=tp_mesh, num_tokens=num_tokens) + + module.forward = tp_forward + return module + + class MoeTensorParalellMegaMoeExperts(MoeExpertsParallel): """TP layer for DeepGEMM Mega MoE experts. @@ -785,6 +946,7 @@ class ParallelInterface(GeneralInterface): "sequence_parallel": SequenceParallel(use_local_output=True), "grouped_gemm": MoEParamShard(Shard(0), shards_expert_dim=True), "ep_router": EpRouterParallel(), + "ep_dispatch_experts": EpDispatchExpertsParallel(), "megamoe_router": RouterParallelMegaMoe(), "moe_tp_experts": MoeExpertsParallel(), "megamoe_experts": MoeTensorParalellMegaMoeExperts(), @@ -849,22 +1011,23 @@ def resolve_parallel_plans( return tp_plan, ep_plan -def apply_tensor_parallelism(model: nn.Module, tp_mesh: DeviceMesh, plan: dict[str, str] | None = None): - plan = model.tp_plan if plan is None else plan - _validate_parallel_plan_styles(plan) +def apply_tensor_parallelism(model, tp_mesh, tp_plan=None): + """Apply parameter sharding and forward hooks on the TP mesh.""" + tp_plan = model.tp_plan if tp_plan is None else tp_plan + _validate_parallel_plan_styles(tp_plan) for name, module in model.named_modules(): # Create DTensor placeholders so the loader knows which shard belongs to this rank. for p_name, _ in list(module.named_parameters(recurse=False)): full = f"{name}.{p_name}" if name else p_name - style_name = _get_parameter_plan(parameter_name=full, plan=plan, is_weight=True) + style_name = _get_parameter_plan(parameter_name=full, plan=tp_plan, is_weight=True) if style_name is not None and style_name in ALL_PARALLEL_STYLES: style = ALL_PARALLEL_STYLES[style_name] style.validate_param(module, p_name, tp_mesh, parameter_name=full) style.shard_param(module, p_name, tp_mesh) - # Install the input/output transforms required by this module's style. - style_name = _get_parameter_plan(parameter_name=name, plan=plan, is_weight=False) + # Install the input/output transforms required by this module's TP style. + style_name = _get_parameter_plan(parameter_name=name, plan=tp_plan, is_weight=False) if style_name is not None and style_name in ALL_PARALLEL_STYLES: if style_name == "mla_kv_a_proj": # MLA needs to know the qk_rope_head_dim to split the projection output into KV and RoPE parts. @@ -876,6 +1039,30 @@ def apply_tensor_parallelism(model: nn.Module, tp_mesh: DeviceMesh, plan: dict[s return model +def apply_expert_parallelism(model: nn.Module, ep_mesh: DeviceMesh, tp_mesh: DeviceMesh, plan: dict[str, str]): + """Shard experts on EP; use TP to split shared tokens and reconstruct outputs around dispatch.""" + for name, module in model.named_modules(): + for p_name, _ in list(module.named_parameters(recurse=False)): + full = f"{name}.{p_name}" if name else p_name + style_name = _get_parameter_plan(parameter_name=full, plan=plan, is_weight=True) + if style_name is not None and style_name in ALL_PARALLEL_STYLES: + style = ALL_PARALLEL_STYLES[style_name] + style.validate_param(module, p_name, ep_mesh, parameter_name=full) + style.shard_param(module, p_name, ep_mesh) + + # Dispatch hooks need both meshes to redistribute tokens between TP and EP ranks. + style_name = _get_parameter_plan(parameter_name=name, plan=plan, is_weight=False) + if style_name is not None and style_name in ALL_PARALLEL_STYLES: + style = ALL_PARALLEL_STYLES[style_name] + if style_name == "ep_dispatch_experts": + style.install_forward(module, ep_mesh=ep_mesh, tp_mesh=tp_mesh) + else: + style.install_forward(module, ep_mesh) + module._is_hooked = True + + return model + + def gather_state_dict_for_save( state_dict: dict[str, torch.Tensor], _tp_plan: dict[str, str], diff --git a/src/transformers/models/qwen3_moe/configuration_qwen3_moe.py b/src/transformers/models/qwen3_moe/configuration_qwen3_moe.py index 9a7d4b4c8b5b..62cf801edfcf 100644 --- a/src/transformers/models/qwen3_moe/configuration_qwen3_moe.py +++ b/src/transformers/models/qwen3_moe/configuration_qwen3_moe.py @@ -67,14 +67,14 @@ class Qwen3MoeConfig(PreTrainedConfig): "layers.*.mlp.up_proj": "colwise", "layers.*.mlp.down_proj": "rowwise", } - # Expert-only EP plan: only shards MoE experts, not attention. - # Attention is left unsharded — FSDP2 handles attention weight distribution. - # This allows EP to scale beyond num_kv_heads (not constrained by 4 for Qwen3-30B). + # Token dispatch by default, so `ep_size` can exceed `tp_size` (bounded by `num_key_value_heads`, 4 on + # Qwen3-30B): with `tp_size=1` the attention is left to FSDP2. For router masking with all-reduce, set + # `ep_size=tp_size` and pass `ep_plan={"layers.*.mlp.gate": "ep_router", "layers.*.mlp.experts": "moe_tp_experts"}` + # (prefixed with `model.` on the causal LM). base_model_ep_plan = { - "layers.*.mlp.gate": "ep_router", "layers.*.mlp.experts.gate_up_proj": "grouped_gemm", "layers.*.mlp.experts.down_proj": "grouped_gemm", - "layers.*.mlp.experts": "moe_tp_experts", + "layers.*.mlp.experts": "ep_dispatch_experts", } base_model_pp_plan = { "embed_tokens": (["input_ids"], ["inputs_embeds"]), diff --git a/tests/tensor_parallel/test_tensor_parallel.py b/tests/tensor_parallel/test_tensor_parallel.py index a265cbd4087f..3fed8106c98d 100644 --- a/tests/tensor_parallel/test_tensor_parallel.py +++ b/tests/tensor_parallel/test_tensor_parallel.py @@ -23,6 +23,7 @@ from transformers.distributed.tensor_parallel import ( ALL_PARALLEL_STYLES, ColwiseParallel, + EpDispatchExpertsParallel, PackedColwiseParallel, PackedRowwiseParallel, RowwiseParallel, @@ -31,7 +32,7 @@ # Qwen3 MoE's predefined plans, as resolved on `Qwen3MoeModel` (no `model.` prefix). -DENSE_TP_PLAN = { +TP_DENSE_PLAN = { "layers.*.self_attn.q_proj": "colwise", "layers.*.self_attn.k_proj": "colwise", "layers.*.self_attn.v_proj": "colwise", @@ -42,17 +43,17 @@ "layers.*.mlp.up_proj": "colwise", "layers.*.mlp.down_proj": "rowwise", } -EXPERT_TP_PLAN = { +TP_EXPERT_PLAN = { "layers.*.mlp.experts.gate_up_proj": "packed_colwise", "layers.*.mlp.experts.down_proj": "rowwise", "layers.*.mlp.experts": "moe_tp_experts", } EP_PLAN = { - "layers.*.mlp.gate": "ep_router", "layers.*.mlp.experts.gate_up_proj": "grouped_gemm", "layers.*.mlp.experts.down_proj": "grouped_gemm", - "layers.*.mlp.experts": "moe_tp_experts", + "layers.*.mlp.experts": "ep_dispatch_experts", } +EP_PLAN_MASKED = EP_PLAN | {"layers.*.mlp.gate": "ep_router", "layers.*.mlp.experts": "moe_tp_experts"} @require_torch @@ -74,6 +75,10 @@ def setUp(self): with torch.device("meta"): self.model = Qwen3MoeModel(self.config) + def _reset_plans(self): + self.model.tp_plan = TP_DENSE_PLAN | TP_EXPERT_PLAN + self.model.ep_plan = EP_PLAN.copy() + def test_ep_plan_setter(self): self.model.ep_plan = None self.assertEqual(self.model.ep_plan, {}) @@ -91,7 +96,7 @@ def test_disabled_parallelism_has_no_plans(self): def test_tp_only_keeps_experts_in_tp_plan(self): tp_plan, ep_plan = tensor_parallel.resolve_parallel_plans(self.model, DistributedConfig(tp_size=4)) - self.assertEqual(tp_plan, DENSE_TP_PLAN | EXPERT_TP_PLAN) + self.assertEqual(tp_plan, TP_DENSE_PLAN | TP_EXPERT_PLAN) self.assertEqual(ep_plan, {}) def test_ep_takes_experts_and_router_out_of_tp_plan(self): @@ -101,7 +106,7 @@ def test_ep_takes_experts_and_router_out_of_tp_plan(self): ): with self.subTest(config=config): tp_plan, ep_plan = tensor_parallel.resolve_parallel_plans(self.model, config) - self.assertEqual(tp_plan, DENSE_TP_PLAN) + self.assertEqual(tp_plan, TP_DENSE_PLAN | TP_EXPERT_PLAN) self.assertEqual(ep_plan, EP_PLAN) def test_legacy_flag_is_an_alias_for_ep_size(self): @@ -125,8 +130,8 @@ def test_legacy_flag_is_an_alias_for_ep_size(self): def test_ep_plan_is_a_dict_and_round_trips(self): with self.assertRaisesRegex(ValueError, "`ep_plan` must be a dictionary or None"): DistributedConfig(tp_size=4, ep_size=4, ep_plan="auto") - config = DistributedConfig(tp_size=4, ep_size=4, ep_plan={"layers.*.mlp.gate": "ep_router"}) - self.assertEqual(config.to_dict()["ep_plan"], {"layers.*.mlp.gate": "ep_router"}) + config = DistributedConfig(tp_size=4, ep_size=4, ep_plan=EP_PLAN_MASKED) + self.assertEqual(config.to_dict()["ep_plan"], EP_PLAN_MASKED) self.assertEqual(DistributedConfig.from_dict(config.to_dict()), config) def test_overrides_merge_into_the_predefined_plans(self): @@ -137,7 +142,7 @@ def test_overrides_merge_into_the_predefined_plans(self): ep_plan={"layers.*.mlp.experts.down_proj": "rowwise"}, ) tp_plan, ep_plan = tensor_parallel.resolve_parallel_plans(self.model, config) - self.assertEqual(tp_plan, DENSE_TP_PLAN | {"layers.*.self_attn.q_proj": "colwise_rep"}) + self.assertEqual(tp_plan, TP_DENSE_PLAN | {"layers.*.self_attn.q_proj": "colwise_rep"}) self.assertEqual(ep_plan, EP_PLAN | {"layers.*.mlp.experts.down_proj": "rowwise"}) # The merged plans are stored on the model, the config defaults are untouched. self.assertEqual(self.model.tp_plan["layers.*.self_attn.q_proj"], "colwise_rep") @@ -149,7 +154,7 @@ def test_overrides_merge_into_the_predefined_plans(self): self.assertEqual(config.ep_plan, {"layers.*.mlp.experts.down_proj": "rowwise"}) # The merged EP plan stays on the model but is not applied while EP is disabled. tp_plan, ep_plan = tensor_parallel.resolve_parallel_plans(self.model, DistributedConfig(tp_size=4)) - self.assertEqual(tp_plan, DENSE_TP_PLAN | EXPERT_TP_PLAN | {"layers.*.self_attn.q_proj": "colwise_rep"}) + self.assertEqual(tp_plan, TP_DENSE_PLAN | TP_EXPERT_PLAN | {"layers.*.self_attn.q_proj": "colwise_rep"}) self.assertEqual(ep_plan, {}) self.assertEqual(self.model.ep_plan["layers.*.mlp.experts.down_proj"], "rowwise") @@ -158,10 +163,11 @@ def test_ep_rules_take_precedence_over_tp_rules_for_the_same_modules(self): tp_size=4, ep_size=4, tp_plan={"layers.*.mlp.experts.gate_up_proj": "packed_rowwise", "layers.*.mlp.gate": "colwise"}, + ep_plan=EP_PLAN_MASKED, ) tp_plan, ep_plan = tensor_parallel.resolve_parallel_plans(self.model, config) - self.assertEqual(tp_plan, DENSE_TP_PLAN) - self.assertEqual(ep_plan, EP_PLAN) + self.assertEqual(tp_plan, TP_DENSE_PLAN) + self.assertEqual(ep_plan, EP_PLAN_MASKED) # The custom TP rules are kept on the model and apply as soon as EP is disabled. tp_plan, ep_plan = tensor_parallel.resolve_parallel_plans(self.model, DistributedConfig(tp_size=4)) self.assertEqual(tp_plan["layers.*.mlp.experts.gate_up_proj"], "packed_rowwise") @@ -173,7 +179,7 @@ def test_ep_requires_an_expert_plan(self): with self.assertRaisesRegex(ValueError, "does not define an expert-parallel plan"): tensor_parallel.resolve_parallel_plans(self.model, DistributedConfig(tp_size=4, ep_size=4)) config = DistributedConfig(tp_size=4, ep_size=4, ep_plan=EP_PLAN) - self.assertEqual(tensor_parallel.resolve_parallel_plans(self.model, config), (DENSE_TP_PLAN, EP_PLAN)) + self.assertEqual(tensor_parallel.resolve_parallel_plans(self.model, config), (TP_DENSE_PLAN, EP_PLAN)) def test_unmatched_override_keys_raise_without_changing_plans(self): original_tp_plan, original_ep_plan = self.model.tp_plan.copy(), self.model.ep_plan.copy() @@ -204,7 +210,7 @@ def test_override_keys_can_match_modules_parameters_or_existing_plan_keys(self): def test_head_model_overrides_need_the_model_prefix(self): with torch.device("meta"): model = Qwen3MoeForCausalLM(self.config) - config = DistributedConfig(tp_size=4, ep_size=4, ep_plan={"layers.*.mlp.gate": "ep_router"}) + config = DistributedConfig(tp_size=4, ep_size=4, ep_plan=EP_PLAN_MASKED) with self.assertRaisesRegex(ValueError, "match nothing in Qwen3MoeForCausalLM"): tensor_parallel.resolve_parallel_plans(model, config) @@ -212,16 +218,17 @@ def test_head_model_overrides_need_the_model_prefix(self): tp_size=4, ep_size=4, tp_plan={"model.layers.*.self_attn.q_proj": "colwise_rep"}, - ep_plan={"model.layers.*.mlp.gate": "ep_router"}, + ep_plan={f"model.{k}": v for k, v in EP_PLAN_MASKED.items()}, ) tp_plan, ep_plan = tensor_parallel.resolve_parallel_plans(model, config) - expected_tp_plan = {f"model.{k}": v for k, v in DENSE_TP_PLAN.items()} | {"lm_head": "colwise_gather_output"} + expected_tp_plan = {f"model.{k}": v for k, v in TP_DENSE_PLAN.items()} | {"lm_head": "colwise_gather_output"} self.assertEqual(tp_plan, expected_tp_plan | config.tp_plan) - self.assertEqual(ep_plan, {f"model.{k}": v for k, v in EP_PLAN.items()}) + self.assertEqual(ep_plan, {f"model.{k}": v for k, v in EP_PLAN_MASKED.items()}) def test_masked_ep_shards_and_installs_hooks_on_the_tp_mesh(self): tp_mesh = object() - _, ep_plan = tensor_parallel.resolve_parallel_plans(self.model, DistributedConfig(tp_size=4, ep_size=4)) + config = DistributedConfig(tp_size=4, ep_size=4, ep_plan=EP_PLAN_MASKED) + _, ep_plan = tensor_parallel.resolve_parallel_plans(self.model, config) experts, router = self.model.layers[0].mlp.experts, self.model.layers[0].mlp.gate with ( patch.object(ALL_PARALLEL_STYLES["grouped_gemm"], "validate_param") as validate, From 865d5d96f92d10af1750493a6a92c414cf55b60a Mon Sep 17 00:00:00 2001 From: 3outeille Date: Fri, 18 Sep 2026 17:17:11 +0000 Subject: [PATCH 55/86] clean --- tests/tensor_parallel/test_tensor_parallel.py | 1 - 1 file changed, 1 deletion(-) diff --git a/tests/tensor_parallel/test_tensor_parallel.py b/tests/tensor_parallel/test_tensor_parallel.py index 3fed8106c98d..90d4302bee2b 100644 --- a/tests/tensor_parallel/test_tensor_parallel.py +++ b/tests/tensor_parallel/test_tensor_parallel.py @@ -23,7 +23,6 @@ from transformers.distributed.tensor_parallel import ( ALL_PARALLEL_STYLES, ColwiseParallel, - EpDispatchExpertsParallel, PackedColwiseParallel, PackedRowwiseParallel, RowwiseParallel, From cac1b2b04e3f65a5a2d100bc621a31ffe5a7599f Mon Sep 17 00:00:00 2001 From: 3outeille Date: Sat, 19 Sep 2026 15:42:33 +0000 Subject: [PATCH 56/86] fix --- tests/tensor_parallel/test_tensor_parallel.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/tests/tensor_parallel/test_tensor_parallel.py b/tests/tensor_parallel/test_tensor_parallel.py index 90d4302bee2b..b9e8a3349254 100644 --- a/tests/tensor_parallel/test_tensor_parallel.py +++ b/tests/tensor_parallel/test_tensor_parallel.py @@ -105,7 +105,7 @@ def test_ep_takes_experts_and_router_out_of_tp_plan(self): ): with self.subTest(config=config): tp_plan, ep_plan = tensor_parallel.resolve_parallel_plans(self.model, config) - self.assertEqual(tp_plan, TP_DENSE_PLAN | TP_EXPERT_PLAN) + self.assertEqual(tp_plan, TP_DENSE_PLAN) self.assertEqual(ep_plan, EP_PLAN) def test_legacy_flag_is_an_alias_for_ep_size(self): From c59446d99e27c156186dfb675937b279b544fe2c Mon Sep 17 00:00:00 2001 From: 3outeille Date: Mon, 28 Sep 2026 15:43:24 +0000 Subject: [PATCH 57/86] make check repo --- src/transformers/models/mellum/configuration_mellum.py | 10 +++++----- 1 file changed, 5 insertions(+), 5 deletions(-) diff --git a/src/transformers/models/mellum/configuration_mellum.py b/src/transformers/models/mellum/configuration_mellum.py index 49deff0e2e69..9b3374607297 100644 --- a/src/transformers/models/mellum/configuration_mellum.py +++ b/src/transformers/models/mellum/configuration_mellum.py @@ -64,14 +64,14 @@ class MellumConfig(PreTrainedConfig): "layers.*.mlp.up_proj": "colwise", "layers.*.mlp.down_proj": "rowwise", } - # Expert-only EP plan: only shards MoE experts, not attention. - # Attention is left unsharded — FSDP2 handles attention weight distribution. - # This allows EP to scale beyond num_kv_heads (not constrained by 4 for Qwen3-30B). + # Token dispatch by default, so `ep_size` can exceed `tp_size` (bounded by `num_key_value_heads`, 4 on + # Qwen3-30B): with `tp_size=1` the attention is left to FSDP2. For router masking with all-reduce, set + # `ep_size=tp_size` and pass `ep_plan={"layers.*.mlp.gate": "ep_router", "layers.*.mlp.experts": "moe_tp_experts"}` + # (prefixed with `model.` on the causal LM). base_model_ep_plan = { - "layers.*.mlp.gate": "ep_router", "layers.*.mlp.experts.gate_up_proj": "grouped_gemm", "layers.*.mlp.experts.down_proj": "grouped_gemm", - "layers.*.mlp.experts": "moe_tp_experts", + "layers.*.mlp.experts": "ep_dispatch_experts", } base_model_pp_plan = { "embed_tokens": (["input_ids"], ["inputs_embeds"]), From abf336cc17354ef1cb19a183e9b31602dfc63f31 Mon Sep 17 00:00:00 2001 From: 3outeille Date: Mon, 28 Sep 2026 18:17:45 +0000 Subject: [PATCH 58/86] Drop the load-time ep_size == tp_size guard that rejects token dispatch The check ran before the expert plan was resolved, so it rejected every token-dispatch layout (ep_size != tp_size). DistributedConfig._validate_resolved_ep_plan already enforces ep_size == tp_size for the masked plan once the plan is known. --- src/transformers/distributed/mixin.py | 5 ----- 1 file changed, 5 deletions(-) diff --git a/src/transformers/distributed/mixin.py b/src/transformers/distributed/mixin.py index 584e2e5a6fb4..36a5b9c246f8 100644 --- a/src/transformers/distributed/mixin.py +++ b/src/transformers/distributed/mixin.py @@ -161,11 +161,6 @@ def prepare_distribute_model( if isinstance(distributed_config, dict): distributed_config = DistributedConfig.from_dict(distributed_config) - if distributed_config.ep_size > 1 and distributed_config.ep_size != distributed_config.tp_size: - raise ValueError( - "All-reduce expert parallelism requires `ep_size=tp_size` and identical tokens per EP group." - ) - if distributed_config.tp_size == 1 and distributed_config.fsdp_size == 1 and distributed_config.pp_size == 1: return distributed_config, device_map, None From 6244a1c38b7a43d378ad4804fc872347bd40ad7c Mon Sep 17 00:00:00 2001 From: 3outeille Date: Tue, 29 Sep 2026 12:53:15 +0000 Subject: [PATCH 59/86] remove all to all warning by using functional version --- .../distributed/tensor_parallel.py | 19 +++---------------- 1 file changed, 3 insertions(+), 16 deletions(-) diff --git a/src/transformers/distributed/tensor_parallel.py b/src/transformers/distributed/tensor_parallel.py index d6c14907efe6..c8f5fbecc577 100644 --- a/src/transformers/distributed/tensor_parallel.py +++ b/src/transformers/distributed/tensor_parallel.py @@ -38,7 +38,7 @@ if is_torch_distributed_available(): import torch.distributed as dist - from torch.distributed.nn.functional import all_to_all_single + from torch.distributed._functional_collectives import all_to_all_single from torch.distributed.tensor import DTensor, Partial, Replicate, Shard, distribute_tensor from torch.distributed.tensor.placement_types import _StridedShard @@ -789,7 +789,6 @@ def _dispatch_tokens( Also returns the sort order and the per-rank split sizes that `_combine_tokens` needs to reverse the exchange. """ - hidden_dim = hidden_states.size(-1) num_top_k = top_k_index.size(-1) # Sorting the selected pairs by expert groups them by owner rank, since each rank owns a contiguous range of @@ -805,13 +804,7 @@ def _dispatch_tokens( recv_counts = torch.empty_like(send_counts) torch.distributed.all_to_all_single(recv_counts, send_counts, group=ep_group) send_sizes, recv_sizes = torch.stack([send_counts.sum(dim=1), recv_counts.sum(dim=1)]).tolist() - recv_tokens = all_to_all_single( - send_tokens.new_empty(sum(recv_sizes), hidden_dim), - send_tokens, - output_split_sizes=recv_sizes, - input_split_sizes=send_sizes, - group=ep_group, - ) + recv_tokens = all_to_all_single(send_tokens, recv_sizes, send_sizes, ep_group) recv_expert_ids = torch.arange(num_local_experts, device=hidden_states.device).repeat(ep_size) recv_expert_ids = recv_expert_ids.repeat_interleave(recv_counts.reshape(-1), output_size=sum(recv_sizes)) return recv_tokens, recv_expert_ids, order, send_sizes, recv_sizes @@ -845,13 +838,7 @@ def _combine_tokens( """Return expert outputs to the token owners and combine them with routing weights.""" num_tokens, num_top_k = top_k_weights.shape hidden_dim = expert_output.size(-1) - recv_out = all_to_all_single( - expert_output.new_empty(order.numel(), hidden_dim), - expert_output, - output_split_sizes=send_sizes, - input_split_sizes=recv_sizes, - group=ep_group, - ) + recv_out = all_to_all_single(expert_output, send_sizes, recv_sizes, ep_group) # Restore the original (token, top-k slot) order, then apply routing weights. token_outputs = torch.empty_like(recv_out) token_outputs[order] = recv_out From c72bbf72c4c087dc89cc9c2af7f5f576ac1cce27 Mon Sep 17 00:00:00 2001 From: 3outeille Date: Wed, 30 Sep 2026 22:50:55 +0000 Subject: [PATCH 60/86] use ep plan instead of group_gemm is_expert --- src/transformers/distributed/fsdp.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/src/transformers/distributed/fsdp.py b/src/transformers/distributed/fsdp.py index d651ae61f398..60366b6ceb87 100644 --- a/src/transformers/distributed/fsdp.py +++ b/src/transformers/distributed/fsdp.py @@ -19,7 +19,7 @@ from ..utils import is_torch_available, is_torch_distributed_available, is_torch_greater_or_equal, logging, strtobool from ..utils.quantization_config import QuantizationMethod -from .tensor_parallel import replace_layer_number_by_wildcard +from .tensor_parallel import _get_parameter_plan, replace_layer_number_by_wildcard from .utils import _is_torch_distributed_initialized @@ -210,8 +210,8 @@ def apply_fully_sharded_data_parallelism(model: nn.Module, mesh_manager: MeshMan fsdp_policy_kwargs = _get_fsdp_policy_kwargs(distributed_config) if distributed_config.ep_size > 1 and "ep_dispatch_experts" in model.ep_plan.values(): expert_mesh = mesh_manager.get_mesh("efsdp") - for module in model.modules(): - if getattr(module, "_is_expert_parallel", False): + for module_name, module in model.named_modules(): + if _get_parameter_plan(module_name, model.ep_plan, is_weight=False) == "ep_dispatch_experts": fully_shard(module, mesh=expert_mesh, reshard_after_forward=True, **fsdp_policy_kwargs) # An expert group spans several data-parallel batches, so an expert's gradient sums over # all of them. FSDP2 would divide by the efsdp group size; dividing by fsdp_size instead From c31a1b7a9f65e1ef311469d4fa87b19cfe6a7f7c Mon Sep 17 00:00:00 2001 From: 3outeille Date: Wed, 30 Sep 2026 23:35:39 +0000 Subject: [PATCH 61/86] use dtensor instead of manuall all_reduce --- .../distributed/tensor_parallel.py | 36 ++++++++++--------- 1 file changed, 20 insertions(+), 16 deletions(-) diff --git a/src/transformers/distributed/tensor_parallel.py b/src/transformers/distributed/tensor_parallel.py index c8f5fbecc577..c7d7b12bf134 100644 --- a/src/transformers/distributed/tensor_parallel.py +++ b/src/transformers/distributed/tensor_parallel.py @@ -853,30 +853,34 @@ def transform_inputs_pre_forward(self, module, args, kwargs, mesh, *, tp_mesh=No top_k_weights = top_k_weights.to_local() if tp_mesh is None or tp_mesh.size() == 1: return (hidden_states, top_k_index, top_k_weights), kwargs - - tp_group = tp_mesh.get_group() - hidden_states = _AllReduceBackward.apply(hidden_states, tp_group) - top_k_weights = _AllReduceBackward.apply(top_k_weights, tp_group) - # TP ranks share the same batch. Slice tokens here so each is dispatched only once, then restore the full output in the post hook - num_tokens, tp_rank, tp_size = hidden_states.size(0), tp_mesh.get_local_rank(), tp_mesh.size() - rows = slice(num_tokens * tp_rank // tp_size, num_tokens * (tp_rank + 1) // tp_size) - return (hidden_states[rows], top_k_index[rows], top_k_weights[rows]), kwargs + # TP ranks share the same batch, so keep only this rank's rows and each token is dispatched once. + # Replicate -> Shard(0) is a local chunk (no communication); its backward all-gathers the row gradients. + hidden_states = DTensor.from_local(hidden_states, tp_mesh, [Replicate()], run_check=False) + top_k_index = DTensor.from_local(top_k_index, tp_mesh, [Replicate()], run_check=False) + top_k_weights = DTensor.from_local(top_k_weights, tp_mesh, [Replicate()], run_check=False) + + hidden_states = hidden_states.redistribute(tp_mesh, [Shard(0)]).to_local() + top_k_index = top_k_index.redistribute(tp_mesh, [Shard(0)]).to_local() + top_k_weights = top_k_weights.redistribute(tp_mesh, [Shard(0)]).to_local() + return (hidden_states, top_k_index, top_k_weights), kwargs def transform_output_post_forward(self, module, output, mesh, *, tp_mesh=None, num_tokens=None): if tp_mesh is None or tp_mesh.size() == 1: return output - tp_rank, tp_size = tp_mesh.get_local_rank(), tp_mesh.size() - rows = slice(num_tokens * tp_rank // tp_size, num_tokens * (tp_rank + 1) // tp_size) - full_output = output.new_zeros(num_tokens, output.size(-1)) - full_output[rows] = output - return _AllReduceForward.apply(full_output, tp_mesh.get_group()) + # Shard(0) -> Replicate is one all-gather of the row slices, which also handles uneven and empty slices. + hidden_dim = output.size(-1) + output = DTensor.from_local( + output.contiguous(), tp_mesh, [Shard(0)], shape=(num_tokens, hidden_dim), stride=(hidden_dim, 1) + ) + return output.full_tensor() def install_forward(self, module, ep_mesh, *, tp_mesh=None): + """Experts stay whole (Shard(0) on `ep_mesh`); `tp_mesh` is only the group of ranks holding the same batch.""" experts_forward = module.forward ep_group, ep_size = ep_mesh.get_group(), ep_mesh.size() - def tp_forward(hidden_states, top_k_index, top_k_weights): - # Read the full token count before the pre hook slices the inputs on TP. + def ep_forward(hidden_states, top_k_index, top_k_weights): + # Read the full token count before the pre hook slices the inputs across the batch replicas. num_tokens = hidden_states.size(0) (hidden_states, top_k_index, top_k_weights), _ = self.transform_inputs_pre_forward( module, (hidden_states, top_k_index, top_k_weights), {}, ep_mesh, tp_mesh=tp_mesh @@ -892,7 +896,7 @@ def tp_forward(hidden_states, top_k_index, top_k_weights): return self.transform_output_post_forward(module, output, ep_mesh, tp_mesh=tp_mesh, num_tokens=num_tokens) - module.forward = tp_forward + module.forward = ep_forward return module From 00ee316a1e48df63b031adfbb25d5535d4d8c0dc Mon Sep 17 00:00:00 2001 From: 3outeille Date: Thu, 1 Oct 2026 00:10:58 +0000 Subject: [PATCH 62/86] better comment --- src/transformers/distributed/fsdp.py | 10 +++++++--- 1 file changed, 7 insertions(+), 3 deletions(-) diff --git a/src/transformers/distributed/fsdp.py b/src/transformers/distributed/fsdp.py index 60366b6ceb87..7d9a73223101 100644 --- a/src/transformers/distributed/fsdp.py +++ b/src/transformers/distributed/fsdp.py @@ -213,9 +213,13 @@ def apply_fully_sharded_data_parallelism(model: nn.Module, mesh_manager: MeshMan for module_name, module in model.named_modules(): if _get_parameter_plan(module_name, model.ep_plan, is_weight=False) == "ep_dispatch_experts": fully_shard(module, mesh=expert_mesh, reshard_after_forward=True, **fsdp_policy_kwargs) - # An expert group spans several data-parallel batches, so an expert's gradient sums over - # all of them. FSDP2 would divide by the efsdp group size; dividing by fsdp_size instead - # gives the same per-batch average the dense modules get on the fsdp mesh, even when efsdp has a single rank. + # Dense parameters on `fsdp` get a per-batch average: gradients summed over `fsdp_size` batches, + # then divided by `fsdp_size`. Experts must match that scale: + # - an EP group holds `ep_size / tp_size` distinct batches (TP ranks share a batch, and each token is + # dispatched exactly once), so an expert's local gradient already sums over that many batches; + # - the `efsdp` reduce then sums `efsdp_size` copies of that expert, one per EP group. + # The expert gradient therefore covers `ep_size / tp_size * efsdp_size = fsdp_size` batches, so divide + # by `fsdp_size` rather than `efsdp_size`. module.set_gradient_divide_factor(float(distributed_config.fsdp_size)) if torch.distributed.get_backend(expert_mesh.get_group()) != "nccl": # Non-NCCL backends need to sum first, then apply the division otherwise it runtime error. From ec8278e1a991376ae6d71ac8a6e2c1c7068a4793 Mon Sep 17 00:00:00 2001 From: 3outeille Date: Thu, 1 Oct 2026 13:50:13 +0000 Subject: [PATCH 63/86] linting --- src/transformers/distributed/fsdp.py | 12 +++++------- src/transformers/distributed/tensor_parallel.py | 2 +- 2 files changed, 6 insertions(+), 8 deletions(-) diff --git a/src/transformers/distributed/fsdp.py b/src/transformers/distributed/fsdp.py index 7d9a73223101..dbc756111dcb 100644 --- a/src/transformers/distributed/fsdp.py +++ b/src/transformers/distributed/fsdp.py @@ -213,13 +213,11 @@ def apply_fully_sharded_data_parallelism(model: nn.Module, mesh_manager: MeshMan for module_name, module in model.named_modules(): if _get_parameter_plan(module_name, model.ep_plan, is_weight=False) == "ep_dispatch_experts": fully_shard(module, mesh=expert_mesh, reshard_after_forward=True, **fsdp_policy_kwargs) - # Dense parameters on `fsdp` get a per-batch average: gradients summed over `fsdp_size` batches, - # then divided by `fsdp_size`. Experts must match that scale: - # - an EP group holds `ep_size / tp_size` distinct batches (TP ranks share a batch, and each token is - # dispatched exactly once), so an expert's local gradient already sums over that many batches; - # - the `efsdp` reduce then sums `efsdp_size` copies of that expert, one per EP group. - # The expert gradient therefore covers `ep_size / tp_size * efsdp_size = fsdp_size` batches, so divide - # by `fsdp_size` rather than `efsdp_size`. + # Dense parameters on fsdp get a per-batch average: gradients summed over fsdp_size batches, + # then divided by fsdp_size. Experts must match that scale: + # - an EP group holds ep_size / tp_size distinct batches + # - the efsdp reduce then sums efsdp_size copies of that expert. + # The expert gradient therefore covers ep_size / tp_size * efsdp_size = fsdp_size batches module.set_gradient_divide_factor(float(distributed_config.fsdp_size)) if torch.distributed.get_backend(expert_mesh.get_group()) != "nccl": # Non-NCCL backends need to sum first, then apply the division otherwise it runtime error. diff --git a/src/transformers/distributed/tensor_parallel.py b/src/transformers/distributed/tensor_parallel.py index c7d7b12bf134..5690addbc1da 100644 --- a/src/transformers/distributed/tensor_parallel.py +++ b/src/transformers/distributed/tensor_parallel.py @@ -858,7 +858,7 @@ def transform_inputs_pre_forward(self, module, args, kwargs, mesh, *, tp_mesh=No hidden_states = DTensor.from_local(hidden_states, tp_mesh, [Replicate()], run_check=False) top_k_index = DTensor.from_local(top_k_index, tp_mesh, [Replicate()], run_check=False) top_k_weights = DTensor.from_local(top_k_weights, tp_mesh, [Replicate()], run_check=False) - + hidden_states = hidden_states.redistribute(tp_mesh, [Shard(0)]).to_local() top_k_index = top_k_index.redistribute(tp_mesh, [Shard(0)]).to_local() top_k_weights = top_k_weights.redistribute(tp_mesh, [Shard(0)]).to_local() From 39c7e41386fb6947e1f20289cc75c4ffbf28bc9b Mon Sep 17 00:00:00 2001 From: 3outeille Date: Wed, 16 Sep 2026 15:27:01 +0000 Subject: [PATCH 64/86] Train models sharded at load time with the Trainer, token dispatch included A model wrapped by its own FSDP2 (`DistributedConfig` with `fsdp_size > 1` or expert-parallel token dispatch) already owns placement and gradient reduction. Prepare it with `accelerator.prepare_model(..., evaluation_mode=True)` so Accelerate applies autocast and compilation without wrapping the DTensor parameters in DDP, which it rejects, or sharding them again. This is what token dispatch with `tp_size=1` needed: there is no `ParallelismConfig` to tell Accelerate about the model's TP, so it fell into the DDP branch. Loss scaling and token counting derive the replication factor from `get_tp_size()`, which now also falls back to Accelerate's parallelism config for tensor parallelism configured outside `DistributedConfig`. `tests/trainer/distributed/test_trainer_distributed_expert_parallel.py` trains a tiny Qwen3 MoE for four steps under masking on a `(fsdp, tp)` mesh, token dispatch with a batch per rank, and token dispatch with TP pairs, each with the same global batch as a single-process run, and compares the logged losses and gradient norms and the saved weights. --- docs/source/en/expert_parallelism.md | 3 +- src/transformers/trainer.py | 31 +++-- .../scripts/expert_parallel_train.py | 110 ++++++++++++++++++ ...est_trainer_distributed_expert_parallel.py | 106 +++++++++++++++++ 4 files changed, 239 insertions(+), 11 deletions(-) create mode 100644 tests/trainer/distributed/scripts/expert_parallel_train.py create mode 100644 tests/trainer/distributed/test_trainer_distributed_expert_parallel.py diff --git a/docs/source/en/expert_parallelism.md b/docs/source/en/expert_parallelism.md index 61d646c56f23..3ee76f533c2a 100644 --- a/docs/source/en/expert_parallelism.md +++ b/docs/source/en/expert_parallelism.md @@ -160,8 +160,7 @@ These configurations each use eight GPUs: | `DistributedConfig(tp_size=2, fsdp_size=4, ep_size=4)` | Dispatch with a TP pair per batch, each pair slicing its batch in two; two batches per expert group. | | `DistributedConfig(tp_size=8, ep_size=8)` | Dispatch with every rank sharing one batch, sliced in eight, or masking and all-reduce for a masked plan. | -> [!WARNING] -> The [`Trainer`] does not account for token dispatch yet: batch and token counting assume the all-reduce layout, where the ranks of a TP group share a batch and `fsdp_size` data-parallel shards exist. Trainer support for dispatch comes in a follow-up. +The [`Trainer`] trains these layouts as loaded: it leaves the placement and gradient reduction to the model's own FSDP2 and expert-parallel wrappers instead of wrapping it again, gives each device mesh its own optimizer param group and gradient-norm term, and counts tokens once per TP group. The effective global batch size is `per_device_train_batch_size * fsdp_size * gradient_accumulation_steps`, whichever the layout. [`~Trainer.save_model`] gathers the sharded weights into a regular checkpoint. ## Combining with FSDP2 diff --git a/src/transformers/trainer.py b/src/transformers/trainer.py index bfc759c5692c..57f84725fd96 100755 --- a/src/transformers/trainer.py +++ b/src/transformers/trainer.py @@ -1723,7 +1723,17 @@ def _prepare_for_training(self, max_steps, train_dataloader, resume_from_checkpo use_accelerator_prepare = model is self.model # prepare using `accelerator` prepare - if use_accelerator_prepare: + if ( + getattr(model, "_is_fsdp_managed_module", False) + and not self.is_fsdp_enabled + and not self.is_deepspeed_enabled + ): + # Sharded at load time (`DistributedConfig` with FSDP2 or expert-parallel token dispatch): the model + # already owns placement and gradient reduction. Prepare autocast and compilation only, without + # asking Accelerate to wrap the DTensor parameters in DDP or to shard them again. + model = self.accelerator.prepare_model(model, device_placement=False, evaluation_mode=True) + self.optimizer = self.accelerator.prepare(self.optimizer) + elif use_accelerator_prepare: if delay_optimizer_creation: # TODO: check if we can move this somewhere else if self.is_fsdp_enabled and _is_peft_model(self.model): @@ -2174,11 +2184,9 @@ def compute_loss( and (self.model_accepts_loss_kwargs or self.compute_loss_func) and num_items_in_batch is not None ): - # TP and EP-as-TP ranks see replicated batches; `num_processes` over-counts - # them by `tp_size`. Mirror the divisor used in `_get_num_items_in_batch`. - loss_scale = self.accelerator.num_processes - if (pc := getattr(self.accelerator, "parallelism_config", None)) is not None: - loss_scale //= pc.tp_size + # TP ranks (and the expert-parallel ranks sharing their batch) see replicated batches; `num_processes` + # over-counts them by `tp_size`. Mirror the divisor used in `_get_num_items_in_batch`. + loss_scale = self.accelerator.num_processes // self.get_tp_size() loss *= loss_scale if self.args.n_gpu <= 1 else self.args.n_gpu return (loss, outputs) if return_outputs else loss @@ -2327,8 +2335,9 @@ def _get_num_items_in_batch(self, batch_samples: list, device: torch.device) -> # In the DataParallel case, convert the scalar tensor into a 2-dim tensor with the same value repeated num_items_in_batch = num_items_in_batch.unsqueeze(0).expand(self.args.n_gpu, -1) # Divide by number of devices with the same batch - if pc := getattr(self.accelerator, "parallelism_config", None): - num_items_in_batch = num_items_in_batch // pc.non_data_parallel_size + num_items_in_batch = num_items_in_batch // ( + self.get_tp_size() * self.get_cp_size() * self.get_sp_size() + ) return num_items_in_batch @@ -2573,7 +2582,11 @@ def get_tp_size(self) -> int: if self.is_deepspeed_enabled and (deepspeed_config := getattr(self.args, "hf_deepspeed_config", None)): return deepspeed_config.config.get("tensor_parallel", {}).get("autotp_size", 1) - # 3. Default fallback + # 3. Fall back to accelerate, for tensor parallelism configured outside `DistributedConfig` + if (pc := getattr(self.accelerator, "parallelism_config", None)) is not None: + return pc.tp_size + + # 4. Default fallback return 1 def _wrap_model(self, model: nn.Module, training: bool = True, dataloader: DataLoader | None = None) -> nn.Module: diff --git a/tests/trainer/distributed/scripts/expert_parallel_train.py b/tests/trainer/distributed/scripts/expert_parallel_train.py new file mode 100644 index 000000000000..e1a5ae8b637e --- /dev/null +++ b/tests/trainer/distributed/scripts/expert_parallel_train.py @@ -0,0 +1,110 @@ +# Copyright 2026 The HuggingFace Team. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +""" +Worker script for the expert-parallel `Trainer` tests: train a tiny MoE for a few steps under one layout. + +Launched via ``torchrun`` from ``test_trainer_distributed_expert_parallel.py``. Every layout consumes the same +global batch per step, so the logged losses and gradient norms and the saved weights must match the single-process +``reference`` run. +""" + +import argparse +import json +import os + +import torch +from torch.utils.data import Dataset + +from transformers import AutoModelForCausalLM, Trainer, TrainingArguments +from transformers.distributed import DistributedConfig + + +GLOBAL_BATCH_SIZE = 8 +SEQ_LEN = 16 +NUM_STEPS = 4 + +# `world_size=4` layouts. `dp` is the number of distinct batches per step; each rank trains on +# `GLOBAL_BATCH_SIZE // dp` samples. The masked plan is Qwen3 MoE's dispatch plan with the two forward rules +# overridden; dispatch uses the model's default plan. +MASKED_PLAN = {"model.layers.*.mlp.gate": "ep_router", "model.layers.*.mlp.experts": "moe_tp_experts"} +LAYOUTS = { + "reference": {"world_size": 1, "dp": 1, "config": None}, + "masked": { + "world_size": 4, + "dp": 2, + "config": {"tp_size": 2, "fsdp_size": 2, "ep_size": 2, "ep_plan": MASKED_PLAN}, + }, + "dispatch": {"world_size": 4, "dp": 4, "config": {"tp_size": 1, "fsdp_size": 4, "ep_size": 2}}, + "dispatch_tp": {"world_size": 4, "dp": 2, "config": {"tp_size": 2, "fsdp_size": 2, "ep_size": 4}}, +} + + +class TokenDataset(Dataset): + def __init__(self, vocab_size: int): + generator = torch.Generator().manual_seed(0) + self.input_ids = torch.randint(0, vocab_size, (NUM_STEPS * GLOBAL_BATCH_SIZE, SEQ_LEN), generator=generator) + + def __len__(self): + return self.input_ids.size(0) + + def __getitem__(self, index): + return {"input_ids": self.input_ids[index], "labels": self.input_ids[index].clone()} + + +def main(): + parser = argparse.ArgumentParser() + parser.add_argument("--layout", choices=LAYOUTS, required=True) + parser.add_argument("--model_dir", required=True) + parser.add_argument("--output_dir", required=True) + args = parser.parse_args() + layout = LAYOUTS[args.layout] + assert int(os.environ.get("WORLD_SIZE", 1)) == layout["world_size"], args.layout + + distributed_config = DistributedConfig(**layout["config"]) if layout["config"] else None + model = AutoModelForCausalLM.from_pretrained( + args.model_dir, dtype=torch.float32, distributed_config=distributed_config + ) + if distributed_config is None: + model = model.to("cuda") + + training_args = TrainingArguments( + output_dir=os.path.join(args.output_dir, "trainer"), + per_device_train_batch_size=GLOBAL_BATCH_SIZE // layout["dp"], + max_steps=NUM_STEPS, + learning_rate=1e-3, + max_grad_norm=1.0, + logging_steps=1, + save_strategy="no", + report_to=[], + seed=0, + data_seed=0, + dataloader_drop_last=True, + remove_unused_columns=False, + average_tokens_across_devices=True, + disable_tqdm=True, + ) + trainer = Trainer(model=model, args=training_args, train_dataset=TokenDataset(model.config.vocab_size)) + trainer.train() + # Gathers the sharded weights; every rank takes part, the main process writes. + trainer.save_model(os.path.join(args.output_dir, "model")) + + if trainer.is_world_process_zero(): + steps = [log for log in trainer.state.log_history if "loss" in log] + with open(os.path.join(args.output_dir, "results.json"), "w") as f: + json.dump({"loss": [s["loss"] for s in steps], "grad_norm": [s["grad_norm"] for s in steps]}, f) + + +if __name__ == "__main__": + main() diff --git a/tests/trainer/distributed/test_trainer_distributed_expert_parallel.py b/tests/trainer/distributed/test_trainer_distributed_expert_parallel.py new file mode 100644 index 000000000000..f3b64bb53af9 --- /dev/null +++ b/tests/trainer/distributed/test_trainer_distributed_expert_parallel.py @@ -0,0 +1,106 @@ +# Copyright 2026 The HuggingFace Team. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +import json +import os + +from parameterized import parameterized + +from transformers import Qwen3MoeConfig, Qwen3MoeForCausalLM, is_torch_available +from transformers.testing_utils import ( + TestCasePlus, + backend_device_count, + execute_subprocess_async, + get_torch_dist_unique_port, + require_torch_multi_accelerator, + slow, + torch_device, +) + + +if is_torch_available(): + import torch + from safetensors.torch import load_file + + +SCRIPT = os.path.join(os.path.dirname(__file__), "scripts", "expert_parallel_train.py") +LAYOUTS = ("masked", "dispatch", "dispatch_tp") + + +@slow +@require_torch_multi_accelerator +class TestTrainerExpertParallel(TestCasePlus): + """ + `Trainer` with a model sharded at load time by `DistributedConfig`: router masking with all-reduce on a + `(fsdp, tp)` mesh, token dispatch with an independent batch per rank, and token dispatch with TP groups sharing a + batch. Every layout consumes the same global batch per step as the single-process reference, so the logged + losses and gradient norms and the saved weights have to match it. + """ + + def _run(self, layout, model_dir, world_size): + output_dir = self.get_auto_remove_tmp_dir() + cmd = [ + "torchrun", + f"--nproc_per_node={world_size}", + "--nnodes=1", + f"--master_port={get_torch_dist_unique_port()}", + SCRIPT, + f"--layout={layout}", + f"--model_dir={model_dir}", + f"--output_dir={output_dir}", + ] + execute_subprocess_async(cmd, env=self.get_env()) + with open(os.path.join(output_dir, "results.json")) as f: + results = json.load(f) + return results, load_file(os.path.join(output_dir, "model", "model.safetensors")) + + def _model_dir(self): + torch.manual_seed(0) + config = Qwen3MoeConfig( + vocab_size=128, + hidden_size=32, + intermediate_size=64, + moe_intermediate_size=32, + num_hidden_layers=2, + num_attention_heads=4, + num_key_value_heads=2, + head_dim=8, + num_experts=4, + num_experts_per_tok=2, + max_position_embeddings=64, + ) + model_dir = self.get_auto_remove_tmp_dir() + Qwen3MoeForCausalLM(config).save_pretrained(model_dir) + return model_dir + + @parameterized.expand([(layout,) for layout in LAYOUTS]) + def test_matches_single_process_reference(self, layout): + if backend_device_count(torch_device) < 4: + self.skipTest("Requires 4 accelerators") + model_dir = self._model_dir() + reference, reference_weights = self._run("reference", model_dir, world_size=1) + results, weights = self._run(layout, model_dir, world_size=4) + + self.assertEqual(len(results["loss"]), len(reference["loss"])) + torch.testing.assert_close( + torch.tensor(results["loss"]), torch.tensor(reference["loss"]), rtol=1e-4, atol=1e-4 + ) + torch.testing.assert_close( + torch.tensor(results["grad_norm"]), torch.tensor(reference["grad_norm"]), rtol=1e-3, atol=1e-4 + ) + self.assertEqual(set(weights), set(reference_weights)) + # Adam turns rounding-level gradient differences into `lr`-sized weight differences where the gradient is + # near zero, so the weights get a tolerance above the learning rate. + for key, tensor in reference_weights.items(): + torch.testing.assert_close(weights[key], tensor, rtol=1e-3, atol=2e-3, msg=lambda m: f"{key}: {m}") From dfec8d1ef72dbd248e1598cf67c13cd4da28d542 Mon Sep 17 00:00:00 2001 From: 3outeille Date: Fri, 25 Sep 2026 15:58:47 +0000 Subject: [PATCH 65/86] Use ep-dispatch version of expert_parallelism.md --- docs/source/en/expert_parallelism.md | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/docs/source/en/expert_parallelism.md b/docs/source/en/expert_parallelism.md index 3ee76f533c2a..61d646c56f23 100644 --- a/docs/source/en/expert_parallelism.md +++ b/docs/source/en/expert_parallelism.md @@ -160,7 +160,8 @@ These configurations each use eight GPUs: | `DistributedConfig(tp_size=2, fsdp_size=4, ep_size=4)` | Dispatch with a TP pair per batch, each pair slicing its batch in two; two batches per expert group. | | `DistributedConfig(tp_size=8, ep_size=8)` | Dispatch with every rank sharing one batch, sliced in eight, or masking and all-reduce for a masked plan. | -The [`Trainer`] trains these layouts as loaded: it leaves the placement and gradient reduction to the model's own FSDP2 and expert-parallel wrappers instead of wrapping it again, gives each device mesh its own optimizer param group and gradient-norm term, and counts tokens once per TP group. The effective global batch size is `per_device_train_batch_size * fsdp_size * gradient_accumulation_steps`, whichever the layout. [`~Trainer.save_model`] gathers the sharded weights into a regular checkpoint. +> [!WARNING] +> The [`Trainer`] does not account for token dispatch yet: batch and token counting assume the all-reduce layout, where the ranks of a TP group share a batch and `fsdp_size` data-parallel shards exist. Trainer support for dispatch comes in a follow-up. ## Combining with FSDP2 From 628da80a5bf38917a2ec312c2925cc2000c41bb0 Mon Sep 17 00:00:00 2001 From: 3outeille Date: Fri, 25 Sep 2026 16:05:31 +0000 Subject: [PATCH 66/86] Sync replicated trainable parameters when training a model sharded at load time A PEFT adapter attached after fully_shard is a plain tensor next to DTensor base weights. FSDP2 only reduces what it sharded and the DDP wrap is skipped on this path, so each rank trained its own adapter. Broadcast these parameters from rank 0 and average their gradient at the end of each accumulation window. --- src/transformers/trainer.py | 31 +++++++++++++++++++++++++++++++ 1 file changed, 31 insertions(+) diff --git a/src/transformers/trainer.py b/src/transformers/trainer.py index 57f84725fd96..7ca56548d061 100755 --- a/src/transformers/trainer.py +++ b/src/transformers/trainer.py @@ -1733,6 +1733,7 @@ def _prepare_for_training(self, max_steps, train_dataloader, resume_from_checkpo # asking Accelerate to wrap the DTensor parameters in DDP or to shard them again. model = self.accelerator.prepare_model(model, device_placement=False, evaluation_mode=True) self.optimizer = self.accelerator.prepare(self.optimizer) + self._sync_replicated_trainable_parameters(model) elif use_accelerator_prepare: if delay_optimizer_creation: # TODO: check if we can move this somewhere else @@ -2589,6 +2590,36 @@ def get_tp_size(self) -> int: # 4. Default fallback return 1 + def _sync_replicated_trainable_parameters(self, model: nn.Module) -> None: + """ + Keep the trainable parameters that FSDP2 does not manage identical across ranks. + + A parameter added after `fully_shard` (a PEFT adapter attached to a model sharded at load time) stays a plain + replicated tensor next to the DTensor base weights. FSDP2 reduce-scatters only what it sharded and the DDP + wrap is skipped for such a model, so each rank would init its own copy and train it on its own batch. + Broadcast these parameters from rank 0, then average their gradient over all ranks at the end of each + accumulation window. + """ + if not dist.is_available() or not dist.is_initialized() or dist.get_world_size() == 1: + return + + from torch.distributed.tensor import DTensor + + def average_gradient(param): + # Averaging the accumulated micro-batch gradients once gives the same result as after every backward. + if self.accelerator.sync_gradients: + dist.all_reduce(param.grad, op=dist.ReduceOp.AVG) + + for param in model.parameters(): + if not param.requires_grad or isinstance(param.data, DTensor): + continue + with torch.no_grad(): + dist.broadcast(param.data, src=0) + # `train()` can run more than once on the same model: register the hook only once. + if not getattr(param, "_replicated_grad_hook_registered", False): + param.register_post_accumulate_grad_hook(average_gradient) + param._replicated_grad_hook_registered = True + def _wrap_model(self, model: nn.Module, training: bool = True, dataloader: DataLoader | None = None) -> nn.Module: """Wrap `model` for distributed training if needed (DDP, FSDP, SageMaker, etc.).""" # train/eval could be run multiple-times - if already wrapped, don't re-wrap it again From e72d719bb8f1406c63e9dfaf53b66ac42cdcd949 Mon Sep 17 00:00:00 2001 From: 3outeille Date: Mon, 28 Sep 2026 16:41:44 +0000 Subject: [PATCH 67/86] Fix utf-8 encoding in expert parallel trainer tests --- tests/trainer/distributed/scripts/expert_parallel_train.py | 2 +- .../distributed/test_trainer_distributed_expert_parallel.py | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/tests/trainer/distributed/scripts/expert_parallel_train.py b/tests/trainer/distributed/scripts/expert_parallel_train.py index e1a5ae8b637e..6e02544103e6 100644 --- a/tests/trainer/distributed/scripts/expert_parallel_train.py +++ b/tests/trainer/distributed/scripts/expert_parallel_train.py @@ -102,7 +102,7 @@ def main(): if trainer.is_world_process_zero(): steps = [log for log in trainer.state.log_history if "loss" in log] - with open(os.path.join(args.output_dir, "results.json"), "w") as f: + with open(os.path.join(args.output_dir, "results.json"), "w", encoding="utf-8") as f: json.dump({"loss": [s["loss"] for s in steps], "grad_norm": [s["grad_norm"] for s in steps]}, f) diff --git a/tests/trainer/distributed/test_trainer_distributed_expert_parallel.py b/tests/trainer/distributed/test_trainer_distributed_expert_parallel.py index f3b64bb53af9..2e4bb051a60a 100644 --- a/tests/trainer/distributed/test_trainer_distributed_expert_parallel.py +++ b/tests/trainer/distributed/test_trainer_distributed_expert_parallel.py @@ -61,7 +61,7 @@ def _run(self, layout, model_dir, world_size): f"--output_dir={output_dir}", ] execute_subprocess_async(cmd, env=self.get_env()) - with open(os.path.join(output_dir, "results.json")) as f: + with open(os.path.join(output_dir, "results.json"), encoding="utf-8") as f: results = json.load(f) return results, load_file(os.path.join(output_dir, "model", "model.safetensors")) From e9bcdc57702d4b6a18b7fdb2e32c33e180e5573f Mon Sep 17 00:00:00 2001 From: 3outeille Date: Mon, 28 Sep 2026 17:58:00 +0000 Subject: [PATCH 68/86] [distributed] Default MoE expert-parallel plans to token dispatch Switch `base_model_ep_plan` from router masking with all-reduce (`ep_router` + `moe_tp_experts`) to all-to-all token dispatch (`ep_dispatch_experts`) for 41 MoE models, as done for Qwen3 MoE. Dispatch finds each expert's owner from the global expert ids, so the router rule is dropped. Models kept on masking: - Llama 4: no experts rule, non-standard experts module. - gemma4 (and diffusion_gemma), granitemoe_swa, hy_v4, inkling, mimo_v2_flash, youtu, zaya: their EP mixin tests already fail with masking because TP-sharded parameters used outside a TP style (attention sinks, norm weights, per-layer embeddings) mix Tensor and DTensor once TP composes with FSDP2, which dispatch always applies. Switch them once that composition is fixed. --- docs/source/en/expert_parallelism.md | 8 ++++---- src/transformers/models/afmoe/configuration_afmoe.py | 3 +-- src/transformers/models/axk1/configuration_axk1.py | 3 +-- src/transformers/models/axk2/configuration_axk2.py | 3 +-- .../models/cohere2_moe/configuration_cohere2_moe.py | 3 +-- .../deepseek_ocr2/configuration_deepseek_ocr2.py | 3 +-- .../models/deepseek_v2/configuration_deepseek_v2.py | 3 +-- .../models/deepseek_v2/modular_deepseek_v2.py | 3 +-- .../models/deepseek_v3/configuration_deepseek_v3.py | 3 +-- .../deepseek_v32/configuration_deepseek_v32.py | 3 +-- .../models/deepseek_v4/configuration_deepseek_v4.py | 12 ++++++------ src/transformers/models/dots1/configuration_dots1.py | 3 +-- src/transformers/models/dots1/modular_dots1.py | 3 +-- .../ernie4_5_moe/configuration_ernie4_5_moe.py | 3 +-- .../ernie4_5_vl_moe/configuration_ernie4_5_vl_moe.py | 3 +-- .../models/exaone_moe/configuration_exaone_moe.py | 3 +-- .../models/exaone_moe/modular_exaone_moe.py | 3 +-- .../models/flex_olmo/configuration_flex_olmo.py | 3 +-- .../models/flex_olmo/modular_flex_olmo.py | 3 +-- .../models/glm4_moe/configuration_glm4_moe.py | 3 +-- src/transformers/models/glm4_moe/modular_glm4_moe.py | 3 +-- .../glm4_moe_lite/configuration_glm4_moe_lite.py | 3 +-- .../models/glm4_moe_lite/modular_glm4_moe_lite.py | 3 +-- .../models/glm4v_moe/configuration_glm4v_moe.py | 3 +-- .../models/glm5_next/configuration_glm5_next.py | 3 +-- .../models/glm_moe_dsa/configuration_glm_moe_dsa.py | 3 +-- .../models/gpt_oss/configuration_gpt_oss.py | 3 +-- .../hunyuan_v1_moe/configuration_hunyuan_v1_moe.py | 3 +-- src/transformers/models/hy_v3/configuration_hy_v3.py | 3 +-- src/transformers/models/hy_v3/modular_hy_v3.py | 3 +-- .../models/kimi_linear/configuration_kimi_linear.py | 3 +-- .../models/laguna/configuration_laguna.py | 3 +-- .../models/lfm2_moe/configuration_lfm2_moe.py | 3 +-- .../models/minimax/configuration_minimax.py | 3 +-- src/transformers/models/minimax/modular_minimax.py | 3 +-- .../models/minimax_m2/configuration_minimax_m2.py | 3 +-- .../models/minimax_m2/modular_minimax_m2.py | 3 +-- .../minimax_m3_vl/configuration_minimax_m3_vl.py | 3 +-- .../models/mistral4/configuration_mistral4.py | 3 +-- .../models/mixtral/configuration_mixtral.py | 3 +-- src/transformers/models/olmoe/configuration_olmoe.py | 3 +-- .../configuration_openai_privacy_filter.py | 3 +-- .../models/phimoe/configuration_phimoe.py | 3 +-- .../models/qwen2_moe/configuration_qwen2_moe.py | 3 +-- .../models/qwen3_5_moe/configuration_qwen3_5_moe.py | 3 +-- .../models/qwen3_next/configuration_qwen3_next.py | 3 +-- .../qwen3_omni_moe/configuration_qwen3_omni_moe.py | 3 +-- .../models/qwen3_omni_moe/modular_qwen3_omni_moe.py | 3 +-- .../qwen3_vl_moe/configuration_qwen3_vl_moe.py | 3 +-- .../models/qwen3_vl_moe/modular_qwen3_vl_moe.py | 3 +-- .../models/qwen4_exp/configuration_qwen4_exp.py | 3 +-- .../models/solar_open/configuration_solar_open.py | 3 +-- .../models/step3p7/configuration_step3p7.py | 3 +-- src/transformers/models/youtu/configuration_youtu.py | 1 + src/transformers/models/youtu/modular_youtu.py | 7 +++++++ src/transformers/models/zaya/configuration_zaya.py | 1 + src/transformers/models/zaya/modular_zaya.py | 7 +++++++ 57 files changed, 77 insertions(+), 112 deletions(-) diff --git a/docs/source/en/expert_parallelism.md b/docs/source/en/expert_parallelism.md index 61d646c56f23..d79e67f465cb 100644 --- a/docs/source/en/expert_parallelism.md +++ b/docs/source/en/expert_parallelism.md @@ -19,7 +19,7 @@ rendered properly in your Markdown viewer. ## DistributedConfig -Enable expert parallelism with the [`DistributedConfig`] class and the `ep_size` argument. Most models route with masking and all-reduce, which requires `ep_size=tp_size` so every rank in an expert group receives the same tokens. Models whose plan uses [token dispatch](#token-dispatch), such as Qwen3 MoE, can set `ep_size` independently of `tp_size`. +Enable expert parallelism with the [`DistributedConfig`] class and the `ep_size` argument. Most MoE models default to [token dispatch](#token-dispatch), so `ep_size` can be set independently of `tp_size`. A few, such as Llama 4 and Gemma 4, still default to masking and all-reduce (`"ep_router"` and `"moe_tp_experts"` in `model.ep_plan`). Masking is also available on any model with an `ep_plan` override, and requires `ep_size=tp_size` so every rank in an expert group receives the same tokens. ```py import os @@ -74,7 +74,7 @@ distributed_config = DistributedConfig( Providing a plan does not infer parallel sizes: set `tp_size` and `ep_size` explicitly. -Qwen3 MoE defaults to `"ep_dispatch_experts"`. To use masking and all-reduce instead, set `ep_size=tp_size` and override both the router and the expert forward rules: +Most MoE models default to `"ep_dispatch_experts"`. To use masking and all-reduce instead, set `ep_size=tp_size` and override both the router and the expert forward rules (the router module name depends on the model, e.g. `mlp.router` on gpt-oss): ```py distributed_config = DistributedConfig( @@ -87,7 +87,7 @@ distributed_config = DistributedConfig( ) ``` -Conversely, override the expert forward rule of a model whose plan uses masking with `"ep_dispatch_experts"` to use token dispatch. The router rule is then ignored, since dispatch needs the global expert ids to find each expert's owner. +Conversely, when a plan combines `"ep_dispatch_experts"` with an `"ep_router"` rule, the router rule is ignored, since dispatch needs the global expert ids to find each expert's owner. Non-expert rules in the EP plan are ignored too: with dispatch, the EP plan only shards the experts. ## Token dispatch @@ -155,7 +155,7 @@ These configurations each use eight GPUs: | Configuration | Result | | :--- | :--- | -| `DistributedConfig(tp_size=4, fsdp_size=2, ep_size=4)` | Dispatch with TP groups of four, each slicing its batch in four (Qwen3 MoE default plan); unless you specify a ep_plan to use the legacy masked EP | +| `DistributedConfig(tp_size=4, fsdp_size=2, ep_size=4)` | Dispatch with TP groups of four, each slicing its batch in four (default plan); unless you specify a ep_plan to use the legacy masked EP | | `DistributedConfig(tp_size=1, fsdp_size=8, ep_size=4)` | Dispatch with an independent batch on each rank, no slicing, and experts FSDP-sharded across pairs of ranks. | | `DistributedConfig(tp_size=2, fsdp_size=4, ep_size=4)` | Dispatch with a TP pair per batch, each pair slicing its batch in two; two batches per expert group. | | `DistributedConfig(tp_size=8, ep_size=8)` | Dispatch with every rank sharing one batch, sliced in eight, or masking and all-reduce for a masked plan. | diff --git a/src/transformers/models/afmoe/configuration_afmoe.py b/src/transformers/models/afmoe/configuration_afmoe.py index df3980e51917..40083d09e68a 100644 --- a/src/transformers/models/afmoe/configuration_afmoe.py +++ b/src/transformers/models/afmoe/configuration_afmoe.py @@ -62,10 +62,9 @@ class AfmoeConfig(PreTrainedConfig): "norm": (["hidden_states"], ["hidden_states"]), } base_model_ep_plan = { - "layers.*.mlp.router": "ep_router", "layers.*.mlp.experts.gate_up_proj": "grouped_gemm", "layers.*.mlp.experts.down_proj": "grouped_gemm", - "layers.*.mlp.experts": "moe_tp_experts", + "layers.*.mlp.experts": "ep_dispatch_experts", } vocab_size: int = 200192 diff --git a/src/transformers/models/axk1/configuration_axk1.py b/src/transformers/models/axk1/configuration_axk1.py index d034f5bf59da..485434b894f4 100644 --- a/src/transformers/models/axk1/configuration_axk1.py +++ b/src/transformers/models/axk1/configuration_axk1.py @@ -69,10 +69,9 @@ class AXK1Config(PreTrainedConfig): "norm": (["hidden_states"], ["hidden_states"]), } base_model_ep_plan = { - "layers.*.mlp.gate": "ep_router", "layers.*.mlp.experts.gate_up_proj": "grouped_gemm", "layers.*.mlp.experts.down_proj": "grouped_gemm", - "layers.*.mlp.experts": "moe_tp_experts", + "layers.*.mlp.experts": "ep_dispatch_experts", } attribute_map = { "num_local_experts": "n_routed_experts", diff --git a/src/transformers/models/axk2/configuration_axk2.py b/src/transformers/models/axk2/configuration_axk2.py index 4cb2746491a2..2fb7a12ac4f1 100644 --- a/src/transformers/models/axk2/configuration_axk2.py +++ b/src/transformers/models/axk2/configuration_axk2.py @@ -82,10 +82,9 @@ class AXK2Config(PreTrainedConfig): "norm": (["hidden_states"], ["hidden_states"]), } base_model_ep_plan = { - "layers.*.mlp.gate": "ep_router", "layers.*.mlp.experts.gate_up_proj": "grouped_gemm", "layers.*.mlp.experts.down_proj": "grouped_gemm", - "layers.*.mlp.experts": "moe_tp_experts", + "layers.*.mlp.experts": "ep_dispatch_experts", } attribute_map = {"num_local_experts": "n_routed_experts"} diff --git a/src/transformers/models/cohere2_moe/configuration_cohere2_moe.py b/src/transformers/models/cohere2_moe/configuration_cohere2_moe.py index 32ef6fdf66ad..12db6efafaed 100644 --- a/src/transformers/models/cohere2_moe/configuration_cohere2_moe.py +++ b/src/transformers/models/cohere2_moe/configuration_cohere2_moe.py @@ -85,10 +85,9 @@ class Cohere2MoeConfig(PreTrainedConfig): "norm": (["hidden_states"], ["hidden_states"]), } base_model_ep_plan = { - "layers.*.mlp.gate": "ep_router", "layers.*.mlp.experts.gate_up_proj": "grouped_gemm", "layers.*.mlp.experts.down_proj": "grouped_gemm", - "layers.*.mlp.experts": "moe_tp_experts", + "layers.*.mlp.experts": "ep_dispatch_experts", } vocab_size: int = 256000 diff --git a/src/transformers/models/deepseek_ocr2/configuration_deepseek_ocr2.py b/src/transformers/models/deepseek_ocr2/configuration_deepseek_ocr2.py index d12eb50849db..0d61c8465b49 100644 --- a/src/transformers/models/deepseek_ocr2/configuration_deepseek_ocr2.py +++ b/src/transformers/models/deepseek_ocr2/configuration_deepseek_ocr2.py @@ -238,10 +238,9 @@ class DeepseekOcr2TextConfig(PreTrainedConfig): mlp_bias: bool = False head_dim: int | None = None base_model_ep_plan = { - "layers.*.mlp.gate": "ep_router", "layers.*.mlp.experts.gate_up_proj": "grouped_gemm", "layers.*.mlp.experts.down_proj": "grouped_gemm", - "layers.*.mlp.experts": "moe_tp_experts", + "layers.*.mlp.experts": "ep_dispatch_experts", } attribute_map = { "num_experts": "n_routed_experts", diff --git a/src/transformers/models/deepseek_v2/configuration_deepseek_v2.py b/src/transformers/models/deepseek_v2/configuration_deepseek_v2.py index 24d8043e150e..e0124e55856a 100644 --- a/src/transformers/models/deepseek_v2/configuration_deepseek_v2.py +++ b/src/transformers/models/deepseek_v2/configuration_deepseek_v2.py @@ -95,10 +95,9 @@ class DeepseekV2Config(PreTrainedConfig): mlp_bias: bool = False head_dim: int | None = None base_model_ep_plan = { - "layers.*.mlp.gate": "ep_router", "layers.*.mlp.experts.gate_up_proj": "grouped_gemm", "layers.*.mlp.experts.down_proj": "grouped_gemm", - "layers.*.mlp.experts": "moe_tp_experts", + "layers.*.mlp.experts": "ep_dispatch_experts", } attribute_map = { "num_experts": "n_routed_experts", diff --git a/src/transformers/models/deepseek_v2/modular_deepseek_v2.py b/src/transformers/models/deepseek_v2/modular_deepseek_v2.py index 9717b12b17d4..271c7a7c4815 100644 --- a/src/transformers/models/deepseek_v2/modular_deepseek_v2.py +++ b/src/transformers/models/deepseek_v2/modular_deepseek_v2.py @@ -84,10 +84,9 @@ class DeepseekV2Config(LlamaConfig): "layers.*.mlp.down_proj": "rowwise", } base_model_ep_plan = { - "layers.*.mlp.gate": "ep_router", "layers.*.mlp.experts.gate_up_proj": "grouped_gemm", "layers.*.mlp.experts.down_proj": "grouped_gemm", - "layers.*.mlp.experts": "moe_tp_experts", + "layers.*.mlp.experts": "ep_dispatch_experts", } model_type = "deepseek_v2" diff --git a/src/transformers/models/deepseek_v3/configuration_deepseek_v3.py b/src/transformers/models/deepseek_v3/configuration_deepseek_v3.py index 06af5dbec4db..61c7a5220e94 100644 --- a/src/transformers/models/deepseek_v3/configuration_deepseek_v3.py +++ b/src/transformers/models/deepseek_v3/configuration_deepseek_v3.py @@ -69,10 +69,9 @@ class DeepseekV3Config(PreTrainedConfig): "norm": (["hidden_states"], ["hidden_states"]), } base_model_ep_plan = { - "layers.*.mlp.gate": "ep_router", "layers.*.mlp.experts.gate_up_proj": "grouped_gemm", "layers.*.mlp.experts.down_proj": "grouped_gemm", - "layers.*.mlp.experts": "moe_tp_experts", + "layers.*.mlp.experts": "ep_dispatch_experts", } attribute_map = { diff --git a/src/transformers/models/deepseek_v32/configuration_deepseek_v32.py b/src/transformers/models/deepseek_v32/configuration_deepseek_v32.py index 73bb4af5a27b..752ac1935a32 100644 --- a/src/transformers/models/deepseek_v32/configuration_deepseek_v32.py +++ b/src/transformers/models/deepseek_v32/configuration_deepseek_v32.py @@ -77,10 +77,9 @@ class DeepseekV32Config(PreTrainedConfig): "norm": (["hidden_states"], ["hidden_states"]), } base_model_ep_plan = { - "layers.*.mlp.gate": "ep_router", "layers.*.mlp.experts.gate_up_proj": "grouped_gemm", "layers.*.mlp.experts.down_proj": "grouped_gemm", - "layers.*.mlp.experts": "moe_tp_experts", + "layers.*.mlp.experts": "ep_dispatch_experts", } attribute_map = {"num_local_experts": "n_routed_experts"} diff --git a/src/transformers/models/deepseek_v4/configuration_deepseek_v4.py b/src/transformers/models/deepseek_v4/configuration_deepseek_v4.py index 78e7f11291b1..b4f617c39e6a 100644 --- a/src/transformers/models/deepseek_v4/configuration_deepseek_v4.py +++ b/src/transformers/models/deepseek_v4/configuration_deepseek_v4.py @@ -114,10 +114,11 @@ class DeepseekV4Config(PreTrainedConfig): base_model_ep_plan = { # V4 ships EP only (no `base_model_tp_plan` — the runtime picks one plan or # the other, never both, and V4 is MoE so EP is the only sensible config). - # MoE parallelism: route on the gate, run the routed experts as a grouped-GEMM - # kernel sharded along the expert axis, and wrap the experts module with - # `moe_tp_experts` so its output gets all-reduced across ranks. Same shape as - # gpt-oss. Main attention stays replicated: V4 is shared-KV MQA + a CSA / HCA + # MoE parallelism: run the routed experts as a grouped-GEMM kernel sharded along + # the expert axis, with all-to-all token dispatch (`ep_dispatch_experts`). Same + # shape as gpt-oss. Token dispatch only keeps the expert rules, so the indexer rules + # below apply to router masking with all-reduce (`ep_router` on the gate and + # `moe_tp_experts` on the experts). Main attention stays replicated: V4 is shared-KV MQA + a CSA / HCA # compressor branch — both broadcast a single KV head across all attention # heads via `repeat_kv`, so colwise-sharding `q_b_proj` would leave KV # replicated and `repeat_kv` would no longer match the rank-local query head @@ -127,10 +128,9 @@ class DeepseekV4Config(PreTrainedConfig): # hidden_states), so head-sharding is well-formed; `q_b_proj` and the # `scorer.weights_proj` go colwise, and the `scorer` output is all-reduced # so every rank sees the same `index_scores` and picks the same top-k. - "layers.*.mlp.gate": "ep_router", "layers.*.mlp.experts.gate_up_proj": "grouped_gemm", "layers.*.mlp.experts.down_proj": "grouped_gemm", - "layers.*.mlp.experts": "moe_tp_experts", + "layers.*.mlp.experts": "ep_dispatch_experts", "layers.*.self_attn.compressor.indexer.q_b_proj": "colwise", "layers.*.self_attn.compressor.indexer.scorer.weights_proj": "colwise", "layers.*.self_attn.compressor.indexer.scorer": "all_reduce", diff --git a/src/transformers/models/dots1/configuration_dots1.py b/src/transformers/models/dots1/configuration_dots1.py index 39981615ca14..64c8a9680d7e 100644 --- a/src/transformers/models/dots1/configuration_dots1.py +++ b/src/transformers/models/dots1/configuration_dots1.py @@ -71,10 +71,9 @@ class Dots1Config(PreTrainedConfig): "norm": (["hidden_states"], ["hidden_states"]), } base_model_ep_plan = { - "layers.*.mlp.gate": "ep_router", "layers.*.mlp.experts.gate_up_proj": "grouped_gemm", "layers.*.mlp.experts.down_proj": "grouped_gemm", - "layers.*.mlp.experts": "moe_tp_experts", + "layers.*.mlp.experts": "ep_dispatch_experts", } attribute_map = { diff --git a/src/transformers/models/dots1/modular_dots1.py b/src/transformers/models/dots1/modular_dots1.py index 51238c7f5357..5dae488f0617 100644 --- a/src/transformers/models/dots1/modular_dots1.py +++ b/src/transformers/models/dots1/modular_dots1.py @@ -85,10 +85,9 @@ class Dots1Config(PreTrainedConfig): "norm": (["hidden_states"], ["hidden_states"]), } base_model_ep_plan = { - "layers.*.mlp.gate": "ep_router", "layers.*.mlp.experts.gate_up_proj": "grouped_gemm", "layers.*.mlp.experts.down_proj": "grouped_gemm", - "layers.*.mlp.experts": "moe_tp_experts", + "layers.*.mlp.experts": "ep_dispatch_experts", } attribute_map = { diff --git a/src/transformers/models/ernie4_5_moe/configuration_ernie4_5_moe.py b/src/transformers/models/ernie4_5_moe/configuration_ernie4_5_moe.py index 47c7a53dff84..142ff7a2b25c 100644 --- a/src/transformers/models/ernie4_5_moe/configuration_ernie4_5_moe.py +++ b/src/transformers/models/ernie4_5_moe/configuration_ernie4_5_moe.py @@ -83,10 +83,9 @@ class Ernie4_5_MoeConfig(PreTrainedConfig): "norm": (["hidden_states"], ["hidden_states"]), } base_model_ep_plan = { - "layers.*.mlp.gate": "ep_router", "layers.*.mlp.experts.gate_up_proj": "grouped_gemm", "layers.*.mlp.experts.down_proj": "grouped_gemm", - "layers.*.mlp.experts": "moe_tp_experts", + "layers.*.mlp.experts": "ep_dispatch_experts", } vocab_size: int = 103424 diff --git a/src/transformers/models/ernie4_5_vl_moe/configuration_ernie4_5_vl_moe.py b/src/transformers/models/ernie4_5_vl_moe/configuration_ernie4_5_vl_moe.py index 762b910ee666..b920a2f42652 100644 --- a/src/transformers/models/ernie4_5_vl_moe/configuration_ernie4_5_vl_moe.py +++ b/src/transformers/models/ernie4_5_vl_moe/configuration_ernie4_5_vl_moe.py @@ -103,10 +103,9 @@ class Ernie4_5_VLMoeTextConfig(PreTrainedConfig): "norm": (["hidden_states"], ["hidden_states"]), } base_model_ep_plan = { - "layers.*.mlp.gate": "ep_router", "layers.*.mlp.experts.gate_up_proj": "grouped_gemm", "layers.*.mlp.experts.down_proj": "grouped_gemm", - "layers.*.mlp.experts": "moe_tp_experts", + "layers.*.mlp.experts": "ep_dispatch_experts", } vocab_size: int = 103424 diff --git a/src/transformers/models/exaone_moe/configuration_exaone_moe.py b/src/transformers/models/exaone_moe/configuration_exaone_moe.py index b3e9698028ed..587536b2ff3c 100644 --- a/src/transformers/models/exaone_moe/configuration_exaone_moe.py +++ b/src/transformers/models/exaone_moe/configuration_exaone_moe.py @@ -104,10 +104,9 @@ class ExaoneMoeConfig(PreTrainedConfig): layer_types: list[str] | None = None base_model_ep_plan = { - "layers.*.mlp.gate": "ep_router", "layers.*.mlp.experts.gate_up_proj": "grouped_gemm", "layers.*.mlp.experts.down_proj": "grouped_gemm", - "layers.*.mlp.experts": "moe_tp_experts", + "layers.*.mlp.experts": "ep_dispatch_experts", } mlp_layer_types: list[str] | None = None first_k_dense_replace: int = 1 diff --git a/src/transformers/models/exaone_moe/modular_exaone_moe.py b/src/transformers/models/exaone_moe/modular_exaone_moe.py index ff65e7f15ac1..70baba0e8154 100644 --- a/src/transformers/models/exaone_moe/modular_exaone_moe.py +++ b/src/transformers/models/exaone_moe/modular_exaone_moe.py @@ -80,10 +80,9 @@ class ExaoneMoeConfig(Exaone4Config): ```""" base_model_ep_plan = { - "layers.*.mlp.gate": "ep_router", "layers.*.mlp.experts.gate_up_proj": "grouped_gemm", "layers.*.mlp.experts.down_proj": "grouped_gemm", - "layers.*.mlp.experts": "moe_tp_experts", + "layers.*.mlp.experts": "ep_dispatch_experts", } vocab_size: int = 102400 diff --git a/src/transformers/models/flex_olmo/configuration_flex_olmo.py b/src/transformers/models/flex_olmo/configuration_flex_olmo.py index b4fff038c4a4..921ef2b2597c 100644 --- a/src/transformers/models/flex_olmo/configuration_flex_olmo.py +++ b/src/transformers/models/flex_olmo/configuration_flex_olmo.py @@ -64,10 +64,9 @@ class FlexOlmoConfig(PreTrainedConfig): "norm": (["hidden_states"], ["hidden_states"]), } base_model_ep_plan = { - "layers.*.mlp.gate": "ep_router", "layers.*.mlp.experts.gate_up_proj": "grouped_gemm", "layers.*.mlp.experts.down_proj": "grouped_gemm", - "layers.*.mlp.experts": "moe_tp_experts", + "layers.*.mlp.experts": "ep_dispatch_experts", } vocab_size: int = 100352 diff --git a/src/transformers/models/flex_olmo/modular_flex_olmo.py b/src/transformers/models/flex_olmo/modular_flex_olmo.py index 97a63d1a1d5f..956bea08670c 100644 --- a/src/transformers/models/flex_olmo/modular_flex_olmo.py +++ b/src/transformers/models/flex_olmo/modular_flex_olmo.py @@ -74,10 +74,9 @@ class FlexOlmoConfig(PreTrainedConfig): "norm": (["hidden_states"], ["hidden_states"]), } base_model_ep_plan = { - "layers.*.mlp.gate": "ep_router", "layers.*.mlp.experts.gate_up_proj": "grouped_gemm", "layers.*.mlp.experts.down_proj": "grouped_gemm", - "layers.*.mlp.experts": "moe_tp_experts", + "layers.*.mlp.experts": "ep_dispatch_experts", } vocab_size: int = 100352 diff --git a/src/transformers/models/glm4_moe/configuration_glm4_moe.py b/src/transformers/models/glm4_moe/configuration_glm4_moe.py index 81968cd73ea2..c1220df0d3b4 100644 --- a/src/transformers/models/glm4_moe/configuration_glm4_moe.py +++ b/src/transformers/models/glm4_moe/configuration_glm4_moe.py @@ -78,10 +78,9 @@ class Glm4MoeConfig(PreTrainedConfig): "norm": (["hidden_states"], ["hidden_states"]), } base_model_ep_plan = { - "layers.*.mlp.gate": "ep_router", "layers.*.mlp.experts.gate_up_proj": "grouped_gemm", "layers.*.mlp.experts.down_proj": "grouped_gemm", - "layers.*.mlp.experts": "moe_tp_experts", + "layers.*.mlp.experts": "ep_dispatch_experts", } attribute_map = { diff --git a/src/transformers/models/glm4_moe/modular_glm4_moe.py b/src/transformers/models/glm4_moe/modular_glm4_moe.py index 3582f33d59b9..a0103d7ef0df 100644 --- a/src/transformers/models/glm4_moe/modular_glm4_moe.py +++ b/src/transformers/models/glm4_moe/modular_glm4_moe.py @@ -91,10 +91,9 @@ class Glm4MoeConfig(PreTrainedConfig): "norm": (["hidden_states"], ["hidden_states"]), } base_model_ep_plan = { - "layers.*.mlp.gate": "ep_router", "layers.*.mlp.experts.gate_up_proj": "grouped_gemm", "layers.*.mlp.experts.down_proj": "grouped_gemm", - "layers.*.mlp.experts": "moe_tp_experts", + "layers.*.mlp.experts": "ep_dispatch_experts", } attribute_map = { diff --git a/src/transformers/models/glm4_moe_lite/configuration_glm4_moe_lite.py b/src/transformers/models/glm4_moe_lite/configuration_glm4_moe_lite.py index bde313cf6487..3c4c5ac485d6 100644 --- a/src/transformers/models/glm4_moe_lite/configuration_glm4_moe_lite.py +++ b/src/transformers/models/glm4_moe_lite/configuration_glm4_moe_lite.py @@ -69,10 +69,9 @@ class Glm4MoeLiteConfig(PreTrainedConfig): "norm": (["hidden_states"], ["hidden_states"]), } base_model_ep_plan = { - "layers.*.mlp.gate": "ep_router", "layers.*.mlp.experts.gate_up_proj": "grouped_gemm", "layers.*.mlp.experts.down_proj": "grouped_gemm", - "layers.*.mlp.experts": "moe_tp_experts", + "layers.*.mlp.experts": "ep_dispatch_experts", } attribute_map = { diff --git a/src/transformers/models/glm4_moe_lite/modular_glm4_moe_lite.py b/src/transformers/models/glm4_moe_lite/modular_glm4_moe_lite.py index 07ce9ec807eb..2dc9dca45b89 100644 --- a/src/transformers/models/glm4_moe_lite/modular_glm4_moe_lite.py +++ b/src/transformers/models/glm4_moe_lite/modular_glm4_moe_lite.py @@ -77,10 +77,9 @@ class Glm4MoeLiteConfig(PreTrainedConfig): "norm": (["hidden_states"], ["hidden_states"]), } base_model_ep_plan = { - "layers.*.mlp.gate": "ep_router", "layers.*.mlp.experts.gate_up_proj": "grouped_gemm", "layers.*.mlp.experts.down_proj": "grouped_gemm", - "layers.*.mlp.experts": "moe_tp_experts", + "layers.*.mlp.experts": "ep_dispatch_experts", } attribute_map = { diff --git a/src/transformers/models/glm4v_moe/configuration_glm4v_moe.py b/src/transformers/models/glm4v_moe/configuration_glm4v_moe.py index 40d31b3d4dae..193cde3e776b 100644 --- a/src/transformers/models/glm4v_moe/configuration_glm4v_moe.py +++ b/src/transformers/models/glm4v_moe/configuration_glm4v_moe.py @@ -67,10 +67,9 @@ class Glm4vMoeTextConfig(PreTrainedConfig): "norm": (["hidden_states"], ["hidden_states"]), } base_model_ep_plan = { - "layers.*.mlp.gate": "ep_router", "layers.*.mlp.experts.gate_up_proj": "grouped_gemm", "layers.*.mlp.experts.down_proj": "grouped_gemm", - "layers.*.mlp.experts": "moe_tp_experts", + "layers.*.mlp.experts": "ep_dispatch_experts", } attribute_map = { "num_local_experts": "n_routed_experts", diff --git a/src/transformers/models/glm5_next/configuration_glm5_next.py b/src/transformers/models/glm5_next/configuration_glm5_next.py index b7970fb00aba..c5b4efd48037 100644 --- a/src/transformers/models/glm5_next/configuration_glm5_next.py +++ b/src/transformers/models/glm5_next/configuration_glm5_next.py @@ -90,10 +90,9 @@ class Glm5NextTextConfig(PreTrainedConfig): "norm": (["hidden_states"], ["hidden_states"]), } base_model_ep_plan = { - "layers.*.mlp.gate": "ep_router", "layers.*.mlp.experts.gate_up_proj": "grouped_gemm", "layers.*.mlp.experts.down_proj": "grouped_gemm", - "layers.*.mlp.experts": "moe_tp_experts", + "layers.*.mlp.experts": "ep_dispatch_experts", } attribute_map = {"num_local_experts": "n_routed_experts"} diff --git a/src/transformers/models/glm_moe_dsa/configuration_glm_moe_dsa.py b/src/transformers/models/glm_moe_dsa/configuration_glm_moe_dsa.py index 975909241602..9ff12c6a58c7 100644 --- a/src/transformers/models/glm_moe_dsa/configuration_glm_moe_dsa.py +++ b/src/transformers/models/glm_moe_dsa/configuration_glm_moe_dsa.py @@ -82,10 +82,9 @@ class GlmMoeDsaConfig(PreTrainedConfig): "norm": (["hidden_states"], ["hidden_states"]), } base_model_ep_plan = { - "layers.*.mlp.gate": "ep_router", "layers.*.mlp.experts.gate_up_proj": "grouped_gemm", "layers.*.mlp.experts.down_proj": "grouped_gemm", - "layers.*.mlp.experts": "moe_tp_experts", + "layers.*.mlp.experts": "ep_dispatch_experts", } attribute_map = {"num_local_experts": "n_routed_experts"} diff --git a/src/transformers/models/gpt_oss/configuration_gpt_oss.py b/src/transformers/models/gpt_oss/configuration_gpt_oss.py index 47c029a5bca9..5aa25ed9ee1f 100644 --- a/src/transformers/models/gpt_oss/configuration_gpt_oss.py +++ b/src/transformers/models/gpt_oss/configuration_gpt_oss.py @@ -33,12 +33,11 @@ class GptOssConfig(PreTrainedConfig): "norm": (["hidden_states"], ["hidden_states"]), } base_model_ep_plan = { - "layers.*.mlp.router": "ep_router", "layers.*.mlp.experts.gate_up_proj": "grouped_gemm", "layers.*.mlp.experts.gate_up_proj_bias": "grouped_gemm", "layers.*.mlp.experts.down_proj": "grouped_gemm", "layers.*.mlp.experts.down_proj_bias": "grouped_gemm", - "layers.*.mlp.experts": "moe_tp_experts", + "layers.*.mlp.experts": "ep_dispatch_experts", } num_hidden_layers: int = 36 diff --git a/src/transformers/models/hunyuan_v1_moe/configuration_hunyuan_v1_moe.py b/src/transformers/models/hunyuan_v1_moe/configuration_hunyuan_v1_moe.py index f72c6155c5ea..2a7f9e51c41a 100644 --- a/src/transformers/models/hunyuan_v1_moe/configuration_hunyuan_v1_moe.py +++ b/src/transformers/models/hunyuan_v1_moe/configuration_hunyuan_v1_moe.py @@ -34,10 +34,9 @@ class HunYuanMoEV1Config(PreTrainedConfig): model_type = "hunyuan_v1_moe" keys_to_ignore_at_inference = ["past_key_values"] base_model_ep_plan = { - "layers.*.mlp.gate": "ep_router", "layers.*.mlp.experts.gate_up_proj": "grouped_gemm", "layers.*.mlp.experts.down_proj": "grouped_gemm", - "layers.*.mlp.experts": "moe_tp_experts", + "layers.*.mlp.experts": "ep_dispatch_experts", } attribute_map = { "num_experts_per_tok": "moe_topk", diff --git a/src/transformers/models/hy_v3/configuration_hy_v3.py b/src/transformers/models/hy_v3/configuration_hy_v3.py index 4399a0f4d546..6266f301095c 100644 --- a/src/transformers/models/hy_v3/configuration_hy_v3.py +++ b/src/transformers/models/hy_v3/configuration_hy_v3.py @@ -73,10 +73,9 @@ class HYV3Config(PreTrainedConfig): "norm": (["hidden_states"], ["hidden_states"]), } base_model_ep_plan = { - "layers.*.mlp.gate": "ep_router", "layers.*.mlp.experts.gate_up_proj": "grouped_gemm", "layers.*.mlp.experts.down_proj": "grouped_gemm", - "layers.*.mlp.experts": "moe_tp_experts", + "layers.*.mlp.experts": "ep_dispatch_experts", } vocab_size: int = 120832 diff --git a/src/transformers/models/hy_v3/modular_hy_v3.py b/src/transformers/models/hy_v3/modular_hy_v3.py index c24758bb5100..7aa409123b40 100644 --- a/src/transformers/models/hy_v3/modular_hy_v3.py +++ b/src/transformers/models/hy_v3/modular_hy_v3.py @@ -98,10 +98,9 @@ class HYV3Config(PreTrainedConfig): "norm": (["hidden_states"], ["hidden_states"]), } base_model_ep_plan = { - "layers.*.mlp.gate": "ep_router", "layers.*.mlp.experts.gate_up_proj": "grouped_gemm", "layers.*.mlp.experts.down_proj": "grouped_gemm", - "layers.*.mlp.experts": "moe_tp_experts", + "layers.*.mlp.experts": "ep_dispatch_experts", } vocab_size: int = 120832 diff --git a/src/transformers/models/kimi_linear/configuration_kimi_linear.py b/src/transformers/models/kimi_linear/configuration_kimi_linear.py index 27f7adb2a3aa..5aee105e3ee0 100644 --- a/src/transformers/models/kimi_linear/configuration_kimi_linear.py +++ b/src/transformers/models/kimi_linear/configuration_kimi_linear.py @@ -60,10 +60,9 @@ class KimiLinearConfig(PreTrainedConfig): "norm": (["hidden_states"], ["hidden_states"]), } base_model_ep_plan = { - "layers.*.mlp.gate": "ep_router", "layers.*.mlp.experts.gate_up_proj": "grouped_gemm", "layers.*.mlp.experts.down_proj": "grouped_gemm", - "layers.*.mlp.experts": "moe_tp_experts", + "layers.*.mlp.experts": "ep_dispatch_experts", } attribute_map = { "max_position_embeddings": "model_max_length", diff --git a/src/transformers/models/laguna/configuration_laguna.py b/src/transformers/models/laguna/configuration_laguna.py index e10b6fbd4d80..024937999e4b 100644 --- a/src/transformers/models/laguna/configuration_laguna.py +++ b/src/transformers/models/laguna/configuration_laguna.py @@ -84,10 +84,9 @@ class LagunaConfig(PreTrainedConfig): "norm": (["hidden_states"], ["hidden_states"]), } base_model_ep_plan = { - "layers.*.mlp.gate": "ep_router", "layers.*.mlp.experts.gate_up_proj": "grouped_gemm", "layers.*.mlp.experts.down_proj": "grouped_gemm", - "layers.*.mlp.experts": "moe_tp_experts", + "layers.*.mlp.experts": "ep_dispatch_experts", } vocab_size: int = 100352 diff --git a/src/transformers/models/lfm2_moe/configuration_lfm2_moe.py b/src/transformers/models/lfm2_moe/configuration_lfm2_moe.py index 7b940d740aaa..e3ade4acd74d 100644 --- a/src/transformers/models/lfm2_moe/configuration_lfm2_moe.py +++ b/src/transformers/models/lfm2_moe/configuration_lfm2_moe.py @@ -49,10 +49,9 @@ class Lfm2MoeConfig(PreTrainedConfig): model_type = "lfm2_moe" keys_to_ignore_at_inference = ["past_key_values"] base_model_ep_plan = { - "layers.*.feed_forward.gate": "ep_router", "layers.*.feed_forward.experts.gate_up_proj": "grouped_gemm", "layers.*.feed_forward.experts.down_proj": "grouped_gemm", - "layers.*.feed_forward.experts": "moe_tp_experts", + "layers.*.feed_forward.experts": "ep_dispatch_experts", } default_theta = 1000000.0 diff --git a/src/transformers/models/minimax/configuration_minimax.py b/src/transformers/models/minimax/configuration_minimax.py index 709139a584e5..a7fd614a1d62 100644 --- a/src/transformers/models/minimax/configuration_minimax.py +++ b/src/transformers/models/minimax/configuration_minimax.py @@ -76,10 +76,9 @@ class MiniMaxConfig(PreTrainedConfig): "norm": (["hidden_states"], ["hidden_states"]), } base_model_ep_plan = { - "layers.*.mlp.gate": "ep_router", "layers.*.mlp.experts.gate_up_proj": "grouped_gemm", "layers.*.mlp.experts.down_proj": "grouped_gemm", - "layers.*.mlp.experts": "moe_tp_experts", + "layers.*.mlp.experts": "ep_dispatch_experts", } attribute_map = {"num_experts": "num_local_experts"} diff --git a/src/transformers/models/minimax/modular_minimax.py b/src/transformers/models/minimax/modular_minimax.py index 58b39e59911e..37c41a63c126 100644 --- a/src/transformers/models/minimax/modular_minimax.py +++ b/src/transformers/models/minimax/modular_minimax.py @@ -103,10 +103,9 @@ class MiniMaxConfig(PreTrainedConfig): "norm": (["hidden_states"], ["hidden_states"]), } base_model_ep_plan = { - "layers.*.mlp.gate": "ep_router", "layers.*.mlp.experts.gate_up_proj": "grouped_gemm", "layers.*.mlp.experts.down_proj": "grouped_gemm", - "layers.*.mlp.experts": "moe_tp_experts", + "layers.*.mlp.experts": "ep_dispatch_experts", } attribute_map = {"num_experts": "num_local_experts"} diff --git a/src/transformers/models/minimax_m2/configuration_minimax_m2.py b/src/transformers/models/minimax_m2/configuration_minimax_m2.py index 6e4fcc421b88..3e52eb9f736e 100644 --- a/src/transformers/models/minimax_m2/configuration_minimax_m2.py +++ b/src/transformers/models/minimax_m2/configuration_minimax_m2.py @@ -62,10 +62,9 @@ class MiniMaxM2Config(PreTrainedConfig): "norm": (["hidden_states"], ["hidden_states"]), } base_model_ep_plan = { - "layers.*.mlp.gate": "ep_router", "layers.*.mlp.experts.gate_up_proj": "grouped_gemm", "layers.*.mlp.experts.down_proj": "grouped_gemm", - "layers.*.mlp.experts": "moe_tp_experts", + "layers.*.mlp.experts": "ep_dispatch_experts", } attribute_map = { diff --git a/src/transformers/models/minimax_m2/modular_minimax_m2.py b/src/transformers/models/minimax_m2/modular_minimax_m2.py index 0722185c1c18..57090e31a927 100644 --- a/src/transformers/models/minimax_m2/modular_minimax_m2.py +++ b/src/transformers/models/minimax_m2/modular_minimax_m2.py @@ -81,10 +81,9 @@ class MiniMaxM2Config(PreTrainedConfig): "norm": (["hidden_states"], ["hidden_states"]), } base_model_ep_plan = { - "layers.*.mlp.gate": "ep_router", "layers.*.mlp.experts.gate_up_proj": "grouped_gemm", "layers.*.mlp.experts.down_proj": "grouped_gemm", - "layers.*.mlp.experts": "moe_tp_experts", + "layers.*.mlp.experts": "ep_dispatch_experts", } attribute_map = { diff --git a/src/transformers/models/minimax_m3_vl/configuration_minimax_m3_vl.py b/src/transformers/models/minimax_m3_vl/configuration_minimax_m3_vl.py index 536d2c87d0fa..956785f69a06 100644 --- a/src/transformers/models/minimax_m3_vl/configuration_minimax_m3_vl.py +++ b/src/transformers/models/minimax_m3_vl/configuration_minimax_m3_vl.py @@ -70,10 +70,9 @@ class MiniMaxM3VLTextConfig(PreTrainedConfig): "norm": (["hidden_states"], ["hidden_states"]), } base_model_ep_plan = { - "layers.*.mlp.gate": "ep_router", "layers.*.mlp.experts.gate_up_proj": "grouped_gemm", "layers.*.mlp.experts.down_proj": "grouped_gemm", - "layers.*.mlp.experts": "moe_tp_experts", + "layers.*.mlp.experts": "ep_dispatch_experts", } attribute_map = { diff --git a/src/transformers/models/mistral4/configuration_mistral4.py b/src/transformers/models/mistral4/configuration_mistral4.py index 04ad81fe6cc4..7b921c4aa7a6 100644 --- a/src/transformers/models/mistral4/configuration_mistral4.py +++ b/src/transformers/models/mistral4/configuration_mistral4.py @@ -63,10 +63,9 @@ class Mistral4Config(PreTrainedConfig): "norm": (["hidden_states"], ["hidden_states"]), } base_model_ep_plan = { - "layers.*.mlp.gate": "ep_router", "layers.*.mlp.experts.gate_up_proj": "grouped_gemm", "layers.*.mlp.experts.down_proj": "grouped_gemm", - "layers.*.mlp.experts": "moe_tp_experts", + "layers.*.mlp.experts": "ep_dispatch_experts", } attribute_map = { diff --git a/src/transformers/models/mixtral/configuration_mixtral.py b/src/transformers/models/mixtral/configuration_mixtral.py index f0705cbf0b9c..c1158852cf55 100644 --- a/src/transformers/models/mixtral/configuration_mixtral.py +++ b/src/transformers/models/mixtral/configuration_mixtral.py @@ -57,10 +57,9 @@ class MixtralConfig(PreTrainedConfig): "norm": (["hidden_states"], ["hidden_states"]), } base_model_ep_plan = { - "layers.*.mlp.gate": "ep_router", "layers.*.mlp.experts.gate_up_proj": "grouped_gemm", "layers.*.mlp.experts.down_proj": "grouped_gemm", - "layers.*.mlp.experts": "moe_tp_experts", + "layers.*.mlp.experts": "ep_dispatch_experts", } attribute_map = {"num_experts": "num_local_experts"} diff --git a/src/transformers/models/olmoe/configuration_olmoe.py b/src/transformers/models/olmoe/configuration_olmoe.py index 1bdb43a14332..699703b0d94d 100644 --- a/src/transformers/models/olmoe/configuration_olmoe.py +++ b/src/transformers/models/olmoe/configuration_olmoe.py @@ -55,10 +55,9 @@ class OlmoeConfig(PreTrainedConfig): "layers.*.mlp.experts": "moe_tp_experts", } base_model_ep_plan = { - "layers.*.mlp.gate": "ep_router", "layers.*.mlp.experts.gate_up_proj": "grouped_gemm", "layers.*.mlp.experts.down_proj": "grouped_gemm", - "layers.*.mlp.experts": "moe_tp_experts", + "layers.*.mlp.experts": "ep_dispatch_experts", } vocab_size: int = 50304 diff --git a/src/transformers/models/openai_privacy_filter/configuration_openai_privacy_filter.py b/src/transformers/models/openai_privacy_filter/configuration_openai_privacy_filter.py index e7aaefde4bca..bbd1b002895d 100644 --- a/src/transformers/models/openai_privacy_filter/configuration_openai_privacy_filter.py +++ b/src/transformers/models/openai_privacy_filter/configuration_openai_privacy_filter.py @@ -57,12 +57,11 @@ class OpenAIPrivacyFilterConfig(PreTrainedConfig): "norm": (["hidden_states"], ["hidden_states"]), } base_model_ep_plan = { - "layers.*.mlp.router": "ep_router", "layers.*.mlp.experts.gate_up_proj": "grouped_gemm", "layers.*.mlp.experts.gate_up_proj_bias": "grouped_gemm", "layers.*.mlp.experts.down_proj": "grouped_gemm", "layers.*.mlp.experts.down_proj_bias": "grouped_gemm", - "layers.*.mlp.experts": "moe_tp_experts", + "layers.*.mlp.experts": "ep_dispatch_experts", } num_hidden_layers: int = 8 num_local_experts: int = 128 diff --git a/src/transformers/models/phimoe/configuration_phimoe.py b/src/transformers/models/phimoe/configuration_phimoe.py index 37673ceca5a2..89afec64ce3d 100644 --- a/src/transformers/models/phimoe/configuration_phimoe.py +++ b/src/transformers/models/phimoe/configuration_phimoe.py @@ -47,10 +47,9 @@ class PhimoeConfig(PreTrainedConfig): model_type = "phimoe" keys_to_ignore_at_inference = ["past_key_values"] base_model_ep_plan = { - "layers.*.mlp.router": "ep_router", "layers.*.mlp.experts.gate_up_proj": "grouped_gemm", "layers.*.mlp.experts.down_proj": "grouped_gemm", - "layers.*.mlp.experts": "moe_tp_experts", + "layers.*.mlp.experts": "ep_dispatch_experts", } default_theta = 1000000.0 diff --git a/src/transformers/models/qwen2_moe/configuration_qwen2_moe.py b/src/transformers/models/qwen2_moe/configuration_qwen2_moe.py index 318eedbe203d..6f034ef7dde6 100644 --- a/src/transformers/models/qwen2_moe/configuration_qwen2_moe.py +++ b/src/transformers/models/qwen2_moe/configuration_qwen2_moe.py @@ -68,10 +68,9 @@ class Qwen2MoeConfig(PreTrainedConfig): "norm": (["hidden_states"], ["hidden_states"]), } base_model_ep_plan = { - "layers.*.mlp.gate": "ep_router", "layers.*.mlp.experts.gate_up_proj": "grouped_gemm", "layers.*.mlp.experts.down_proj": "grouped_gemm", - "layers.*.mlp.experts": "moe_tp_experts", + "layers.*.mlp.experts": "ep_dispatch_experts", } vocab_size: int = 151936 diff --git a/src/transformers/models/qwen3_5_moe/configuration_qwen3_5_moe.py b/src/transformers/models/qwen3_5_moe/configuration_qwen3_5_moe.py index 89283522a3dd..0ae57e89f9f5 100644 --- a/src/transformers/models/qwen3_5_moe/configuration_qwen3_5_moe.py +++ b/src/transformers/models/qwen3_5_moe/configuration_qwen3_5_moe.py @@ -81,10 +81,9 @@ class Qwen3_5MoeTextConfig(PreTrainedConfig): "norm": (["hidden_states"], ["hidden_states"]), } base_model_ep_plan = { - "layers.*.mlp.gate": "ep_router", "layers.*.mlp.experts.gate_up_proj": "grouped_gemm", "layers.*.mlp.experts.down_proj": "grouped_gemm", - "layers.*.mlp.experts": "moe_tp_experts", + "layers.*.mlp.experts": "ep_dispatch_experts", } vocab_size: int = 248320 diff --git a/src/transformers/models/qwen3_next/configuration_qwen3_next.py b/src/transformers/models/qwen3_next/configuration_qwen3_next.py index 076ccc380ee6..bcba1a251bea 100644 --- a/src/transformers/models/qwen3_next/configuration_qwen3_next.py +++ b/src/transformers/models/qwen3_next/configuration_qwen3_next.py @@ -84,10 +84,9 @@ class Qwen3NextConfig(PreTrainedConfig): "norm": (["hidden_states"], ["hidden_states"]), } base_model_ep_plan = { - "layers.*.mlp.gate": "ep_router", "layers.*.mlp.experts.gate_up_proj": "grouped_gemm", "layers.*.mlp.experts.down_proj": "grouped_gemm", - "layers.*.mlp.experts": "moe_tp_experts", + "layers.*.mlp.experts": "ep_dispatch_experts", } vocab_size: int = 151936 diff --git a/src/transformers/models/qwen3_omni_moe/configuration_qwen3_omni_moe.py b/src/transformers/models/qwen3_omni_moe/configuration_qwen3_omni_moe.py index 41d44aedb7d5..6a3c3b87ee7f 100644 --- a/src/transformers/models/qwen3_omni_moe/configuration_qwen3_omni_moe.py +++ b/src/transformers/models/qwen3_omni_moe/configuration_qwen3_omni_moe.py @@ -370,10 +370,9 @@ class Qwen3OmniMoeTalkerTextConfig(PreTrainedConfig): "layers.*.mlp.down_proj": "rowwise", } base_model_ep_plan = { - "layers.*.mlp.gate": "ep_router", "layers.*.mlp.experts.gate_up_proj": "grouped_gemm", "layers.*.mlp.experts.down_proj": "grouped_gemm", - "layers.*.mlp.experts": "moe_tp_experts", + "layers.*.mlp.experts": "ep_dispatch_experts", } base_model_pp_plan = { "embed_tokens": (["input_ids"], ["inputs_embeds"]), diff --git a/src/transformers/models/qwen3_omni_moe/modular_qwen3_omni_moe.py b/src/transformers/models/qwen3_omni_moe/modular_qwen3_omni_moe.py index c8ebaafcc6d4..e10bb164d1ac 100644 --- a/src/transformers/models/qwen3_omni_moe/modular_qwen3_omni_moe.py +++ b/src/transformers/models/qwen3_omni_moe/modular_qwen3_omni_moe.py @@ -403,10 +403,9 @@ def __post_init__(self, **kwargs): @strict class Qwen3OmniMoeTalkerTextConfig(Qwen3MoeConfig): base_model_ep_plan = { - "layers.*.mlp.gate": "ep_router", "layers.*.mlp.experts.gate_up_proj": "grouped_gemm", "layers.*.mlp.experts.down_proj": "grouped_gemm", - "layers.*.mlp.experts": "moe_tp_experts", + "layers.*.mlp.experts": "ep_dispatch_experts", } vocab_size: int = 3072 diff --git a/src/transformers/models/qwen3_vl_moe/configuration_qwen3_vl_moe.py b/src/transformers/models/qwen3_vl_moe/configuration_qwen3_vl_moe.py index 26b4793fb73c..ede38489950f 100644 --- a/src/transformers/models/qwen3_vl_moe/configuration_qwen3_vl_moe.py +++ b/src/transformers/models/qwen3_vl_moe/configuration_qwen3_vl_moe.py @@ -65,10 +65,9 @@ class Qwen3VLMoeTextConfig(PreTrainedConfig): "layers.*.mlp.down_proj": "rowwise", } base_model_ep_plan = { - "layers.*.mlp.gate": "ep_router", "layers.*.mlp.experts.gate_up_proj": "grouped_gemm", "layers.*.mlp.experts.down_proj": "grouped_gemm", - "layers.*.mlp.experts": "moe_tp_experts", + "layers.*.mlp.experts": "ep_dispatch_experts", } base_model_pp_plan = { "embed_tokens": (["input_ids"], ["inputs_embeds"]), diff --git a/src/transformers/models/qwen3_vl_moe/modular_qwen3_vl_moe.py b/src/transformers/models/qwen3_vl_moe/modular_qwen3_vl_moe.py index 286439a60521..0902fd903091 100644 --- a/src/transformers/models/qwen3_vl_moe/modular_qwen3_vl_moe.py +++ b/src/transformers/models/qwen3_vl_moe/modular_qwen3_vl_moe.py @@ -92,10 +92,9 @@ class Qwen3VLMoeTextConfig(Qwen3MoeConfig): "layers.*.mlp.down_proj": "rowwise", } base_model_ep_plan = { - "layers.*.mlp.gate": "ep_router", "layers.*.mlp.experts.gate_up_proj": "grouped_gemm", "layers.*.mlp.experts.down_proj": "grouped_gemm", - "layers.*.mlp.experts": "moe_tp_experts", + "layers.*.mlp.experts": "ep_dispatch_experts", } base_model_pp_plan = { "embed_tokens": (["input_ids"], ["inputs_embeds"]), diff --git a/src/transformers/models/qwen4_exp/configuration_qwen4_exp.py b/src/transformers/models/qwen4_exp/configuration_qwen4_exp.py index 69a37377a2cb..fcaedba75175 100644 --- a/src/transformers/models/qwen4_exp/configuration_qwen4_exp.py +++ b/src/transformers/models/qwen4_exp/configuration_qwen4_exp.py @@ -100,10 +100,9 @@ class Qwen4ExpTextConfig(PreTrainedConfig): } base_model_pp_plan = None base_model_ep_plan = { - "layers.*.mlp.gate": "ep_router", "layers.*.mlp.experts.gate_up_proj": "grouped_gemm", "layers.*.mlp.experts.down_proj": "grouped_gemm", - "layers.*.mlp.experts": "moe_tp_experts", + "layers.*.mlp.experts": "ep_dispatch_experts", } vocab_size: int = 248320 diff --git a/src/transformers/models/solar_open/configuration_solar_open.py b/src/transformers/models/solar_open/configuration_solar_open.py index cd57174b627c..dbc02e34ce12 100644 --- a/src/transformers/models/solar_open/configuration_solar_open.py +++ b/src/transformers/models/solar_open/configuration_solar_open.py @@ -52,10 +52,9 @@ class SolarOpenConfig(PreTrainedConfig): "norm": (["hidden_states"], ["hidden_states"]), } base_model_ep_plan = { - "layers.*.mlp.gate": "ep_router", "layers.*.mlp.experts.gate_up_proj": "grouped_gemm", "layers.*.mlp.experts.down_proj": "grouped_gemm", - "layers.*.mlp.experts": "moe_tp_experts", + "layers.*.mlp.experts": "ep_dispatch_experts", } attribute_map = { "num_local_experts": "n_routed_experts", diff --git a/src/transformers/models/step3p7/configuration_step3p7.py b/src/transformers/models/step3p7/configuration_step3p7.py index aaebb87166bd..bb6e084eac49 100644 --- a/src/transformers/models/step3p7/configuration_step3p7.py +++ b/src/transformers/models/step3p7/configuration_step3p7.py @@ -128,10 +128,9 @@ class Step3p7TextConfig(PreTrainedConfig): "norm": (["hidden_states"], ["hidden_states"]), } base_model_ep_plan = { - "layers.*.mlp.gate": "ep_router", "layers.*.mlp.experts.gate_up_proj": "grouped_gemm", "layers.*.mlp.experts.down_proj": "grouped_gemm", - "layers.*.mlp.experts": "moe_tp_experts", + "layers.*.mlp.experts": "ep_dispatch_experts", } attribute_map = { "num_local_experts": "n_routed_experts", diff --git a/src/transformers/models/youtu/configuration_youtu.py b/src/transformers/models/youtu/configuration_youtu.py index ffe0fa4fc355..ac4feeaade6c 100644 --- a/src/transformers/models/youtu/configuration_youtu.py +++ b/src/transformers/models/youtu/configuration_youtu.py @@ -60,6 +60,7 @@ class YoutuConfig(PreTrainedConfig): "layers": (["hidden_states", "attention_mask"], ["hidden_states"]), "norm": (["hidden_states"], ["hidden_states"]), } + # Router masking with all-reduce until TP + FSDP composes for this model (token dispatch always applies FSDP2). base_model_ep_plan = { "layers.*.mlp.gate": "ep_router", "layers.*.mlp.experts.gate_up_proj": "grouped_gemm", diff --git a/src/transformers/models/youtu/modular_youtu.py b/src/transformers/models/youtu/modular_youtu.py index 6ad7d962c5c5..1d5040e06a86 100644 --- a/src/transformers/models/youtu/modular_youtu.py +++ b/src/transformers/models/youtu/modular_youtu.py @@ -65,6 +65,13 @@ class YoutuConfig(DeepseekV3Config): "layers.*.mlp.down_proj": "rowwise", } attribute_map = {} + # Router masking with all-reduce until TP + FSDP composes for this model (token dispatch always applies FSDP2). + base_model_ep_plan = { + "layers.*.mlp.gate": "ep_router", + "layers.*.mlp.experts.gate_up_proj": "grouped_gemm", + "layers.*.mlp.experts.down_proj": "grouped_gemm", + "layers.*.mlp.experts": "moe_tp_experts", + } vocab_size: int = 128256 hidden_size: int = 2048 diff --git a/src/transformers/models/zaya/configuration_zaya.py b/src/transformers/models/zaya/configuration_zaya.py index 63fe0c5b7ce4..bca6ee3e5b25 100644 --- a/src/transformers/models/zaya/configuration_zaya.py +++ b/src/transformers/models/zaya/configuration_zaya.py @@ -52,6 +52,7 @@ class ZayaConfig(PreTrainedConfig): model_type = "zaya" keys_to_ignore_at_inference = ["past_key_values"] + # Router masking with all-reduce until TP + FSDP composes for this model (token dispatch always applies FSDP2). base_model_ep_plan = { "layers.*.mlp.gate": "ep_router", "layers.*.mlp.experts.gate_up_proj": "grouped_gemm", diff --git a/src/transformers/models/zaya/modular_zaya.py b/src/transformers/models/zaya/modular_zaya.py index 100b551c27db..87370ca6c778 100644 --- a/src/transformers/models/zaya/modular_zaya.py +++ b/src/transformers/models/zaya/modular_zaya.py @@ -71,6 +71,13 @@ class ZayaConfig(LagunaConfig): """ model_type = "zaya" + # Router masking with all-reduce until TP + FSDP composes for this model (token dispatch always applies FSDP2). + base_model_ep_plan = { + "layers.*.mlp.gate": "ep_router", + "layers.*.mlp.experts.gate_up_proj": "grouped_gemm", + "layers.*.mlp.experts.down_proj": "grouped_gemm", + "layers.*.mlp.experts": "moe_tp_experts", + } vocab_size: int = 262272 moe_intermediate_size: int = 2048 From 4497d747e5ff0af8cef7cb2630cb6bdde3d54743 Mon Sep 17 00:00:00 2001 From: 3outeille Date: Mon, 28 Sep 2026 16:44:10 +0000 Subject: [PATCH 69/86] remove --- .../scripts/expert_parallel_train.py | 110 ------------------ ...est_trainer_distributed_expert_parallel.py | 106 ----------------- 2 files changed, 216 deletions(-) delete mode 100644 tests/trainer/distributed/scripts/expert_parallel_train.py delete mode 100644 tests/trainer/distributed/test_trainer_distributed_expert_parallel.py diff --git a/tests/trainer/distributed/scripts/expert_parallel_train.py b/tests/trainer/distributed/scripts/expert_parallel_train.py deleted file mode 100644 index 6e02544103e6..000000000000 --- a/tests/trainer/distributed/scripts/expert_parallel_train.py +++ /dev/null @@ -1,110 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -""" -Worker script for the expert-parallel `Trainer` tests: train a tiny MoE for a few steps under one layout. - -Launched via ``torchrun`` from ``test_trainer_distributed_expert_parallel.py``. Every layout consumes the same -global batch per step, so the logged losses and gradient norms and the saved weights must match the single-process -``reference`` run. -""" - -import argparse -import json -import os - -import torch -from torch.utils.data import Dataset - -from transformers import AutoModelForCausalLM, Trainer, TrainingArguments -from transformers.distributed import DistributedConfig - - -GLOBAL_BATCH_SIZE = 8 -SEQ_LEN = 16 -NUM_STEPS = 4 - -# `world_size=4` layouts. `dp` is the number of distinct batches per step; each rank trains on -# `GLOBAL_BATCH_SIZE // dp` samples. The masked plan is Qwen3 MoE's dispatch plan with the two forward rules -# overridden; dispatch uses the model's default plan. -MASKED_PLAN = {"model.layers.*.mlp.gate": "ep_router", "model.layers.*.mlp.experts": "moe_tp_experts"} -LAYOUTS = { - "reference": {"world_size": 1, "dp": 1, "config": None}, - "masked": { - "world_size": 4, - "dp": 2, - "config": {"tp_size": 2, "fsdp_size": 2, "ep_size": 2, "ep_plan": MASKED_PLAN}, - }, - "dispatch": {"world_size": 4, "dp": 4, "config": {"tp_size": 1, "fsdp_size": 4, "ep_size": 2}}, - "dispatch_tp": {"world_size": 4, "dp": 2, "config": {"tp_size": 2, "fsdp_size": 2, "ep_size": 4}}, -} - - -class TokenDataset(Dataset): - def __init__(self, vocab_size: int): - generator = torch.Generator().manual_seed(0) - self.input_ids = torch.randint(0, vocab_size, (NUM_STEPS * GLOBAL_BATCH_SIZE, SEQ_LEN), generator=generator) - - def __len__(self): - return self.input_ids.size(0) - - def __getitem__(self, index): - return {"input_ids": self.input_ids[index], "labels": self.input_ids[index].clone()} - - -def main(): - parser = argparse.ArgumentParser() - parser.add_argument("--layout", choices=LAYOUTS, required=True) - parser.add_argument("--model_dir", required=True) - parser.add_argument("--output_dir", required=True) - args = parser.parse_args() - layout = LAYOUTS[args.layout] - assert int(os.environ.get("WORLD_SIZE", 1)) == layout["world_size"], args.layout - - distributed_config = DistributedConfig(**layout["config"]) if layout["config"] else None - model = AutoModelForCausalLM.from_pretrained( - args.model_dir, dtype=torch.float32, distributed_config=distributed_config - ) - if distributed_config is None: - model = model.to("cuda") - - training_args = TrainingArguments( - output_dir=os.path.join(args.output_dir, "trainer"), - per_device_train_batch_size=GLOBAL_BATCH_SIZE // layout["dp"], - max_steps=NUM_STEPS, - learning_rate=1e-3, - max_grad_norm=1.0, - logging_steps=1, - save_strategy="no", - report_to=[], - seed=0, - data_seed=0, - dataloader_drop_last=True, - remove_unused_columns=False, - average_tokens_across_devices=True, - disable_tqdm=True, - ) - trainer = Trainer(model=model, args=training_args, train_dataset=TokenDataset(model.config.vocab_size)) - trainer.train() - # Gathers the sharded weights; every rank takes part, the main process writes. - trainer.save_model(os.path.join(args.output_dir, "model")) - - if trainer.is_world_process_zero(): - steps = [log for log in trainer.state.log_history if "loss" in log] - with open(os.path.join(args.output_dir, "results.json"), "w", encoding="utf-8") as f: - json.dump({"loss": [s["loss"] for s in steps], "grad_norm": [s["grad_norm"] for s in steps]}, f) - - -if __name__ == "__main__": - main() diff --git a/tests/trainer/distributed/test_trainer_distributed_expert_parallel.py b/tests/trainer/distributed/test_trainer_distributed_expert_parallel.py deleted file mode 100644 index 2e4bb051a60a..000000000000 --- a/tests/trainer/distributed/test_trainer_distributed_expert_parallel.py +++ /dev/null @@ -1,106 +0,0 @@ -# Copyright 2026 The HuggingFace Team. All rights reserved. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import json -import os - -from parameterized import parameterized - -from transformers import Qwen3MoeConfig, Qwen3MoeForCausalLM, is_torch_available -from transformers.testing_utils import ( - TestCasePlus, - backend_device_count, - execute_subprocess_async, - get_torch_dist_unique_port, - require_torch_multi_accelerator, - slow, - torch_device, -) - - -if is_torch_available(): - import torch - from safetensors.torch import load_file - - -SCRIPT = os.path.join(os.path.dirname(__file__), "scripts", "expert_parallel_train.py") -LAYOUTS = ("masked", "dispatch", "dispatch_tp") - - -@slow -@require_torch_multi_accelerator -class TestTrainerExpertParallel(TestCasePlus): - """ - `Trainer` with a model sharded at load time by `DistributedConfig`: router masking with all-reduce on a - `(fsdp, tp)` mesh, token dispatch with an independent batch per rank, and token dispatch with TP groups sharing a - batch. Every layout consumes the same global batch per step as the single-process reference, so the logged - losses and gradient norms and the saved weights have to match it. - """ - - def _run(self, layout, model_dir, world_size): - output_dir = self.get_auto_remove_tmp_dir() - cmd = [ - "torchrun", - f"--nproc_per_node={world_size}", - "--nnodes=1", - f"--master_port={get_torch_dist_unique_port()}", - SCRIPT, - f"--layout={layout}", - f"--model_dir={model_dir}", - f"--output_dir={output_dir}", - ] - execute_subprocess_async(cmd, env=self.get_env()) - with open(os.path.join(output_dir, "results.json"), encoding="utf-8") as f: - results = json.load(f) - return results, load_file(os.path.join(output_dir, "model", "model.safetensors")) - - def _model_dir(self): - torch.manual_seed(0) - config = Qwen3MoeConfig( - vocab_size=128, - hidden_size=32, - intermediate_size=64, - moe_intermediate_size=32, - num_hidden_layers=2, - num_attention_heads=4, - num_key_value_heads=2, - head_dim=8, - num_experts=4, - num_experts_per_tok=2, - max_position_embeddings=64, - ) - model_dir = self.get_auto_remove_tmp_dir() - Qwen3MoeForCausalLM(config).save_pretrained(model_dir) - return model_dir - - @parameterized.expand([(layout,) for layout in LAYOUTS]) - def test_matches_single_process_reference(self, layout): - if backend_device_count(torch_device) < 4: - self.skipTest("Requires 4 accelerators") - model_dir = self._model_dir() - reference, reference_weights = self._run("reference", model_dir, world_size=1) - results, weights = self._run(layout, model_dir, world_size=4) - - self.assertEqual(len(results["loss"]), len(reference["loss"])) - torch.testing.assert_close( - torch.tensor(results["loss"]), torch.tensor(reference["loss"]), rtol=1e-4, atol=1e-4 - ) - torch.testing.assert_close( - torch.tensor(results["grad_norm"]), torch.tensor(reference["grad_norm"]), rtol=1e-3, atol=1e-4 - ) - self.assertEqual(set(weights), set(reference_weights)) - # Adam turns rounding-level gradient differences into `lr`-sized weight differences where the gradient is - # near zero, so the weights get a tolerance above the learning rate. - for key, tensor in reference_weights.items(): - torch.testing.assert_close(weights[key], tensor, rtol=1e-3, atol=2e-3, msg=lambda m: f"{key}: {m}") From 9b41975a034cc5cffc22067b4b779925badd22d4 Mon Sep 17 00:00:00 2001 From: 3outeille Date: Mon, 28 Sep 2026 19:15:24 +0000 Subject: [PATCH 70/86] fix ruff --- .../deepseek_v4/configuration_deepseek_v4.py | 24 ++++--------------- 1 file changed, 5 insertions(+), 19 deletions(-) diff --git a/src/transformers/models/deepseek_v4/configuration_deepseek_v4.py b/src/transformers/models/deepseek_v4/configuration_deepseek_v4.py index b4f617c39e6a..e361b66ccc44 100644 --- a/src/transformers/models/deepseek_v4/configuration_deepseek_v4.py +++ b/src/transformers/models/deepseek_v4/configuration_deepseek_v4.py @@ -111,29 +111,15 @@ class DeepseekV4Config(PreTrainedConfig): "layers": (["hidden_states", "attention_mask"], ["hidden_states"]), "norm": (["hidden_states"], ["hidden_states"]), } + # Main attention and the shared MLP stay replicated. Shared-KV MQA and the CSA / HCA compressor + # broadcast one KV head via `repeat_kv`, so colwise `q_b_proj` would mismatch the rank-local + # head count; the shared MLP is too small to be worth sharding. base_model_ep_plan = { - # V4 ships EP only (no `base_model_tp_plan` — the runtime picks one plan or - # the other, never both, and V4 is MoE so EP is the only sensible config). - # MoE parallelism: run the routed experts as a grouped-GEMM kernel sharded along - # the expert axis, with all-to-all token dispatch (`ep_dispatch_experts`). Same - # shape as gpt-oss. Token dispatch only keeps the expert rules, so the indexer rules - # below apply to router masking with all-reduce (`ep_router` on the gate and - # `moe_tp_experts` on the experts). Main attention stays replicated: V4 is shared-KV MQA + a CSA / HCA - # compressor branch — both broadcast a single KV head across all attention - # heads via `repeat_kv`, so colwise-sharding `q_b_proj` would leave KV - # replicated and `repeat_kv` would no longer match the rank-local query head - # count. The shared MLP also stays replicated — it's small and not worth - # sharding. The Lightning Indexer is the one carve-out: its keys are - # replicated (own compressor at index_head_dim fed by replicated - # hidden_states), so head-sharding is well-formed; `q_b_proj` and the - # `scorer.weights_proj` go colwise, and the `scorer` output is all-reduced - # so every rank sees the same `index_scores` and picks the same top-k. + # V4 ships EP only (no `base_model_tp_plan`). Routed experts run as a grouped-GEMM kernel + # sharded along the expert axis, with all-to-all token dispatch (same as gpt-oss). "layers.*.mlp.experts.gate_up_proj": "grouped_gemm", "layers.*.mlp.experts.down_proj": "grouped_gemm", "layers.*.mlp.experts": "ep_dispatch_experts", - "layers.*.self_attn.compressor.indexer.q_b_proj": "colwise", - "layers.*.self_attn.compressor.indexer.scorer.weights_proj": "colwise", - "layers.*.self_attn.compressor.indexer.scorer": "all_reduce", } vocab_size: int = 129280 From 69ecbf137dd34b8473e97bc17c13220f98cda4ae Mon Sep 17 00:00:00 2001 From: 3outeille Date: Mon, 28 Sep 2026 18:36:23 +0000 Subject: [PATCH 71/86] use attribute instead --- src/transformers/distributed/mixin.py | 10 ++++++++++ src/transformers/trainer.py | 21 ++++++++++----------- 2 files changed, 20 insertions(+), 11 deletions(-) diff --git a/src/transformers/distributed/mixin.py b/src/transformers/distributed/mixin.py index 36a5b9c246f8..29bef1debfe1 100644 --- a/src/transformers/distributed/mixin.py +++ b/src/transformers/distributed/mixin.py @@ -52,6 +52,7 @@ class DistributedMixin: """Distributed orchestration and save/load hooks for [`PreTrainedModel`].""" _device_mesh = None + _is_distributed_loading_by_transformers = False _mesh_manager: TransformersDeviceMesh | None = None _tp_plan: dict[str, str] | None = None _ep_plan: dict[str, str] | None = None @@ -60,6 +61,14 @@ class DistributedMixin: _pp_plan: dict[str, tuple[str, str]] | None = None _fsdp_plan: dict[str, str] | None = None + @property + def is_distributed_loading_by_transformers(self) -> bool: + """ + Whether `from_pretrained(distributed_config=...)` sharded the model at load time (each rank reads only its + shard through `DtensorShardOperation`), rather than a wrapper such as Accelerate sharding it afterwards. + """ + return self._is_distributed_loading_by_transformers + def init_parallel_plans(self) -> None: """Copy class-level plans onto the instance and merge config/children contributions.""" model_cls = type(self) @@ -187,6 +196,7 @@ def maybe_distribute_model( model.config.distributed_config = distributed_config model._mesh_manager = mesh_manager model._device_mesh = mesh_manager.get_mesh(("pp", "fsdp", "tp")) + model._is_distributed_loading_by_transformers = True model._tp_size = distributed_config.tp_size model._fsdp_size = distributed_config.fsdp_size diff --git a/src/transformers/trainer.py b/src/transformers/trainer.py index 7ca56548d061..c1420244524f 100755 --- a/src/transformers/trainer.py +++ b/src/transformers/trainer.py @@ -465,6 +465,10 @@ def __init__( elif len(devices) == 1: self.is_model_parallel = self.args.device != torch.device(devices[0]) + # Sharded at load time by `from_pretrained(distributed_config=...)`, whatever the parallelism: the model owns + # its placement and gradient reduction, so Accelerate must not wrap or shard it again. + self.is_distributed_loading_by_transformers = getattr(model, "is_distributed_loading_by_transformers", False) + self.is_fsdp_xla_enabled = args.fsdp and args.fsdp_config.get("xla", False) if args.fsdp: if self.is_deepspeed_enabled: @@ -486,7 +490,7 @@ def __init__( or is_sagemaker_mp_enabled() # Sharded at load time (`DistributedConfig`): the model manages its own placement, and # `.to()` on FSDP2-managed (possibly CPU-offloaded) parameters raises in `_apply`. - or getattr(model, "_device_mesh", None) is not None + or self.is_distributed_loading_by_transformers ): self.place_model_on_device = False else: @@ -627,7 +631,7 @@ def __init__( # Resolved lazily at the first gradient clip; see `_has_mixed_mesh_grads`. self._mixed_mesh_grads: bool | None = None if ( - getattr(model, "_device_mesh", None) is not None + self.is_distributed_loading_by_transformers and args.save_strategy != SaveStrategy.NO and not args.save_only_model ): @@ -1723,14 +1727,9 @@ def _prepare_for_training(self, max_steps, train_dataloader, resume_from_checkpo use_accelerator_prepare = model is self.model # prepare using `accelerator` prepare - if ( - getattr(model, "_is_fsdp_managed_module", False) - and not self.is_fsdp_enabled - and not self.is_deepspeed_enabled - ): - # Sharded at load time (`DistributedConfig` with FSDP2 or expert-parallel token dispatch): the model - # already owns placement and gradient reduction. Prepare autocast and compilation only, without - # asking Accelerate to wrap the DTensor parameters in DDP or to shard them again. + if self.is_distributed_loading_by_transformers: + # The model already owns placement and gradient reduction. Prepare autocast and compilation only, + # without asking Accelerate to wrap the DTensor parameters in DDP or to shard them again. model = self.accelerator.prepare_model(model, device_placement=False, evaluation_mode=True) self.optimizer = self.accelerator.prepare(self.optimizer) self._sync_replicated_trainable_parameters(model) @@ -4088,7 +4087,7 @@ def save_model(self, output_dir: str | None = None, _internal_call: bool = False remove_dummy_checkpoint(self.args.should_save, output_dir, [WEIGHTS_NAME, SAFE_WEIGHTS_NAME]) self.model_wrapped.save_checkpoint(output_dir) - elif getattr(self.model, "_device_mesh", None) is not None and not _is_peft_model(self.model): + elif self.is_distributed_loading_by_transformers and not _is_peft_model(self.model): # Sharded at load time (`DistributedConfig`): gathering the weights inside `save_pretrained` # is collective, so every rank saves; only the main process writes, the others leave at the # closing barrier. (PEFT models fall through to the adapter-only save below.) From 0cb580c6eaeb7420d917432a70c9ba6f86e80b32 Mon Sep 17 00:00:00 2001 From: 3outeille Date: Mon, 28 Sep 2026 19:42:18 +0000 Subject: [PATCH 72/86] Fix EP plans for youtu (dense, drop inherited plan) and zaya (token dispatch) --- src/transformers/models/youtu/configuration_youtu.py | 7 ------- src/transformers/models/youtu/modular_youtu.py | 8 +------- src/transformers/models/zaya/configuration_zaya.py | 4 +--- src/transformers/models/zaya/modular_zaya.py | 4 +--- 4 files changed, 3 insertions(+), 20 deletions(-) diff --git a/src/transformers/models/youtu/configuration_youtu.py b/src/transformers/models/youtu/configuration_youtu.py index ac4feeaade6c..6d9f2cef1f96 100644 --- a/src/transformers/models/youtu/configuration_youtu.py +++ b/src/transformers/models/youtu/configuration_youtu.py @@ -60,13 +60,6 @@ class YoutuConfig(PreTrainedConfig): "layers": (["hidden_states", "attention_mask"], ["hidden_states"]), "norm": (["hidden_states"], ["hidden_states"]), } - # Router masking with all-reduce until TP + FSDP composes for this model (token dispatch always applies FSDP2). - base_model_ep_plan = { - "layers.*.mlp.gate": "ep_router", - "layers.*.mlp.experts.gate_up_proj": "grouped_gemm", - "layers.*.mlp.experts.down_proj": "grouped_gemm", - "layers.*.mlp.experts": "moe_tp_experts", - } attribute_map = {} vocab_size: int = 128256 diff --git a/src/transformers/models/youtu/modular_youtu.py b/src/transformers/models/youtu/modular_youtu.py index 1d5040e06a86..e92c7dad18ea 100644 --- a/src/transformers/models/youtu/modular_youtu.py +++ b/src/transformers/models/youtu/modular_youtu.py @@ -65,13 +65,6 @@ class YoutuConfig(DeepseekV3Config): "layers.*.mlp.down_proj": "rowwise", } attribute_map = {} - # Router masking with all-reduce until TP + FSDP composes for this model (token dispatch always applies FSDP2). - base_model_ep_plan = { - "layers.*.mlp.gate": "ep_router", - "layers.*.mlp.experts.gate_up_proj": "grouped_gemm", - "layers.*.mlp.experts.down_proj": "grouped_gemm", - "layers.*.mlp.experts": "moe_tp_experts", - } vocab_size: int = 128256 hidden_size: int = 2048 @@ -99,6 +92,7 @@ class YoutuConfig(DeepseekV3Config): pretraining_tp = AttributeError() moe_intermediate_size = AttributeError() num_mtp_layers = AttributeError() + base_model_ep_plan = AttributeError() def __post_init__(self, **kwargs): if self.initializer_range is None: diff --git a/src/transformers/models/zaya/configuration_zaya.py b/src/transformers/models/zaya/configuration_zaya.py index bca6ee3e5b25..aaa26896564f 100644 --- a/src/transformers/models/zaya/configuration_zaya.py +++ b/src/transformers/models/zaya/configuration_zaya.py @@ -52,12 +52,10 @@ class ZayaConfig(PreTrainedConfig): model_type = "zaya" keys_to_ignore_at_inference = ["past_key_values"] - # Router masking with all-reduce until TP + FSDP composes for this model (token dispatch always applies FSDP2). base_model_ep_plan = { - "layers.*.mlp.gate": "ep_router", "layers.*.mlp.experts.gate_up_proj": "grouped_gemm", "layers.*.mlp.experts.down_proj": "grouped_gemm", - "layers.*.mlp.experts": "moe_tp_experts", + "layers.*.mlp.experts": "ep_dispatch_experts", } vocab_size: int = 262272 diff --git a/src/transformers/models/zaya/modular_zaya.py b/src/transformers/models/zaya/modular_zaya.py index 87370ca6c778..c372b77c640f 100644 --- a/src/transformers/models/zaya/modular_zaya.py +++ b/src/transformers/models/zaya/modular_zaya.py @@ -71,12 +71,10 @@ class ZayaConfig(LagunaConfig): """ model_type = "zaya" - # Router masking with all-reduce until TP + FSDP composes for this model (token dispatch always applies FSDP2). base_model_ep_plan = { - "layers.*.mlp.gate": "ep_router", "layers.*.mlp.experts.gate_up_proj": "grouped_gemm", "layers.*.mlp.experts.down_proj": "grouped_gemm", - "layers.*.mlp.experts": "moe_tp_experts", + "layers.*.mlp.experts": "ep_dispatch_experts", } vocab_size: int = 262272 From bcbb851d37ca5101645728c40b02e0a589047f2f Mon Sep 17 00:00:00 2001 From: 3outeille Date: Mon, 5 Oct 2026 05:11:04 +0000 Subject: [PATCH 73/86] fix --- src/transformers/distributed/utils.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/src/transformers/distributed/utils.py b/src/transformers/distributed/utils.py index 39cd8256bd8f..9ba6b42db0fb 100644 --- a/src/transformers/distributed/utils.py +++ b/src/transformers/distributed/utils.py @@ -148,9 +148,9 @@ class TransformersDeviceMesh: If one were ever added, the identity would become pp * efsdp * ep * etp == pp * fsdp * tp and efsdp would shrink by etp (efsdp = fsdp * tp / (ep * etp)) - Why not reuse fsdp mesh ? When ep_size == tp_size, efsdp == fsdp, both in size and in which - ranks are grouped together, so the fsdp axis of the dense mesh would work for experts too. - As soon as ep_size != tp_size the two group different ranks and you need a separate axis. + When ep_size == tp_size, efsdp and fsdp are the same axis: same size and same rank groups. + In that case experts could reuse the dense mesh's fsdp axis. + When ep_size != tp_size, the two axes group different ranks, so experts need their own efsdp axis. Regarding ep value, We decide to default it to node width (8 on most machines) so all-to-all never leaves the node. - On a single node, ep == fsdp * tp thus efsdp = 1, the axis does nothing. From 877afeea8f3c2ad3fb193872575aa2e9d496ce89 Mon Sep 17 00:00:00 2001 From: 3outeille Date: Mon, 5 Oct 2026 05:15:32 +0000 Subject: [PATCH 74/86] revert --- tests/tensor_parallel/test_tensor_parallel.py | 214 +----------------- 1 file changed, 2 insertions(+), 212 deletions(-) diff --git a/tests/tensor_parallel/test_tensor_parallel.py b/tests/tensor_parallel/test_tensor_parallel.py index a265cbd4087f..c83741361fd5 100644 --- a/tests/tensor_parallel/test_tensor_parallel.py +++ b/tests/tensor_parallel/test_tensor_parallel.py @@ -16,9 +16,8 @@ import torch -from transformers import AutoModelForCausalLM, Qwen3MoeConfig, Qwen3MoeForCausalLM, Qwen3MoeModel +from transformers import AutoModelForCausalLM from transformers.distributed import tensor_parallel -from transformers.distributed.configuration_utils import DistributedConfig from transformers.distributed.sharding_utils import DtensorShardOperation from transformers.distributed.tensor_parallel import ( ALL_PARALLEL_STYLES, @@ -27,216 +26,7 @@ PackedRowwiseParallel, RowwiseParallel, ) -from transformers.testing_utils import TestCasePlus, is_tensor_parallel_test, require_torch - - -# Qwen3 MoE's predefined plans, as resolved on `Qwen3MoeModel` (no `model.` prefix). -DENSE_TP_PLAN = { - "layers.*.self_attn.q_proj": "colwise", - "layers.*.self_attn.k_proj": "colwise", - "layers.*.self_attn.v_proj": "colwise", - "layers.*.self_attn.q_norm": "replicated_with_grad_allreduce", - "layers.*.self_attn.k_norm": "replicated_with_grad_allreduce", - "layers.*.self_attn.o_proj": "rowwise", - "layers.*.mlp.gate_proj": "colwise", - "layers.*.mlp.up_proj": "colwise", - "layers.*.mlp.down_proj": "rowwise", -} -EXPERT_TP_PLAN = { - "layers.*.mlp.experts.gate_up_proj": "packed_colwise", - "layers.*.mlp.experts.down_proj": "rowwise", - "layers.*.mlp.experts": "moe_tp_experts", -} -EP_PLAN = { - "layers.*.mlp.gate": "ep_router", - "layers.*.mlp.experts.gate_up_proj": "grouped_gemm", - "layers.*.mlp.experts.down_proj": "grouped_gemm", - "layers.*.mlp.experts": "moe_tp_experts", -} - - -@require_torch -class TestParallelPlanResolution(TestCasePlus): - def setUp(self): - super().setUp() - self.config = Qwen3MoeConfig( - vocab_size=32, - hidden_size=16, - intermediate_size=32, - moe_intermediate_size=8, - num_hidden_layers=1, - num_attention_heads=4, - num_key_value_heads=4, - head_dim=4, - num_experts=4, - num_experts_per_tok=2, - ) - with torch.device("meta"): - self.model = Qwen3MoeModel(self.config) - - def test_ep_plan_setter(self): - self.model.ep_plan = None - self.assertEqual(self.model.ep_plan, {}) - with self.assertRaisesRegex(ValueError, "Can only set a dictionary"): - self.model.ep_plan = "auto" - with self.assertRaisesRegex(ValueError, "Unsupported parallel styles"): - self.model.ep_plan = {"layers.*.mlp.experts": "invalid_style"} - self.model.ep_plan = EP_PLAN - self.assertEqual(self.model.ep_plan, EP_PLAN) - - def test_disabled_parallelism_has_no_plans(self): - for config in (DistributedConfig(), DistributedConfig(fsdp_size=8), DistributedConfig(pp_size=2)): - with self.subTest(config=config): - self.assertEqual(tensor_parallel.resolve_parallel_plans(self.model, config), ({}, {})) - - def test_tp_only_keeps_experts_in_tp_plan(self): - tp_plan, ep_plan = tensor_parallel.resolve_parallel_plans(self.model, DistributedConfig(tp_size=4)) - self.assertEqual(tp_plan, DENSE_TP_PLAN | EXPERT_TP_PLAN) - self.assertEqual(ep_plan, {}) - - def test_ep_takes_experts_and_router_out_of_tp_plan(self): - for config in ( - DistributedConfig(tp_size=4, ep_size=4), - DistributedConfig(tp_size=2, fsdp_size=2, ep_size=2), - ): - with self.subTest(config=config): - tp_plan, ep_plan = tensor_parallel.resolve_parallel_plans(self.model, config) - self.assertEqual(tp_plan, DENSE_TP_PLAN) - self.assertEqual(ep_plan, EP_PLAN) - - def test_legacy_flag_is_an_alias_for_ep_size(self): - with warnings.catch_warnings(record=True) as caught: - warnings.simplefilter("always") - legacy = DistributedConfig(tp_size=4, enable_expert_parallel=True) - self.assertEqual([w.category for w in caught], [FutureWarning]) - self.assertIn("Use ep_size=4 instead", str(caught[0].message)) - with warnings.catch_warnings(record=True) as caught: - warnings.simplefilter("always") - explicit = DistributedConfig(tp_size=4, ep_size=4) - disabled = DistributedConfig(tp_size=4, ep_size=1, enable_expert_parallel=True) - self.assertEqual(caught, []) - self.assertEqual(legacy, explicit) - self.assertFalse(disabled.enable_expert_parallel) - self.assertEqual( - tensor_parallel.resolve_parallel_plans(self.model, legacy), - tensor_parallel.resolve_parallel_plans(self.model, explicit), - ) - - def test_ep_plan_is_a_dict_and_round_trips(self): - with self.assertRaisesRegex(ValueError, "`ep_plan` must be a dictionary or None"): - DistributedConfig(tp_size=4, ep_size=4, ep_plan="auto") - config = DistributedConfig(tp_size=4, ep_size=4, ep_plan={"layers.*.mlp.gate": "ep_router"}) - self.assertEqual(config.to_dict()["ep_plan"], {"layers.*.mlp.gate": "ep_router"}) - self.assertEqual(DistributedConfig.from_dict(config.to_dict()), config) - - def test_overrides_merge_into_the_predefined_plans(self): - config = DistributedConfig( - tp_size=4, - ep_size=4, - tp_plan={"layers.*.self_attn.q_proj": "colwise_rep"}, - ep_plan={"layers.*.mlp.experts.down_proj": "rowwise"}, - ) - tp_plan, ep_plan = tensor_parallel.resolve_parallel_plans(self.model, config) - self.assertEqual(tp_plan, DENSE_TP_PLAN | {"layers.*.self_attn.q_proj": "colwise_rep"}) - self.assertEqual(ep_plan, EP_PLAN | {"layers.*.mlp.experts.down_proj": "rowwise"}) - # The merged plans are stored on the model, the config defaults are untouched. - self.assertEqual(self.model.tp_plan["layers.*.self_attn.q_proj"], "colwise_rep") - self.assertEqual(self.model.ep_plan["layers.*.mlp.experts.down_proj"], "rowwise") - self.assertEqual(self.model.config.base_model_tp_plan["layers.*.self_attn.q_proj"], "colwise") - self.assertEqual(self.model.config.base_model_ep_plan["layers.*.mlp.experts.down_proj"], "grouped_gemm") - # The overrides are not rewritten with the merged plans. - self.assertEqual(config.tp_plan, {"layers.*.self_attn.q_proj": "colwise_rep"}) - self.assertEqual(config.ep_plan, {"layers.*.mlp.experts.down_proj": "rowwise"}) - # The merged EP plan stays on the model but is not applied while EP is disabled. - tp_plan, ep_plan = tensor_parallel.resolve_parallel_plans(self.model, DistributedConfig(tp_size=4)) - self.assertEqual(tp_plan, DENSE_TP_PLAN | EXPERT_TP_PLAN | {"layers.*.self_attn.q_proj": "colwise_rep"}) - self.assertEqual(ep_plan, {}) - self.assertEqual(self.model.ep_plan["layers.*.mlp.experts.down_proj"], "rowwise") - - def test_ep_rules_take_precedence_over_tp_rules_for_the_same_modules(self): - config = DistributedConfig( - tp_size=4, - ep_size=4, - tp_plan={"layers.*.mlp.experts.gate_up_proj": "packed_rowwise", "layers.*.mlp.gate": "colwise"}, - ) - tp_plan, ep_plan = tensor_parallel.resolve_parallel_plans(self.model, config) - self.assertEqual(tp_plan, DENSE_TP_PLAN) - self.assertEqual(ep_plan, EP_PLAN) - # The custom TP rules are kept on the model and apply as soon as EP is disabled. - tp_plan, ep_plan = tensor_parallel.resolve_parallel_plans(self.model, DistributedConfig(tp_size=4)) - self.assertEqual(tp_plan["layers.*.mlp.experts.gate_up_proj"], "packed_rowwise") - self.assertEqual(tp_plan["layers.*.mlp.gate"], "colwise") - self.assertEqual(ep_plan, {}) - - def test_ep_requires_an_expert_plan(self): - self.model.ep_plan = None - with self.assertRaisesRegex(ValueError, "does not define an expert-parallel plan"): - tensor_parallel.resolve_parallel_plans(self.model, DistributedConfig(tp_size=4, ep_size=4)) - config = DistributedConfig(tp_size=4, ep_size=4, ep_plan=EP_PLAN) - self.assertEqual(tensor_parallel.resolve_parallel_plans(self.model, config), (DENSE_TP_PLAN, EP_PLAN)) - - def test_unmatched_override_keys_raise_without_changing_plans(self): - original_tp_plan, original_ep_plan = self.model.tp_plan.copy(), self.model.ep_plan.copy() - for plan_name in ("tp_plan", "ep_plan"): - for key in ("layers.*.mlp.experst", "layers.*.mlp.experts.missing_weight", "model.layers.*.mlp.experts"): - with self.subTest(plan_name=plan_name, key=key): - config = DistributedConfig(tp_size=4, ep_size=4, **{plan_name: {key: "grouped_gemm"}}) - with self.assertRaisesRegex(ValueError, f"`{plan_name}` keys .* match nothing in") as error: - tensor_parallel.resolve_parallel_plans(self.model, config) - self.assertIn(key, str(error.exception)) - self.assertIn("Qwen3MoeModel", str(error.exception)) - self.assertEqual(self.model.tp_plan, original_tp_plan) - self.assertEqual(self.model.ep_plan, original_ep_plan) - - def test_override_keys_can_match_modules_parameters_or_existing_plan_keys(self): - for plan_name in ("tp_plan", "ep_plan"): - # `gate_proj` is in the predefined TP plan even though this MoE model has no such module. - for key in ("layers.*.mlp", "layers.*.self_attn.q_proj.weight", "layers.*.mlp.gate_proj"): - with self.subTest(plan_name=plan_name, key=key): - original = getattr(self.model, plan_name).copy() - if plan_name == "ep_plan": - self.model.ep_plan = original | {"layers.*.mlp.gate_proj": "colwise"} - config = DistributedConfig(tp_size=4, **{plan_name: {key: "colwise_rep"}}) - tensor_parallel.resolve_parallel_plans(self.model, config) - self.assertEqual(getattr(self.model, plan_name)[key], "colwise_rep") - setattr(self.model, plan_name, original) - - def test_head_model_overrides_need_the_model_prefix(self): - with torch.device("meta"): - model = Qwen3MoeForCausalLM(self.config) - config = DistributedConfig(tp_size=4, ep_size=4, ep_plan={"layers.*.mlp.gate": "ep_router"}) - with self.assertRaisesRegex(ValueError, "match nothing in Qwen3MoeForCausalLM"): - tensor_parallel.resolve_parallel_plans(model, config) - - config = DistributedConfig( - tp_size=4, - ep_size=4, - tp_plan={"model.layers.*.self_attn.q_proj": "colwise_rep"}, - ep_plan={"model.layers.*.mlp.gate": "ep_router"}, - ) - tp_plan, ep_plan = tensor_parallel.resolve_parallel_plans(model, config) - expected_tp_plan = {f"model.{k}": v for k, v in DENSE_TP_PLAN.items()} | {"lm_head": "colwise_gather_output"} - self.assertEqual(tp_plan, expected_tp_plan | config.tp_plan) - self.assertEqual(ep_plan, {f"model.{k}": v for k, v in EP_PLAN.items()}) - - def test_masked_ep_shards_and_installs_hooks_on_the_tp_mesh(self): - tp_mesh = object() - _, ep_plan = tensor_parallel.resolve_parallel_plans(self.model, DistributedConfig(tp_size=4, ep_size=4)) - experts, router = self.model.layers[0].mlp.experts, self.model.layers[0].mlp.gate - with ( - patch.object(ALL_PARALLEL_STYLES["grouped_gemm"], "validate_param") as validate, - patch.object(ALL_PARALLEL_STYLES["grouped_gemm"], "shard_param") as shard, - patch.object(ALL_PARALLEL_STYLES["moe_tp_experts"], "install_forward") as install_experts, - patch.object(ALL_PARALLEL_STYLES["ep_router"], "install_forward") as install_router, - ): - result = tensor_parallel.apply_tensor_parallelism(self.model, tp_mesh, ep_plan) - self.assertIs(result, self.model) - self.assertEqual(shard.call_count, 2) - for name in ("gate_up_proj", "down_proj"): - validate.assert_any_call(experts, name, tp_mesh, parameter_name=f"layers.0.mlp.experts.{name}") - shard.assert_any_call(experts, name, tp_mesh) - install_experts.assert_called_once_with(experts, tp_mesh) - install_router.assert_called_once_with(router, tp_mesh) +from transformers.testing_utils import TestCasePlus, is_tensor_parallel_test @is_tensor_parallel_test From 0d8d4f557173c145e16c203c7c10588ece885567 Mon Sep 17 00:00:00 2001 From: 3outeille Date: Mon, 5 Oct 2026 05:42:52 +0000 Subject: [PATCH 75/86] revert --- tests/tensor_parallel/test_tensor_parallel.py | 220 +----------------- 1 file changed, 2 insertions(+), 218 deletions(-) diff --git a/tests/tensor_parallel/test_tensor_parallel.py b/tests/tensor_parallel/test_tensor_parallel.py index b9e8a3349254..c83741361fd5 100644 --- a/tests/tensor_parallel/test_tensor_parallel.py +++ b/tests/tensor_parallel/test_tensor_parallel.py @@ -16,9 +16,8 @@ import torch -from transformers import AutoModelForCausalLM, Qwen3MoeConfig, Qwen3MoeForCausalLM, Qwen3MoeModel +from transformers import AutoModelForCausalLM from transformers.distributed import tensor_parallel -from transformers.distributed.configuration_utils import DistributedConfig from transformers.distributed.sharding_utils import DtensorShardOperation from transformers.distributed.tensor_parallel import ( ALL_PARALLEL_STYLES, @@ -27,222 +26,7 @@ PackedRowwiseParallel, RowwiseParallel, ) -from transformers.testing_utils import TestCasePlus, is_tensor_parallel_test, require_torch - - -# Qwen3 MoE's predefined plans, as resolved on `Qwen3MoeModel` (no `model.` prefix). -TP_DENSE_PLAN = { - "layers.*.self_attn.q_proj": "colwise", - "layers.*.self_attn.k_proj": "colwise", - "layers.*.self_attn.v_proj": "colwise", - "layers.*.self_attn.q_norm": "replicated_with_grad_allreduce", - "layers.*.self_attn.k_norm": "replicated_with_grad_allreduce", - "layers.*.self_attn.o_proj": "rowwise", - "layers.*.mlp.gate_proj": "colwise", - "layers.*.mlp.up_proj": "colwise", - "layers.*.mlp.down_proj": "rowwise", -} -TP_EXPERT_PLAN = { - "layers.*.mlp.experts.gate_up_proj": "packed_colwise", - "layers.*.mlp.experts.down_proj": "rowwise", - "layers.*.mlp.experts": "moe_tp_experts", -} -EP_PLAN = { - "layers.*.mlp.experts.gate_up_proj": "grouped_gemm", - "layers.*.mlp.experts.down_proj": "grouped_gemm", - "layers.*.mlp.experts": "ep_dispatch_experts", -} -EP_PLAN_MASKED = EP_PLAN | {"layers.*.mlp.gate": "ep_router", "layers.*.mlp.experts": "moe_tp_experts"} - - -@require_torch -class TestParallelPlanResolution(TestCasePlus): - def setUp(self): - super().setUp() - self.config = Qwen3MoeConfig( - vocab_size=32, - hidden_size=16, - intermediate_size=32, - moe_intermediate_size=8, - num_hidden_layers=1, - num_attention_heads=4, - num_key_value_heads=4, - head_dim=4, - num_experts=4, - num_experts_per_tok=2, - ) - with torch.device("meta"): - self.model = Qwen3MoeModel(self.config) - - def _reset_plans(self): - self.model.tp_plan = TP_DENSE_PLAN | TP_EXPERT_PLAN - self.model.ep_plan = EP_PLAN.copy() - - def test_ep_plan_setter(self): - self.model.ep_plan = None - self.assertEqual(self.model.ep_plan, {}) - with self.assertRaisesRegex(ValueError, "Can only set a dictionary"): - self.model.ep_plan = "auto" - with self.assertRaisesRegex(ValueError, "Unsupported parallel styles"): - self.model.ep_plan = {"layers.*.mlp.experts": "invalid_style"} - self.model.ep_plan = EP_PLAN - self.assertEqual(self.model.ep_plan, EP_PLAN) - - def test_disabled_parallelism_has_no_plans(self): - for config in (DistributedConfig(), DistributedConfig(fsdp_size=8), DistributedConfig(pp_size=2)): - with self.subTest(config=config): - self.assertEqual(tensor_parallel.resolve_parallel_plans(self.model, config), ({}, {})) - - def test_tp_only_keeps_experts_in_tp_plan(self): - tp_plan, ep_plan = tensor_parallel.resolve_parallel_plans(self.model, DistributedConfig(tp_size=4)) - self.assertEqual(tp_plan, TP_DENSE_PLAN | TP_EXPERT_PLAN) - self.assertEqual(ep_plan, {}) - - def test_ep_takes_experts_and_router_out_of_tp_plan(self): - for config in ( - DistributedConfig(tp_size=4, ep_size=4), - DistributedConfig(tp_size=2, fsdp_size=2, ep_size=2), - ): - with self.subTest(config=config): - tp_plan, ep_plan = tensor_parallel.resolve_parallel_plans(self.model, config) - self.assertEqual(tp_plan, TP_DENSE_PLAN) - self.assertEqual(ep_plan, EP_PLAN) - - def test_legacy_flag_is_an_alias_for_ep_size(self): - with warnings.catch_warnings(record=True) as caught: - warnings.simplefilter("always") - legacy = DistributedConfig(tp_size=4, enable_expert_parallel=True) - self.assertEqual([w.category for w in caught], [FutureWarning]) - self.assertIn("Use ep_size=4 instead", str(caught[0].message)) - with warnings.catch_warnings(record=True) as caught: - warnings.simplefilter("always") - explicit = DistributedConfig(tp_size=4, ep_size=4) - disabled = DistributedConfig(tp_size=4, ep_size=1, enable_expert_parallel=True) - self.assertEqual(caught, []) - self.assertEqual(legacy, explicit) - self.assertFalse(disabled.enable_expert_parallel) - self.assertEqual( - tensor_parallel.resolve_parallel_plans(self.model, legacy), - tensor_parallel.resolve_parallel_plans(self.model, explicit), - ) - - def test_ep_plan_is_a_dict_and_round_trips(self): - with self.assertRaisesRegex(ValueError, "`ep_plan` must be a dictionary or None"): - DistributedConfig(tp_size=4, ep_size=4, ep_plan="auto") - config = DistributedConfig(tp_size=4, ep_size=4, ep_plan=EP_PLAN_MASKED) - self.assertEqual(config.to_dict()["ep_plan"], EP_PLAN_MASKED) - self.assertEqual(DistributedConfig.from_dict(config.to_dict()), config) - - def test_overrides_merge_into_the_predefined_plans(self): - config = DistributedConfig( - tp_size=4, - ep_size=4, - tp_plan={"layers.*.self_attn.q_proj": "colwise_rep"}, - ep_plan={"layers.*.mlp.experts.down_proj": "rowwise"}, - ) - tp_plan, ep_plan = tensor_parallel.resolve_parallel_plans(self.model, config) - self.assertEqual(tp_plan, TP_DENSE_PLAN | {"layers.*.self_attn.q_proj": "colwise_rep"}) - self.assertEqual(ep_plan, EP_PLAN | {"layers.*.mlp.experts.down_proj": "rowwise"}) - # The merged plans are stored on the model, the config defaults are untouched. - self.assertEqual(self.model.tp_plan["layers.*.self_attn.q_proj"], "colwise_rep") - self.assertEqual(self.model.ep_plan["layers.*.mlp.experts.down_proj"], "rowwise") - self.assertEqual(self.model.config.base_model_tp_plan["layers.*.self_attn.q_proj"], "colwise") - self.assertEqual(self.model.config.base_model_ep_plan["layers.*.mlp.experts.down_proj"], "grouped_gemm") - # The overrides are not rewritten with the merged plans. - self.assertEqual(config.tp_plan, {"layers.*.self_attn.q_proj": "colwise_rep"}) - self.assertEqual(config.ep_plan, {"layers.*.mlp.experts.down_proj": "rowwise"}) - # The merged EP plan stays on the model but is not applied while EP is disabled. - tp_plan, ep_plan = tensor_parallel.resolve_parallel_plans(self.model, DistributedConfig(tp_size=4)) - self.assertEqual(tp_plan, TP_DENSE_PLAN | TP_EXPERT_PLAN | {"layers.*.self_attn.q_proj": "colwise_rep"}) - self.assertEqual(ep_plan, {}) - self.assertEqual(self.model.ep_plan["layers.*.mlp.experts.down_proj"], "rowwise") - - def test_ep_rules_take_precedence_over_tp_rules_for_the_same_modules(self): - config = DistributedConfig( - tp_size=4, - ep_size=4, - tp_plan={"layers.*.mlp.experts.gate_up_proj": "packed_rowwise", "layers.*.mlp.gate": "colwise"}, - ep_plan=EP_PLAN_MASKED, - ) - tp_plan, ep_plan = tensor_parallel.resolve_parallel_plans(self.model, config) - self.assertEqual(tp_plan, TP_DENSE_PLAN) - self.assertEqual(ep_plan, EP_PLAN_MASKED) - # The custom TP rules are kept on the model and apply as soon as EP is disabled. - tp_plan, ep_plan = tensor_parallel.resolve_parallel_plans(self.model, DistributedConfig(tp_size=4)) - self.assertEqual(tp_plan["layers.*.mlp.experts.gate_up_proj"], "packed_rowwise") - self.assertEqual(tp_plan["layers.*.mlp.gate"], "colwise") - self.assertEqual(ep_plan, {}) - - def test_ep_requires_an_expert_plan(self): - self.model.ep_plan = None - with self.assertRaisesRegex(ValueError, "does not define an expert-parallel plan"): - tensor_parallel.resolve_parallel_plans(self.model, DistributedConfig(tp_size=4, ep_size=4)) - config = DistributedConfig(tp_size=4, ep_size=4, ep_plan=EP_PLAN) - self.assertEqual(tensor_parallel.resolve_parallel_plans(self.model, config), (TP_DENSE_PLAN, EP_PLAN)) - - def test_unmatched_override_keys_raise_without_changing_plans(self): - original_tp_plan, original_ep_plan = self.model.tp_plan.copy(), self.model.ep_plan.copy() - for plan_name in ("tp_plan", "ep_plan"): - for key in ("layers.*.mlp.experst", "layers.*.mlp.experts.missing_weight", "model.layers.*.mlp.experts"): - with self.subTest(plan_name=plan_name, key=key): - config = DistributedConfig(tp_size=4, ep_size=4, **{plan_name: {key: "grouped_gemm"}}) - with self.assertRaisesRegex(ValueError, f"`{plan_name}` keys .* match nothing in") as error: - tensor_parallel.resolve_parallel_plans(self.model, config) - self.assertIn(key, str(error.exception)) - self.assertIn("Qwen3MoeModel", str(error.exception)) - self.assertEqual(self.model.tp_plan, original_tp_plan) - self.assertEqual(self.model.ep_plan, original_ep_plan) - - def test_override_keys_can_match_modules_parameters_or_existing_plan_keys(self): - for plan_name in ("tp_plan", "ep_plan"): - # `gate_proj` is in the predefined TP plan even though this MoE model has no such module. - for key in ("layers.*.mlp", "layers.*.self_attn.q_proj.weight", "layers.*.mlp.gate_proj"): - with self.subTest(plan_name=plan_name, key=key): - original = getattr(self.model, plan_name).copy() - if plan_name == "ep_plan": - self.model.ep_plan = original | {"layers.*.mlp.gate_proj": "colwise"} - config = DistributedConfig(tp_size=4, **{plan_name: {key: "colwise_rep"}}) - tensor_parallel.resolve_parallel_plans(self.model, config) - self.assertEqual(getattr(self.model, plan_name)[key], "colwise_rep") - setattr(self.model, plan_name, original) - - def test_head_model_overrides_need_the_model_prefix(self): - with torch.device("meta"): - model = Qwen3MoeForCausalLM(self.config) - config = DistributedConfig(tp_size=4, ep_size=4, ep_plan=EP_PLAN_MASKED) - with self.assertRaisesRegex(ValueError, "match nothing in Qwen3MoeForCausalLM"): - tensor_parallel.resolve_parallel_plans(model, config) - - config = DistributedConfig( - tp_size=4, - ep_size=4, - tp_plan={"model.layers.*.self_attn.q_proj": "colwise_rep"}, - ep_plan={f"model.{k}": v for k, v in EP_PLAN_MASKED.items()}, - ) - tp_plan, ep_plan = tensor_parallel.resolve_parallel_plans(model, config) - expected_tp_plan = {f"model.{k}": v for k, v in TP_DENSE_PLAN.items()} | {"lm_head": "colwise_gather_output"} - self.assertEqual(tp_plan, expected_tp_plan | config.tp_plan) - self.assertEqual(ep_plan, {f"model.{k}": v for k, v in EP_PLAN_MASKED.items()}) - - def test_masked_ep_shards_and_installs_hooks_on_the_tp_mesh(self): - tp_mesh = object() - config = DistributedConfig(tp_size=4, ep_size=4, ep_plan=EP_PLAN_MASKED) - _, ep_plan = tensor_parallel.resolve_parallel_plans(self.model, config) - experts, router = self.model.layers[0].mlp.experts, self.model.layers[0].mlp.gate - with ( - patch.object(ALL_PARALLEL_STYLES["grouped_gemm"], "validate_param") as validate, - patch.object(ALL_PARALLEL_STYLES["grouped_gemm"], "shard_param") as shard, - patch.object(ALL_PARALLEL_STYLES["moe_tp_experts"], "install_forward") as install_experts, - patch.object(ALL_PARALLEL_STYLES["ep_router"], "install_forward") as install_router, - ): - result = tensor_parallel.apply_tensor_parallelism(self.model, tp_mesh, ep_plan) - self.assertIs(result, self.model) - self.assertEqual(shard.call_count, 2) - for name in ("gate_up_proj", "down_proj"): - validate.assert_any_call(experts, name, tp_mesh, parameter_name=f"layers.0.mlp.experts.{name}") - shard.assert_any_call(experts, name, tp_mesh) - install_experts.assert_called_once_with(experts, tp_mesh) - install_router.assert_called_once_with(router, tp_mesh) +from transformers.testing_utils import TestCasePlus, is_tensor_parallel_test @is_tensor_parallel_test From cedde47c9baeb327b13c0b1b3fa5993b59a980db Mon Sep 17 00:00:00 2001 From: Ferdinand Mom Date: Mon, 5 Oct 2026 03:29:32 -0700 Subject: [PATCH 76/86] typo --- src/transformers/distributed/utils.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/src/transformers/distributed/utils.py b/src/transformers/distributed/utils.py index 9ba6b42db0fb..a692b8956cc7 100644 --- a/src/transformers/distributed/utils.py +++ b/src/transformers/distributed/utils.py @@ -141,14 +141,14 @@ class TransformersDeviceMesh: expert : (pp, efsdp, ep) experts Both views cover the same world, so pp * fsdp * tp == pp * efsdp * ep. - efsdp is not something you pick, it is whatever is left once ep is fixed: + efsdp is not something you pick, it is whatever is left once ep is fixed. The relationship is as follow: efsdp = fsdp * tp / ep. It is the FSDP axis for expert weights same role `fsdp` plays for the dense params. There is no etp (expert tensor parallel) axis yet meaning experts are never tensor-sharded here. - If one were ever added, the identity would become pp * efsdp * ep * etp == pp * fsdp * tp and efsdp would shrink by etp + If one were ever added, the relationship would become pp * efsdp * ep * etp == pp * fsdp * tp and efsdp would shrink by etp (efsdp = fsdp * tp / (ep * etp)) - When ep_size == tp_size, efsdp and fsdp are the same axis: same size and same rank groups. + When ep_size == tp_size, efsdp and fsdp are the same axis, same size and same rank groups. In that case experts could reuse the dense mesh's fsdp axis. When ep_size != tp_size, the two axes group different ranks, so experts need their own efsdp axis. From 6a23a306fcce35715f7584070534f34f8f915ab0 Mon Sep 17 00:00:00 2001 From: Ferdinand Mom Date: Mon, 5 Oct 2026 03:33:06 -0700 Subject: [PATCH 77/86] clearer --- src/transformers/distributed/utils.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/src/transformers/distributed/utils.py b/src/transformers/distributed/utils.py index a692b8956cc7..ce6eaf1a09e0 100644 --- a/src/transformers/distributed/utils.py +++ b/src/transformers/distributed/utils.py @@ -148,9 +148,9 @@ class TransformersDeviceMesh: If one were ever added, the relationship would become pp * efsdp * ep * etp == pp * fsdp * tp and efsdp would shrink by etp (efsdp = fsdp * tp / (ep * etp)) - When ep_size == tp_size, efsdp and fsdp are the same axis, same size and same rank groups. + When ep_size == tp_size, efsdp == fsdp (given the relationship efsdp = fsdp * tp / ep), thus same axis, same size and same groups of ranks. In that case experts could reuse the dense mesh's fsdp axis. - When ep_size != tp_size, the two axes group different ranks, so experts need their own efsdp axis. + When ep_size != tp_size, egsdp != fsdp, so we can't reuse the same mesh as they don't have the same groups of ranks Regarding ep value, We decide to default it to node width (8 on most machines) so all-to-all never leaves the node. - On a single node, ep == fsdp * tp thus efsdp = 1, the axis does nothing. From b4d81b5b58ddb2c7506dfa3e57d632d746fcec09 Mon Sep 17 00:00:00 2001 From: 3outeille Date: Tue, 6 Oct 2026 00:07:35 +0000 Subject: [PATCH 78/86] clean --- src/transformers/distributed/utils.py | 12 +++++++----- 1 file changed, 7 insertions(+), 5 deletions(-) diff --git a/src/transformers/distributed/utils.py b/src/transformers/distributed/utils.py index 9ba6b42db0fb..613b3c6fb722 100644 --- a/src/transformers/distributed/utils.py +++ b/src/transformers/distributed/utils.py @@ -141,18 +141,20 @@ class TransformersDeviceMesh: expert : (pp, efsdp, ep) experts Both views cover the same world, so pp * fsdp * tp == pp * efsdp * ep. - efsdp is not something you pick, it is whatever is left once ep is fixed: + efsdp is not something you pick, it is whatever is left once ep is fixed. The relationship is as follow: efsdp = fsdp * tp / ep. It is the FSDP axis for expert weights same role `fsdp` plays for the dense params. There is no etp (expert tensor parallel) axis yet meaning experts are never tensor-sharded here. - If one were ever added, the identity would become pp * efsdp * ep * etp == pp * fsdp * tp and efsdp would shrink by etp + If one were ever added, the relationship would become pp * efsdp * ep * etp == pp * fsdp * tp and efsdp would shrink by etp (efsdp = fsdp * tp / (ep * etp)) - When ep_size == tp_size, efsdp and fsdp are the same axis: same size and same rank groups. + When ep_size == tp_size, efsdp == fsdp (given the relationship efsdp = fsdp * tp / ep), thus same axis, same size and same groups of ranks. In that case experts could reuse the dense mesh's fsdp axis. - When ep_size != tp_size, the two axes group different ranks, so experts need their own efsdp axis. + When ep_size != tp_size, efsdp != fsdp, so we can't reuse the same mesh as they don't have the same groups of ranks + This explains why experts need their own efsdp axis. - Regarding ep value, We decide to default it to node width (8 on most machines) so all-to-all never leaves the node. + Regarding ep value, one can decide to default the value to node width (8 on most machines) so that all-to-all never leaves the node. + That has several implications on efsdp value given your setup: - On a single node, ep == fsdp * tp thus efsdp = 1, the axis does nothing. - On several nodes, we still keep ep at node width, since all-to-all across nodes is expensive. However, each node then holds a full copy of the expert group and efsdp is the number of copies, which is where FSDP happens for the experts From 7b43249c618a5cd20088fc9a9204e43be21b8786 Mon Sep 17 00:00:00 2001 From: 3outeille Date: Tue, 6 Oct 2026 00:25:55 +0000 Subject: [PATCH 79/86] cleaning --- src/transformers/configuration_utils.py | 6 ++---- src/transformers/distributed/fsdp.py | 1 + 2 files changed, 3 insertions(+), 4 deletions(-) diff --git a/src/transformers/configuration_utils.py b/src/transformers/configuration_utils.py index ef53a3117477..69c89e0422b3 100755 --- a/src/transformers/configuration_utils.py +++ b/src/transformers/configuration_utils.py @@ -388,11 +388,9 @@ def __post_init__(self, **kwargs): self.per_layer_config = per_layer_config # TODO: to support models whose input embedding module is not named `embed_tokens` (e.g. GPT-NeoX's `embed_in`). - if getattr(self, "tie_word_embeddings", False) and ( - self.base_model_tp_plan is not None or self.base_model_ep_plan is not None - ): + if getattr(self, "tie_word_embeddings", False) and self.base_model_tp_plan is not None: self.base_model_tp_plan = { - **(self.base_model_tp_plan or {}), + **self.base_model_tp_plan, "embed_tokens": "embedding_rowwise", } diff --git a/src/transformers/distributed/fsdp.py b/src/transformers/distributed/fsdp.py index dbc756111dcb..3e01f9a07179 100644 --- a/src/transformers/distributed/fsdp.py +++ b/src/transformers/distributed/fsdp.py @@ -218,6 +218,7 @@ def apply_fully_sharded_data_parallelism(model: nn.Module, mesh_manager: MeshMan # - an EP group holds ep_size / tp_size distinct batches # - the efsdp reduce then sums efsdp_size copies of that expert. # The expert gradient therefore covers ep_size / tp_size * efsdp_size = fsdp_size batches + # (the relationsip is efsdp_size = fsdp_size * tp_size / ep_size) module.set_gradient_divide_factor(float(distributed_config.fsdp_size)) if torch.distributed.get_backend(expert_mesh.get_group()) != "nccl": # Non-NCCL backends need to sum first, then apply the division otherwise it runtime error. From e0f701256b4b4909e22fa127fbdfeb9b520ba4b6 Mon Sep 17 00:00:00 2001 From: 3outeille Date: Tue, 6 Oct 2026 00:28:14 +0000 Subject: [PATCH 80/86] cleaning --- tests/test_tensor_parallel_mixin.py | 1 - 1 file changed, 1 deletion(-) diff --git a/tests/test_tensor_parallel_mixin.py b/tests/test_tensor_parallel_mixin.py index c82e1de4671a..f1898aa0afc2 100644 --- a/tests/test_tensor_parallel_mixin.py +++ b/tests/test_tensor_parallel_mixin.py @@ -392,7 +392,6 @@ def _test_tp_generation_quantized_impl(_rank, model_path, model_class, max_new_t def _load_ep_and_reference_models(model_path, model_class): """Load EP model and non-EP reference model for comparison.""" - # All-reduce EP: every rank sees the same tokens, so TP and EP span the same ranks. model_ep = model_class.from_pretrained( model_path, distributed_config=DistributedConfig(tp_size=dist.get_world_size(), ep_size=dist.get_world_size()), From af7939764c5fb4ea051bf20ff8689ea6f932392e Mon Sep 17 00:00:00 2001 From: 3outeille Date: Tue, 6 Oct 2026 00:50:58 +0000 Subject: [PATCH 81/86] dont torch cat for empty tokens --- .../distributed/tensor_parallel.py | 20 ++++++++++--------- 1 file changed, 11 insertions(+), 9 deletions(-) diff --git a/src/transformers/distributed/tensor_parallel.py b/src/transformers/distributed/tensor_parallel.py index 5690addbc1da..00b44df39b63 100644 --- a/src/transformers/distributed/tensor_parallel.py +++ b/src/transformers/distributed/tensor_parallel.py @@ -814,17 +814,19 @@ def _run_local_experts( experts_forward: Callable, tokens: torch.Tensor, expert_ids: torch.Tensor, - num_local_experts: int, ) -> torch.Tensor: """Run local experts with top-1 routing and unit weights; apply routing weights after combine.""" - # One zero row per expert keeps tokens and all expert weights connected to backward. Without it, - # empty eager experts can skip the reverse all-to-all and FSDP reduction, leaving other ranks waiting. - num_tokens, hidden_dim = tokens.shape - local_expert_ids = torch.arange(num_local_experts, device=tokens.device) - tokens = torch.cat([tokens, tokens.new_zeros(num_local_experts, hidden_dim)]) - expert_ids = torch.cat([expert_ids, local_expert_ids]).unsqueeze(-1) + num_tokens = tokens.shape[0] + if num_tokens == 0: + # Keep the empty tokens connected to backward so the reverse all-to-all and FSDP reduction still run. + dummy = tokens.sum(dim=0, keepdim=True) + dummy_ids = torch.zeros((1, 1), dtype=expert_ids.dtype, device=tokens.device) + dummy_weights = torch.ones((1, 1), dtype=tokens.dtype, device=tokens.device) + return experts_forward(dummy, dummy_ids, dummy_weights)[:0] + + expert_ids = expert_ids.unsqueeze(-1) weights = torch.ones_like(expert_ids, dtype=tokens.dtype) - return experts_forward(tokens, expert_ids, weights)[:num_tokens] + return experts_forward(tokens, expert_ids, weights) def _combine_tokens( self, @@ -889,7 +891,7 @@ def ep_forward(hidden_states, top_k_index, top_k_weights): tokens, expert_ids, order, send_sizes, recv_sizes = self._dispatch_tokens( hidden_states, top_k_index, module.num_experts, ep_group, ep_size ) - expert_output = self._run_local_experts(experts_forward, tokens, expert_ids, module.num_experts) + expert_output = self._run_local_experts(experts_forward, tokens, expert_ids) output = self._combine_tokens( expert_output, top_k_weights, order, send_sizes, recv_sizes, ep_group ).to(hidden_states.dtype) From 865aa9637caa0a55618f44c079bb78f2f1ade0fc Mon Sep 17 00:00:00 2001 From: 3outeille Date: Tue, 6 Oct 2026 00:59:39 +0000 Subject: [PATCH 82/86] check code quality --- src/transformers/distributed/fsdp.py | 8 +++----- 1 file changed, 3 insertions(+), 5 deletions(-) diff --git a/src/transformers/distributed/fsdp.py b/src/transformers/distributed/fsdp.py index 3e01f9a07179..fc8f72690cbe 100644 --- a/src/transformers/distributed/fsdp.py +++ b/src/transformers/distributed/fsdp.py @@ -213,12 +213,10 @@ def apply_fully_sharded_data_parallelism(model: nn.Module, mesh_manager: MeshMan for module_name, module in model.named_modules(): if _get_parameter_plan(module_name, model.ep_plan, is_weight=False) == "ep_dispatch_experts": fully_shard(module, mesh=expert_mesh, reshard_after_forward=True, **fsdp_policy_kwargs) - # Dense parameters on fsdp get a per-batch average: gradients summed over fsdp_size batches, - # then divided by fsdp_size. Experts must match that scale: - # - an EP group holds ep_size / tp_size distinct batches - # - the efsdp reduce then sums efsdp_size copies of that expert. + # an EP group holds ep_size / tp_size distinct batches + # the efsdp reduce the sums on efsdp_size copies of that expert. # The expert gradient therefore covers ep_size / tp_size * efsdp_size = fsdp_size batches - # (the relationsip is efsdp_size = fsdp_size * tp_size / ep_size) + # not efsdp_size batches (the relationship is efsdp_size = fsdp_size * tp_size / ep_size) module.set_gradient_divide_factor(float(distributed_config.fsdp_size)) if torch.distributed.get_backend(expert_mesh.get_group()) != "nccl": # Non-NCCL backends need to sum first, then apply the division otherwise it runtime error. From 7b57dd23ebf38fb3e9354e2e3d3d9613ecf14a11 Mon Sep 17 00:00:00 2001 From: Ferdinand Mom <47445085+3outeille@users.noreply.github.com> Date: Tue, 6 Oct 2026 16:46:52 +0900 Subject: [PATCH 83/86] Update src/transformers/distributed/utils.py Co-authored-by: Arthur <48595927+ArthurZucker@users.noreply.github.com> --- src/transformers/distributed/utils.py | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/src/transformers/distributed/utils.py b/src/transformers/distributed/utils.py index 06803cdf0b55..24348b8e8a78 100644 --- a/src/transformers/distributed/utils.py +++ b/src/transformers/distributed/utils.py @@ -154,7 +154,9 @@ class TransformersDeviceMesh: efsdp is not something you pick, it is whatever is left once ep is fixed. The relationship is as follow: efsdp = fsdp * tp / ep. It is the FSDP axis for expert weights same role `fsdp` plays for the dense params. - There is no etp (expert tensor parallel) axis yet meaning experts are never tensor-sharded here. + There is no etp (expert tensor parallel) axis: with EP on, experts are never tensor-sharded. (for now) + For experts, EP plays the role TP plays for dense layers (weights stay sharded during compute), + and efsdp plays the role of fsdp (weights all-gathered before compute). If one were ever added, the relationship would become pp * efsdp * ep * etp == pp * fsdp * tp and efsdp would shrink by etp (efsdp = fsdp * tp / (ep * etp)) From 7af3cbc9b2aed8f283dbbb02e9b14e02300687d1 Mon Sep 17 00:00:00 2001 From: Ferdinand Mom <47445085+3outeille@users.noreply.github.com> Date: Tue, 6 Oct 2026 16:51:56 +0900 Subject: [PATCH 84/86] Update src/transformers/distributed/configuration_utils.py Co-authored-by: Arthur <48595927+ArthurZucker@users.noreply.github.com> --- .../distributed/configuration_utils.py | 17 ++++++++++++++++- 1 file changed, 16 insertions(+), 1 deletion(-) diff --git a/src/transformers/distributed/configuration_utils.py b/src/transformers/distributed/configuration_utils.py index 9012271ec545..56430ca8f0a8 100644 --- a/src/transformers/distributed/configuration_utils.py +++ b/src/transformers/distributed/configuration_utils.py @@ -63,7 +63,22 @@ class DistributedConfig: @property def efsdp_size(self) -> int: - """Size of the expert FSDP axis in the expert mesh view.""" + """ + The number of ranks that own the same experts and FSDP-shard them between each other. + Dense and expert parameters are laid out over the same world size: + + pp x fsdp x tp == pp x efsdp x ep => efsdp = fsdp x tp // ep + + Experts are not tensor-parallel, so EP and expert-FSDP together cover all the ranks that + dense parameters split between FSDP and TP. Example with 16 ranks and ep=8: + + ep groups : {0..7} {8..15} each group holds all experts, tokens routed within it + efsdp groups : {0,8} {1,9} ... {7,15} each pair holds the same experts, FSDP-sharded + + - ep == tp : efsdp == fsdp + - ep == fsdp * tp: efsdp == 1, every expert lives whole on a single rank + TODO: it's not the case yet, but In each group {0,8}, the experts will be sharded on dim(1) (hidden-dim) by fsdp + """ return self.fsdp_size * self.tp_size // self.ep_size def __post_init__(self): From f37bc0f4d99c39909ca2c19d4f2cf11260943c8c Mon Sep 17 00:00:00 2001 From: 3outeille Date: Tue, 6 Oct 2026 07:56:59 +0000 Subject: [PATCH 85/86] linting --- src/transformers/distributed/utils.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/src/transformers/distributed/utils.py b/src/transformers/distributed/utils.py index 24348b8e8a78..7cc8d383b699 100644 --- a/src/transformers/distributed/utils.py +++ b/src/transformers/distributed/utils.py @@ -153,10 +153,10 @@ class TransformersDeviceMesh: Both views cover the same world, so pp * fsdp * tp == pp * efsdp * ep. efsdp is not something you pick, it is whatever is left once ep is fixed. The relationship is as follow: efsdp = fsdp * tp / ep. It is the FSDP axis for expert weights same role `fsdp` plays for the dense params. + There is no etp (expert tensor parallel) axis: with EP on, experts are never tensor-sharded. (for now) + For experts, EP plays the role TP plays for dense layers (weights stay sharded during compute), + and efsdp plays the role of fsdp (weights all-gathered before compute). - There is no etp (expert tensor parallel) axis: with EP on, experts are never tensor-sharded. (for now) - For experts, EP plays the role TP plays for dense layers (weights stay sharded during compute), - and efsdp plays the role of fsdp (weights all-gathered before compute). If one were ever added, the relationship would become pp * efsdp * ep * etp == pp * fsdp * tp and efsdp would shrink by etp (efsdp = fsdp * tp / (ep * etp)) From a4d069396fc2101b2c5398fdadcded561f728b39 Mon Sep 17 00:00:00 2001 From: 3outeille Date: Tue, 6 Oct 2026 08:00:10 +0000 Subject: [PATCH 86/86] fix comment --- src/transformers/distributed/configuration_utils.py | 13 ++++++++----- 1 file changed, 8 insertions(+), 5 deletions(-) diff --git a/src/transformers/distributed/configuration_utils.py b/src/transformers/distributed/configuration_utils.py index 56430ca8f0a8..729a4964c8f2 100644 --- a/src/transformers/distributed/configuration_utils.py +++ b/src/transformers/distributed/configuration_utils.py @@ -64,20 +64,23 @@ class DistributedConfig: @property def efsdp_size(self) -> int: """ - The number of ranks that own the same experts and FSDP-shard them between each other. + The number of ranks that own the same experts and FSDP-shard them between each other. Dense and expert parameters are laid out over the same world size: pp x fsdp x tp == pp x efsdp x ep => efsdp = fsdp x tp // ep - Experts are not tensor-parallel, so EP and expert-FSDP together cover all the ranks that - dense parameters split between FSDP and TP. Example with 16 ranks and ep=8: + Experts are not tensor-parallel, so EP and expert-FSDP together cover all the ranks that + dense parameters split between FSDP and TP. Example with 16 ranks, ep=8: - ep groups : {0..7} {8..15} each group holds all experts, tokens routed within it + ep groups : {0..7} {8..15} the 8 ranks of a group together hold all experts + (num_experts / 8 each), tokens are routed within it efsdp groups : {0,8} {1,9} ... {7,15} each pair holds the same experts, FSDP-sharded + on dim 0 (the expert dim) and all-gathered for compute - ep == tp : efsdp == fsdp - ep == fsdp * tp: efsdp == 1, every expert lives whole on a single rank - TODO: it's not the case yet, but In each group {0,8}, the experts will be sharded on dim(1) (hidden-dim) by fsdp + + Sharding over `efsdp` is applied by the EP token-dispatch path (`ep_dispatch_experts`). """ return self.fsdp_size * self.tp_size // self.ep_size