From 3acf1aebdbfebbc211c6ce3b710490006bf47973 Mon Sep 17 00:00:00 2001 From: "Ethan (Yusheng) Su" Date: Fri, 2 Oct 2026 09:11:01 +0900 Subject: [PATCH 01/11] Add optional FSDP2 training backend --- docs/sections/basic_usage/training.md | 29 ++ examples/configs/README.md | 1 + specforge/config/schema.py | 1 + specforge/launch.py | 16 + specforge/optimizer.py | 33 +- specforge/training/DESIGN.md | 12 +- specforge/training/assembly.py | 2 + specforge/training/backend.py | 202 +++++++----- specforge/training/controller.py | 1 + specforge/training/disaggregated.py | 4 + specforge/training/fsdp2.py | 91 ++++++ specforge/training/params.py | 14 + specforge/training/trainer.py | 9 +- tests/test_runtime/test_cli_config_build.py | 47 ++- .../test_dflash_offline_launch.py | 27 +- tests/test_runtime/test_disagg_launch.py | 12 +- tests/test_runtime/test_domain_trainer.py | 11 +- tests/test_runtime/test_domino_launch.py | 12 +- .../test_runtime/test_dspark_disagg_launch.py | 12 +- .../test_dspark_offline_launch.py | 12 +- tests/test_runtime/test_equiv_4rank.py | 13 +- tests/test_runtime/test_export.py | 24 +- tests/test_runtime/test_fsdp2_backend.py | 290 ++++++++++++++++++ .../test_runtime/test_offline_launch_fsdp.py | 12 +- 24 files changed, 786 insertions(+), 101 deletions(-) create mode 100644 specforge/training/fsdp2.py create mode 100644 specforge/training/params.py create mode 100644 tests/test_runtime/test_fsdp2_backend.py diff --git a/docs/sections/basic_usage/training.md b/docs/sections/basic_usage/training.md index c7c055b63..3931a8920 100644 --- a/docs/sections/basic_usage/training.md +++ b/docs/sections/basic_usage/training.md @@ -487,6 +487,27 @@ publishes the exact schedule horizon and the consumer trains to EOF. ## Parallel topologies +`training.backend` selects `fsdp` (the existing FSDP1 implementation, default) +or `fsdp2` (PyTorch composable `fully_shard`). For example: + +```yaml +training: + backend: fsdp2 + fsdp_sharding: SHARD_GRAD_OP +``` + +Both backends retain BF16 compute, FP32 optimizer masters, gradient accumulation, +and `training.optimizer_cpu_offload`. `SHARD_GRAD_OP` keeps parameters gathered +between forward and backward; `FULL_SHARD` reshards them after forward. Both +backends use DDP for `NO_SHARD`. Configured sharding takes precedence over the +legacy `FSDP_SHARDING` environment fallback used by direct Python builders. + +FSDP2 uses per-parameter DTensor shards and explicit block-level sharding. Its +memory management avoids FSDP1's CPU all-gather rate limiter and provides a +foundation for future DTensor-based parallelism. Throughput and peak memory +still depend on the model and sharding policy; selecting FSDP2 alone does not +guarantee a speedup or enable tensor parallelism. + The launcher creates every process group from the typed run config: - Online target TP/EP belongs to each external SGLang capture server, not the @@ -658,6 +679,14 @@ The producer itself is not restarted or resumed. Optimizer/FSDP checkpoints currently require the same trainer world size; control-plane ref redistribution does not imply optimizer-state resharding. +Resume also requires the same training backend and sharding strategy. Old +checkpoints without backend metadata are treated as FSDP1 checkpoints. FSDP1's +flat-parameter optimizer shards cannot be loaded into FSDP2's per-parameter +layout. To switch backends, export the draft and start a new run from its model +weights; optimizer moments, scheduler position, and FP32 master precision are +not transferred by that workflow. FSDP2 retains the existing full draft-weight +checkpoint format, so HF and SGLang export commands remain the same. + Training metrics are printed every `training.log_interval` steps and forwarded to the configured tracking backend. diff --git a/examples/configs/README.md b/examples/configs/README.md index 905c4eb10..a2ad8d02a 100644 --- a/examples/configs/README.md +++ b/examples/configs/README.md @@ -273,6 +273,7 @@ Common fields: | Field | Default | What to write | | --- | --- | --- | | `training.strategy` | `eagle3` | `eagle3`, `peagle`, `dflash`, `domino`, or `dspark`. | +| `training.backend` | `fsdp` | `fsdp` for FSDP1 or `fsdp2` for composable sharding. Both retain the DDP path for `NO_SHARD`; resume requires the original backend. | | `training.num_epochs` | `1` | Positive passes over a finite source. | | `training.max_steps` | `null` | Positive hard stop in optimizer steps. If it is set while `total_steps` is omitted, it is also the fallback schedule horizon. | | `training.total_steps` | `null` | Positive optimizer/loss schedule horizon; it does not itself stop an online stream. A finite online disaggregated run may omit both fields: the producer publishes the exact horizon derived from prepared prompts, epochs, DP size, batch size, and accumulation. | diff --git a/specforge/config/schema.py b/specforge/config/schema.py index 32eed42cf..f8728bf32 100644 --- a/specforge/config/schema.py +++ b/specforge/config/schema.py @@ -894,6 +894,7 @@ def _validate_mode(self): class TrainingConfig(StrictConfigModel): strategy: str = "eagle3" + backend: Literal["fsdp", "fsdp2"] = "fsdp" num_epochs: int = Field(default=1, gt=0) max_steps: Optional[int] = Field(default=None, gt=0) total_steps: Optional[int] = Field(default=None, gt=0) diff --git a/specforge/launch.py b/specforge/launch.py index 030e5388f..77f3aca3b 100644 --- a/specforge/launch.py +++ b/specforge/launch.py @@ -46,6 +46,8 @@ def _assemble_trainer( model, target_head, optimizer_factory, + training_backend: str = "fsdp", + fsdp_sharding: Optional[str] = None, run_id: str, output_dir: str, batch_size: int, @@ -121,6 +123,8 @@ def _assemble_trainer( model=model, target_head=target_head, optimizer_factory=optimizer_factory, + training_backend=training_backend, + fsdp_sharding=fsdp_sharding, run_id=run_id, output_dir=output_dir, batch_size=batch_size, @@ -546,6 +550,8 @@ def build_offline_runtime( draft_model, target_head, optimizer_factory, + training_backend: str = "fsdp", + fsdp_sharding: Optional[str] = None, run_id: str, output_dir: str, ttt_length: int = 7, @@ -632,6 +638,8 @@ def refs_for_epoch(epoch): target_head if algorithm.providers.step.uses_external_target_head else None ), optimizer_factory=optimizer_factory, + training_backend=training_backend, + fsdp_sharding=fsdp_sharding, run_id=run_id, output_dir=output_dir, batch_size=batch_size, @@ -672,6 +680,8 @@ def build_disagg_offline_runtime( draft_model, target_head, optimizer_factory, + training_backend: str = "fsdp", + fsdp_sharding: Optional[str] = None, run_id: str, output_dir: str, ttt_length: int = 7, @@ -754,6 +764,8 @@ def refs_for_epoch(epoch): target_head if algorithm.providers.step.uses_external_target_head else None ), optimizer_factory=optimizer_factory, + training_backend=training_backend, + fsdp_sharding=fsdp_sharding, run_id=run_id, output_dir=output_dir, batch_size=batch_size, @@ -1525,6 +1537,8 @@ def build_disagg_online_consumer( channel, draft_model, optimizer_factory, + training_backend: str = "fsdp", + fsdp_sharding: Optional[str] = None, run_id: str, output_dir: str, target_head=None, @@ -1876,6 +1890,8 @@ def stop_distributor_and_drain() -> None: else None ), optimizer_factory=optimizer_factory, + training_backend=training_backend, + fsdp_sharding=fsdp_sharding, run_id=run_id, output_dir=output_dir, batch_size=batch_size, diff --git a/specforge/optimizer.py b/specforge/optimizer.py index 6fd4b6f46..fd33e76e5 100644 --- a/specforge/optimizer.py +++ b/specforge/optimizer.py @@ -5,6 +5,7 @@ import torch.distributed as dist from specforge.lr_scheduler import ConstantWarmupLR, CosineAnnealingWarmupLR +from specforge.training.params import local_tensor from specforge.utils import print_on_rank0 logger = logging.getLogger(__name__) @@ -57,9 +58,9 @@ def __init__( self.offload_master = bool(offload_master) self.fp32_params = [ ( - p.detach().to(device="cpu", dtype=torch.float32).clone() + local_tensor(p).detach().to(device="cpu", dtype=torch.float32).clone() if self.offload_master - else p.detach().clone().to(torch.float32) + else local_tensor(p).detach().clone().to(torch.float32) ) for p in self.model_params ] @@ -132,7 +133,11 @@ def _grad_norm_and_clip_coefficient(self): """Compute the global grad norm from the model params on their own device, where NCCL can reduce it safely, without materialising master gradients first.""" - grads = [p.grad.detach() for p in self.model_params if p.grad is not None] + grads = [ + local_tensor(p.grad).detach() + for p in self.model_params + if p.grad is not None + ] if grads: total_norm_sq = _sum_of_squares(grads) else: @@ -225,14 +230,18 @@ def step(self, *, loss_denominator=None): mp.grad = None continue if self.offload_master: - master_grad = p.grad.detach().to( - device=mp.device, - dtype=torch.float32, + master_grad = ( + local_tensor(p.grad) + .detach() + .to( + device=mp.device, + dtype=torch.float32, + ) ) master_grad.mul_(host_values[1]) else: master_grad = torch.empty_like(mp) - model_grads.append(p.grad.detach()) + model_grads.append(local_tensor(p.grad).detach()) master_grads.append(master_grad) mp.grad = master_grad if master_grads: @@ -245,10 +254,12 @@ def step(self, *, loss_denominator=None): with torch.no_grad(): if self.offload_master: for p, mp in zip(self.model_params, self.fp32_params): - p.data.copy_(mp.data.to(device=p.device, dtype=p.dtype)) + local_tensor(p).data.copy_( + mp.data.to(device=p.device, dtype=p.dtype) + ) elif self.model_params: torch._foreach_copy_( - [p.data for p in self.model_params], + [local_tensor(p).data for p in self.model_params], [mp.data for mp in self.fp32_params], ) for p in self.model_params: @@ -326,7 +337,9 @@ def load_state_dict(self, state_dict): ) with torch.no_grad(): for p, mp in zip(self.model_params, self.fp32_params): - mp.data.copy_(p.detach().to(device=mp.device, dtype=mp.dtype)) + mp.data.copy_( + local_tensor(p).detach().to(device=mp.device, dtype=mp.dtype) + ) def state_dict(self): return { diff --git a/specforge/training/DESIGN.md b/specforge/training/DESIGN.md index 4a41b9433..d47a0ce05 100644 --- a/specforge/training/DESIGN.md +++ b/specforge/training/DESIGN.md @@ -23,11 +23,19 @@ Below that boundary, `TrainerController` owns the epoch loop, optimizer-step counting, interval checkpoints, and durable acknowledgements; `TrainerCore` owns one branch-free train step and the accumulation boundary; `DraftTrainStrategy` owns model-specific validation, forward/loss, target -projection, and checkpoint filtering; `FSDPTrainingBackend` owns wrapping, -backward, optimizer steps, distributed gradient norms, and full training state. +projection, and checkpoint filtering. `DistributedTrainingBackend` shares the +optimizer lifecycle, replicated DDP path, local gradient scaling, and RNG +state. `FSDPTrainingBackend` and `FSDP2TrainingBackend` own their sharding and +model-state APIs; FSDP2 also implements accumulation with +`set_requires_gradient_sync`. `training.backend` selects the implementation, +defaulting to the original `fsdp` backend. Checkpoint rotation and the latest pointer live in `specforge.training.checkpoint`. Resume restores each rank's optimizer/RNG state and repositions fixed offline refs through `FeatureDataLoader.seek()`. +Rank-local checkpoint metadata prevents loading optimizer shards with a +different backend, sharding strategy, or world size. FSDP2's BF16 optimizer +accesses local DTensor storage at use time, retaining ordinary FP32 master +tensors (including CPU offload) and exactly one explicit global-norm reduction. ## Internal mechanics diff --git a/specforge/training/assembly.py b/specforge/training/assembly.py index 07e63405c..4a17fef87 100644 --- a/specforge/training/assembly.py +++ b/specforge/training/assembly.py @@ -553,6 +553,8 @@ def _common_launch_kwargs( algorithm=algorithm, modality=cfg.model.input_modality, optimizer_factory=_optimizer_factory(cfg), + training_backend=t.backend, + fsdp_sharding=t.fsdp_sharding, run_id=cfg.run_id, output_dir=cfg.output_dir, batch_size=t.batch_size, diff --git a/specforge/training/backend.py b/specforge/training/backend.py index f4d673bb4..9279f664a 100644 --- a/specforge/training/backend.py +++ b/specforge/training/backend.py @@ -8,7 +8,7 @@ # http://www.apache.org/licenses/LICENSE-2.0 """TrainingBackend: model wrapping / backward / optimizer step / state dict. -FSDP-only for now. ``ParallelConfig`` carries the process groups created by the +``ParallelConfig`` carries the process groups created by the single distributed lifecycle: trainer TP (fixed at one by public builders) plus draft DP/USP topology. """ @@ -16,7 +16,6 @@ from __future__ import annotations import abc -import contextlib import logging import os from dataclasses import dataclass, field @@ -26,6 +25,8 @@ import torch.distributed as dist import torch.nn as nn +from specforge.training.params import local_tensor + def _foreach_scale_(tensors: List[torch.Tensor], factor: torch.Tensor) -> None: """``tensor.mul_(factor)`` for every tensor, one fused launch per group. @@ -85,14 +86,14 @@ def from_distributed( tp_size: int = 1, sp_ulysses_size: int = 1, sp_ring_size: int = 1, - sharding_strategy: str = "SHARD_GRAD_OP", + sharding_strategy: Optional[str] = None, param_dtype: torch.dtype = torch.bfloat16, ) -> "ParallelConfig": """Carry every group built by :func:`specforge.distributed.init_distributed`.""" - # Env override for the FSDP sharding strategy — e.g. FSDP_SHARDING=NO_SHARD - # runs DDP-style (full params replicated, one grad all-reduce, no param - # all-gather). Default unchanged when the env var is unset. - sharding_strategy = os.environ.get("FSDP_SHARDING", sharding_strategy) + # Typed config is authoritative; direct Python callers may still use + # the legacy environment fallback, e.g. FSDP_SHARDING=NO_SHARD for DDP. + if sharding_strategy is None: + sharding_strategy = os.environ.get("FSDP_SHARDING", "SHARD_GRAD_OP") if not dist.is_initialized(): return cls( world_size=1, @@ -173,12 +174,13 @@ def state_dict(self) -> dict: ... def load_state_dict(self, state: dict) -> None: ... -class FSDPTrainingBackend(TrainingBackend): - """FSDP1 backend for the canonical SpecForge training math: FSDP with - ``use_orig_params=True`` / bf16 mixed precision over the configured process - group, optimizer targeting the inner trainable submodule.""" +class DistributedTrainingBackend(TrainingBackend): + """Shared model/optimizer lifecycle for the FSDP1 and FSDP2 backends. - name = "fsdp" + Subclasses own parameter sharding and full model state dictionaries. + Replicated execution, local-shard gradient scaling, optimizer and RNG state + use the same contract for both implementations. + """ def __init__( self, @@ -237,8 +239,6 @@ def prepare_model( self._wrapped = False self._wrapper_kind = "none" else: - import functools - pc = self.parallel_config ignored_frozen_modules = self._frozen_target_modules(model) # DFlash-family models expose their transformer block class through @@ -277,38 +277,12 @@ def prepare_model( ) self._wrapper_kind = "ddp" else: - from torch.distributed.fsdp import BackwardPrefetch - from torch.distributed.fsdp import FullyShardedDataParallel as FSDP - from torch.distributed.fsdp import MixedPrecision, ShardingStrategy - from torch.distributed.fsdp.wrap import transformer_auto_wrap_policy - - sharding = getattr(ShardingStrategy, pc.sharding_strategy) - fsdp_kwargs = dict( - use_orig_params=True, - mixed_precision=MixedPrecision( - param_dtype=pc.param_dtype, buffer_dtype=torch.float32 - ), - sharding_strategy=sharding, - process_group=pc.fsdp_process_group, - ) - if ignored_frozen_modules: - fsdp_kwargs["ignored_modules"] = ignored_frozen_modules - if block_classes: - fsdp_kwargs.update( - auto_wrap_policy=functools.partial( - transformer_auto_wrap_policy, - transformer_layer_cls=block_classes, - ), - forward_prefetch=True, - backward_prefetch=BackwardPrefetch.BACKWARD_PRE, - limit_all_gathers=True, - ) - model = FSDP(model, **fsdp_kwargs) - self._wrapper_kind = "fsdp" + model = self._shard_model(model, block_classes, ignored_frozen_modules) + self._wrapper_kind = self.name self.module = model self._wrapped = True self.auto_wrap_block_classes = ( - block_classes if self._wrapper_kind == "fsdp" else set() + block_classes if self._wrapper_kind != "ddp" else set() ) self.ignored_frozen_modules = ignored_frozen_modules if self._optimizer_factory is not None: @@ -317,6 +291,9 @@ def prepare_model( self._configure_optimizer_grad_norm() return self.module + @abc.abstractmethod + def _shard_model(self, model, block_classes, ignored_frozen_modules): ... + def set_optimizer(self, optimizer) -> None: self.optimizer = optimizer self._configure_optimizer_grad_norm() @@ -349,7 +326,11 @@ def scale_gradients(self, factor: torch.Tensor) -> None: raise RuntimeError("scale_gradients called before prepare_model") with torch.no_grad(): _foreach_scale_( - [p.grad for p in self.module.parameters() if p.grad is not None], + [ + local_tensor(p.grad) + for p in self.module.parameters() + if p.grad is not None + ], factor, ) @@ -367,22 +348,24 @@ def step( """ if self.optimizer is None: raise RuntimeError( - "FSDPTrainingBackend.step called before optimizer is set" + f"{type(self).__name__}.step called before optimizer is set" ) if loss_denominator is None: return self.optimizer.step() return self.optimizer.step(loss_denominator=loss_denominator) def state_dict(self) -> dict: - """Full training state ``{"model", "optimizer", "rng"}`` for resume. + """Model, optimizer, RNG and sharding metadata for resume. - ``model`` is gathered rank0-only (``{}`` on other ranks when wrapped). + Sharded model weights are gathered on rank zero; FSDP1 may also return + ignored replicated target tables on other ranks. FSDP optimizer state is rank-local; DDP optimizer state is replicated. RNG state is always rank-local. """ if self.module is None: raise RuntimeError("state_dict called before prepare_model") return { + "metadata": self._checkpoint_metadata(), "model": self._module_state_dict(), "optimizer": ( self.optimizer.state_dict() if self.optimizer is not None else None @@ -392,6 +375,18 @@ def state_dict(self) -> dict: def load_state_dict(self, state: dict) -> None: """Restore whichever of module weights / optimizer / RNG the state carries.""" + if state.get("optimizer") is not None: + metadata = state.get("metadata") + if metadata is None: + # Checkpoints predating backend selection were written by FSDP1. + metadata = {"backend": "fsdp"} + for key, current in self._checkpoint_metadata().items(): + if key in metadata and metadata[key] != current: + raise ValueError( + f"optimizer checkpoint {key}={metadata[key]!r} does not " + f"match this run ({current!r}); resume with the original " + "backend and sharding layout, or load model weights only" + ) if state.get("model") is not None: self._load_module_state_dict(state["model"]) if self.optimizer is not None and state.get("optimizer") is not None: @@ -399,39 +394,35 @@ def load_state_dict(self, state: dict) -> None: if state.get("rng") is not None: self._set_rng_state(state["rng"]) - def _full_state_ctx(self, state_dict_config=None): - """FULL_STATE_DICT context for a wrapped module; a no-op when unwrapped.""" - if self._wrapper_kind != "fsdp": - return contextlib.nullcontext() - from torch.distributed.fsdp import FullyShardedDataParallel as FSDP - from torch.distributed.fsdp import StateDictType - - return FSDP.state_dict_type( - self.module, StateDictType.FULL_STATE_DICT, state_dict_config - ) + def _checkpoint_metadata(self) -> dict: + return { + "backend": self.name, + "sharding_strategy": self.parallel_config.sharding_strategy, + "world_size": self.parallel_config.world_size, + } def _module_state_dict(self) -> dict: if self._wrapper_kind == "ddp": if dist.is_initialized() and dist.get_rank() != 0: return {} return self.module.module.state_dict() - if self._wrapper_kind != "fsdp": - return self.module.state_dict() - from torch.distributed.fsdp import FullStateDictConfig - - # gather to rank0 CPU only — materializing the full model on every - # rank's GPU is wasted memory when only rank0 writes it. - cfg = FullStateDictConfig(offload_to_cpu=True, rank0_only=True) - with self._full_state_ctx(cfg): + if not self._wrapped: return self.module.state_dict() + return self._sharded_model_state_dict() def _load_module_state_dict(self, model_state: dict) -> None: - # every rank loads the full state dict read from the shared file. if self._wrapper_kind == "ddp": self.module.module.load_state_dict(model_state) - return - with self._full_state_ctx(): + elif not self._wrapped: self.module.load_state_dict(model_state) + else: + self._load_sharded_model_state_dict(model_state) + + @abc.abstractmethod + def _sharded_model_state_dict(self) -> dict: ... + + @abc.abstractmethod + def _load_sharded_model_state_dict(self, model_state: dict) -> None: ... @staticmethod def _rng_state() -> dict: @@ -472,4 +463,77 @@ def _set_rng_state(rng: dict) -> None: module.set_rng_state(state, module.current_device()) -__all__ = ["ParallelConfig", "TrainingBackend", "FSDPTrainingBackend"] +class FSDPTrainingBackend(DistributedTrainingBackend): + """Original FSDP1 implementation with BF16 compute and original parameters.""" + + name = "fsdp" + + def _shard_model(self, model, block_classes, ignored_frozen_modules): + import functools + + from torch.distributed.fsdp import BackwardPrefetch + from torch.distributed.fsdp import FullyShardedDataParallel as FSDP + from torch.distributed.fsdp import MixedPrecision, ShardingStrategy + from torch.distributed.fsdp.wrap import transformer_auto_wrap_policy + + pc = self.parallel_config + fsdp_kwargs = dict( + use_orig_params=True, + mixed_precision=MixedPrecision( + param_dtype=pc.param_dtype, buffer_dtype=torch.float32 + ), + sharding_strategy=getattr(ShardingStrategy, pc.sharding_strategy), + process_group=pc.fsdp_process_group, + ) + if ignored_frozen_modules: + fsdp_kwargs["ignored_modules"] = ignored_frozen_modules + if block_classes: + fsdp_kwargs.update( + auto_wrap_policy=functools.partial( + transformer_auto_wrap_policy, + transformer_layer_cls=block_classes, + ), + forward_prefetch=True, + backward_prefetch=BackwardPrefetch.BACKWARD_PRE, + limit_all_gathers=True, + ) + return FSDP(model, **fsdp_kwargs) + + def _full_state_ctx(self, state_dict_config=None): + from torch.distributed.fsdp import FullyShardedDataParallel as FSDP + from torch.distributed.fsdp import StateDictType + + return FSDP.state_dict_type( + self.module, StateDictType.FULL_STATE_DICT, state_dict_config + ) + + def _sharded_model_state_dict(self) -> dict: + from torch.distributed.fsdp import FullStateDictConfig + + cfg = FullStateDictConfig(offload_to_cpu=True, rank0_only=True) + with self._full_state_ctx(cfg): + return self.module.state_dict() + + def _load_sharded_model_state_dict(self, model_state: dict) -> None: + with self._full_state_ctx(): + self.module.load_state_dict(model_state) + + +def create_training_backend(name: str, parallel_config: ParallelConfig, **kwargs): + """Select a backend without importing FSDP2 on the default FSDP1 path.""" + if name == "fsdp": + return FSDPTrainingBackend(parallel_config, **kwargs) + if name == "fsdp2": + from specforge.training.fsdp2 import FSDP2TrainingBackend + + return FSDP2TrainingBackend(parallel_config, **kwargs) + raise ValueError(f"unsupported training backend: {name!r}; expected fsdp or fsdp2") + + +__all__ = [ + "ParallelConfig", + "TrainingBackend", + "DistributedTrainingBackend", + "FSDPTrainingBackend", + "create_training_backend", +] diff --git a/specforge/training/controller.py b/specforge/training/controller.py index d711bb9f9..cef88ebc6 100644 --- a/specforge/training/controller.py +++ b/specforge/training/controller.py @@ -1227,6 +1227,7 @@ def save_checkpoint(self, step: int) -> Checkpoint: shared, step, rank_state={ + "metadata": full.get("metadata"), "optimizer": None if replicated_optimizer else full["optimizer"], "rng": full["rng"], }, diff --git a/specforge/training/disaggregated.py b/specforge/training/disaggregated.py index b30eb274c..7752fc268 100644 --- a/specforge/training/disaggregated.py +++ b/specforge/training/disaggregated.py @@ -538,6 +538,8 @@ def produce() -> int: draft_model=bundle.model, target_head=bundle.target_head, optimizer_factory=optimizer_factory(cfg), + training_backend=cfg.training.backend, + fsdp_sharding=cfg.training.fsdp_sharding, run_id=cfg.run_id, output_dir=cfg.output_dir, ttt_length=cfg.training.ttt_length, @@ -817,6 +819,8 @@ def produce() -> int: draft_model=bundle.model, target_head=bundle.target_head, optimizer_factory=optimizer_factory(cfg), + training_backend=cfg.training.backend, + fsdp_sharding=cfg.training.fsdp_sharding, run_id=cfg.run_id, output_dir=cfg.output_dir, batch_size=cfg.training.batch_size, diff --git a/specforge/training/fsdp2.py b/specforge/training/fsdp2.py new file mode 100644 index 000000000..2adf1daba --- /dev/null +++ b/specforge/training/fsdp2.py @@ -0,0 +1,91 @@ +"""Composable FSDP2 backend with the same training contract as FSDP1.""" + +import torch +from torch.distributed.device_mesh import DeviceMesh +from torch.distributed.fsdp import MixedPrecisionPolicy, fully_shard + +from specforge.training.backend import DistributedTrainingBackend + + +class FSDP2TrainingBackend(DistributedTrainingBackend): + name = "fsdp2" + + def _shard_model(self, model, block_classes, ignored_frozen_modules): + pc = self.parallel_config + if pc.sharding_strategy not in ("FULL_SHARD", "SHARD_GRAD_OP"): + raise ValueError(f"unsupported FSDP2 sharding: {pc.sharding_strategy!r}") + device = next(model.parameters()).device + # The existing FSDP group spans WORLD, including sequence-parallel ranks. + # A draft-DP-only mesh would silently change the training reduction. + mesh = DeviceMesh.from_group( + pc.fsdp_process_group or torch.distributed.group.WORLD, + device_type=device.type, + ) + ignored_params = { + p for module in ignored_frozen_modules for p in module.parameters() + } + # FSDP1 keeps floating buffers (e.g. RoPE) in FP32. FSDP2 does not + # manage buffer precision, so preserve that policy explicitly. + for module in model.modules(): + if module in ignored_frozen_modules: + continue + for name, buffer in module.named_buffers(recurse=False): + if buffer.is_floating_point(): + setattr(module, name, buffer.float()) + kwargs = dict( + mesh=mesh, + reshard_after_forward=pc.sharding_strategy == "FULL_SHARD", + ignored_params=ignored_params, + ) + # Bottom-up application preserves the FSDP1 block/root boundaries. + # Leave non-block draft parameters on the composite root: some losses + # call custom draft head methods after draft_model.forward has returned. + for module in reversed(list(model.modules())): + if module is not model and type(module) in block_classes: + fully_shard( + module, + mp_policy=MixedPrecisionPolicy( + param_dtype=pc.param_dtype, cast_forward_inputs=False + ), + **kwargs, + ) + fully_shard( + model, + mp_policy=MixedPrecisionPolicy(param_dtype=pc.param_dtype), + **kwargs, + ) + return model + + def backward(self, loss: torch.Tensor, *, is_boundary: bool = True) -> None: + if self._wrapper_kind != "fsdp2": + return super().backward(loss, is_boundary=is_boundary) + self.module.set_requires_gradient_sync(is_boundary) + try: + loss.backward() + finally: + self.module.set_requires_gradient_sync(True) + + def _sharded_model_state_dict(self) -> dict: + from torch.distributed.checkpoint.state_dict import ( + StateDictOptions, + get_model_state_dict, + ) + + # All ranks participate; full_state_dict + cpu_offload returns the + # gathered ordinary tensors on rank zero and an empty dict elsewhere. + return get_model_state_dict( + self.module, + options=StateDictOptions(full_state_dict=True, cpu_offload=True), + ) + + def _load_sharded_model_state_dict(self, model_state: dict) -> None: + from torch.distributed.checkpoint.state_dict import ( + StateDictOptions, + set_model_state_dict, + ) + + set_model_state_dict( + self.module, + model_state, + options=StateDictOptions(full_state_dict=True, broadcast_from_rank0=True), + ) diff --git a/specforge/training/params.py b/specforge/training/params.py new file mode 100644 index 000000000..cd79c5122 --- /dev/null +++ b/specforge/training/params.py @@ -0,0 +1,14 @@ +"""Local storage access shared by gradient scaling and FP32 master updates.""" + +import torch +from torch.distributed.tensor import DTensor + + +def local_tensor(tensor: torch.Tensor) -> torch.Tensor: + """Return a shard view without gathering or changing a DTensor's placement. + + Callers explicitly reduce local gradient norms over the owning shard group; + mixing DTensor reductions with that collective would count shards twice. + Resolve the view at use time since FSDP replaces storage around forward. + """ + return tensor.to_local() if isinstance(tensor, DTensor) else tensor diff --git a/specforge/training/trainer.py b/specforge/training/trainer.py index 846bd3a7e..68ed08eaa 100644 --- a/specforge/training/trainer.py +++ b/specforge/training/trainer.py @@ -31,7 +31,7 @@ checkpoint_key_fingerprint, ) from specforge.runtime.data_plane import FeatureDataLoader, FeatureStore -from specforge.training.backend import FSDPTrainingBackend, ParallelConfig +from specforge.training.backend import ParallelConfig, create_training_backend from specforge.training.checkpoint import CheckpointManager from specforge.training.controller import TrainerController, TrainerCore @@ -87,6 +87,8 @@ def __init__( model, target_head, optimizer_factory, + training_backend: str = "fsdp", + fsdp_sharding: Optional[str] = None, run_id: str, output_dir: str, batch_size: int, @@ -426,11 +428,14 @@ def set_epoch(epoch): del state, saved_weights parallel = ParallelConfig.from_distributed( + sharding_strategy=fsdp_sharding, tp_size=tp_size, sp_ulysses_size=sp_ulysses_size, sp_ring_size=sp_ring_size, ) - backend = FSDPTrainingBackend(parallel, optimizer_factory=optimizer_factory) + backend = create_training_backend( + training_backend, parallel, optimizer_factory=optimizer_factory + ) # FSDP-wrap the composite model and build the optimizer over the inner draft # AFTER wrapping; the strategy MUST run forward through the wrapped module so # FSDP is actually in the forward/backward path (not bypassed at >1 rank). diff --git a/tests/test_runtime/test_cli_config_build.py b/tests/test_runtime/test_cli_config_build.py index 9fc027fd5..484094bd2 100644 --- a/tests/test_runtime/test_cli_config_build.py +++ b/tests/test_runtime/test_cli_config_build.py @@ -20,6 +20,12 @@ @unittest.skipUnless(CUDA, "cli config build requires CUDA") class TestCliConfigBuild(unittest.TestCase): def test_config_build_matches_programmatic_and_trains(self): + self._check_backend("fsdp") + + def test_config_selects_fsdp2_and_trains(self): + self._check_backend("fsdp2") + + def _check_backend(self, training_backend): torch.manual_seed(0) from tests.test_runtime import _fixtures as fx @@ -52,6 +58,7 @@ def test_config_build_matches_programmatic_and_trains(self): "ttt_length": 3, "max_steps": 4, "max_checkpoints": 2, + "save_interval": 1, "log_interval": 1, }, "run_id": "cli-gate", @@ -61,12 +68,24 @@ def test_config_build_matches_programmatic_and_trains(self): with open(yaml_path, "w") as f: yaml.safe_dump(run_config, f) - cfg = load_config(yaml_path, ["training.max_steps=2"]) # override applies + cfg = load_config( + yaml_path, + [ + "training.max_steps=2", + f"training.backend={training_backend}", + "training.fsdp_sharding=FULL_SHARD", + ], + ) + construction_rng = torch.get_rng_state() run = build_application_run(resolve_run(cfg)) trainer = run.trainer # package-level assembly is the single wiring the CLI executes self.assertIsInstance(trainer, Trainer) + self.assertEqual(trainer.backend.name, training_backend) + self.assertEqual( + trainer.backend.parallel_config.sharding_strategy, "FULL_SHARD" + ) self.assertEqual(trainer.run_id, "cli-gate") self.assertEqual(trainer.max_steps, 2) self.assertEqual(trainer.max_checkpoints, 2) @@ -82,6 +101,32 @@ def test_config_build_matches_programmatic_and_trains(self): os.path.islink(os.path.join(run_config["output_dir"], "cli-gate-latest")) ) + # Resume through the public config/Trainer/CheckpointManager seam, not + # just backend.load_state_dict; require the same final draft weights. + from specforge.training.checkpoint import CheckpointManager + + checkpoint = os.path.join(cfg.output_dir, "cli-gate-step1") + final = CheckpointManager.read_resume_state( + os.path.join(cfg.output_dir, "cli-gate-latest") + ) + self.assertEqual(final["backend"]["metadata"]["backend"], training_backend) + resumed_cfg = cfg.model_copy( + update={ + "training": cfg.training.model_copy(update={"resume_from": checkpoint}) + } + ) + torch.set_rng_state(construction_rng) + resumed = build_application_run(resolve_run(resumed_cfg)) + self.assertEqual(resumed.trainer.global_step, 1) + self.assertEqual(resumed.run(), 2) + actual = CheckpointManager.read_resume_state( + os.path.join(cfg.output_dir, "cli-gate-latest") + ) + for key, expected in final["draft_state_dict"].items(): + torch.testing.assert_close( + actual["draft_state_dict"][key], expected, rtol=0, atol=0 + ) + class TestCliDispatch(unittest.TestCase): def test_train_command_dispatches_one_resolved_run(self): diff --git a/tests/test_runtime/test_dflash_offline_launch.py b/tests/test_runtime/test_dflash_offline_launch.py index 2d1b7c536..aa40c008a 100644 --- a/tests/test_runtime/test_dflash_offline_launch.py +++ b/tests/test_runtime/test_dflash_offline_launch.py @@ -16,12 +16,24 @@ @unittest.skipUnless(CUDA, "DFlash offline launcher requires CUDA") class TestDFlashOfflineLaunch(unittest.TestCase): def test_dflash_trains_from_precomputed_features(self): + self._check_backend("fsdp") + + def test_fsdp2_trains_through_the_same_launcher(self): + self._check_backend("fsdp2") + + def test_dflash2_nested_blocks_train_with_fsdp2(self): + self._check_backend("fsdp2", dflash2=True) + + def _check_backend(self, training_backend, *, dflash2=False): from tests.test_runtime import _fixtures as fx fx.build_single_rank_distributed(port="29567") + from torch.distributed.fsdp import FSDPModule from torch.distributed.fsdp import FullyShardedDataParallel as FSDP + expected_wrapper = FSDPModule if training_backend == "fsdp2" else FSDP + from specforge.launch import build_offline_runtime from specforge.optimizer import BF16Optimizer @@ -41,6 +53,18 @@ def test_dflash_trains_from_precomputed_features(self): attention_backend="sdpa", ) self.assertEqual(width, hidden) + if dflash2: + from specforge.modeling.draft.dflash2 import DFlash2DraftModel + + config = model.draft_model.config + config.dflash_config.update( + conv_kernel_size=3, + conv_group_size=16, + selector_rank=8, + selector_top_k=4, + ) + model.draft_model = DFlash2DraftModel(config).cuda().bfloat16() + model.selector_loss_alpha = 0.1 def optimizer_factory(module): return BF16Optimizer( @@ -57,6 +81,7 @@ def optimizer_factory(module): draft_model=model, target_head=None, optimizer_factory=optimizer_factory, + training_backend=training_backend, run_id="dflash-offline", output_dir=os.path.join(workdir, "out"), max_len=sequence_length, @@ -66,7 +91,7 @@ def optimizer_factory(module): ) module = trainer.core.strategy.trainable_module() - self.assertIsInstance(module, FSDP) + self.assertIsInstance(module, expected_wrapper) self.assertEqual(trainer.fit(), 2) self.assertTrue(all(torch.isfinite(p).all() for p in module.parameters())) diff --git a/tests/test_runtime/test_disagg_launch.py b/tests/test_runtime/test_disagg_launch.py index 1a12f4b9b..150c36a20 100644 --- a/tests/test_runtime/test_disagg_launch.py +++ b/tests/test_runtime/test_disagg_launch.py @@ -148,11 +148,20 @@ def test_auth_required_consumer_must_present_token(self): @unittest.skipUnless(CUDA, "disagg launcher FSDP path requires CUDA") class TestDisaggLaunchFSDP(unittest.TestCase): def test_build_disagg_runtime_trains_through_fsdp(self): + self._check_backend("fsdp") + + def test_fsdp2_trains_through_the_same_launcher(self): + self._check_backend("fsdp2") + + def _check_backend(self, training_backend): torch.manual_seed(0) fx.build_single_rank_distributed(port="29577") + from torch.distributed.fsdp import FSDPModule from torch.distributed.fsdp import FullyShardedDataParallel as FSDP + expected_wrapper = FSDPModule if training_backend == "fsdp2" else FSDP + from specforge.launch import build_disagg_offline_runtime from specforge.optimizer import BF16Optimizer @@ -188,6 +197,7 @@ def optimizer_factory(draft_module): draft_model=eagle3_model, target_head=target_head, optimizer_factory=optimizer_factory, + training_backend=training_backend, run_id="e2e", output_dir=os.path.join(work, "out"), max_len=512, @@ -199,7 +209,7 @@ def optimizer_factory(draft_module): module = trainer.core.strategy.trainable_module() self.assertIsInstance( - module, FSDP, "strategy must hold the FSDP-wrapped module" + module, expected_wrapper, "strategy must hold the FSDP-wrapped module" ) step = trainer.fit() diff --git a/tests/test_runtime/test_domain_trainer.py b/tests/test_runtime/test_domain_trainer.py index d9df397d0..420abb6a6 100644 --- a/tests/test_runtime/test_domain_trainer.py +++ b/tests/test_runtime/test_domain_trainer.py @@ -127,7 +127,9 @@ def ack_train_refs(self, tid, ids, *, global_step, optimizer_durable): tr, FeatureDataLoader=FakeLoader, ParallelConfig=FakeParallel, - FSDPTrainingBackend=FakeBackend, + create_training_backend=lambda name, parallel, **kwargs: FakeBackend( + parallel, **kwargs + ), TrainerCore=FakeCore, TrainerController=FakeController, ): @@ -220,7 +222,12 @@ def test_offline_wiring_matches_assemble_trainer(self): self.assertIsNone(cap["ctrl_kw"]["ack_fn"]) self.assertEqual( cap["parallel_kw"], - {"tp_size": 1, "sp_ulysses_size": 1, "sp_ring_size": 1}, + { + "tp_size": 1, + "sp_ulysses_size": 1, + "sp_ring_size": 1, + "sharding_strategy": None, + }, ) # run identity rides the shared checkpoint payload, validated on resume diff --git a/tests/test_runtime/test_domino_launch.py b/tests/test_runtime/test_domino_launch.py index 52ff16650..9ec561686 100644 --- a/tests/test_runtime/test_domino_launch.py +++ b/tests/test_runtime/test_domino_launch.py @@ -96,12 +96,21 @@ def forward(self, **kwargs): @unittest.skipUnless(CUDA, "Domino offline launcher path requires CUDA") class TestDominoOfflineLaunch(unittest.TestCase): def test_domino_trains_from_precomputed_dflash_features(self): + self._check_backend("fsdp") + + def test_fsdp2_trains_through_the_same_launcher(self): + self._check_backend("fsdp2") + + def _check_backend(self, training_backend): from tests.test_runtime import _fixtures as fx fx.build_single_rank_distributed(port="29571") + from torch.distributed.fsdp import FSDPModule from torch.distributed.fsdp import FullyShardedDataParallel as FSDP + expected_wrapper = FSDPModule if training_backend == "fsdp2" else FSDP + from specforge.launch import build_offline_runtime from specforge.optimizer import BF16Optimizer @@ -137,6 +146,7 @@ def optimizer_factory(module): draft_model=model, target_head=None, optimizer_factory=optimizer_factory, + training_backend=training_backend, run_id="domino-offline", output_dir=os.path.join(workdir, "out"), max_len=sequence_length, @@ -148,7 +158,7 @@ def optimizer_factory(module): ) module = trainer.core.strategy.trainable_module() - self.assertIsInstance(module, FSDP) + self.assertIsInstance(module, expected_wrapper) self.assertEqual(trainer.fit(), 2) self.assertTrue(all(torch.isfinite(p).all() for p in module.parameters())) diff --git a/tests/test_runtime/test_dspark_disagg_launch.py b/tests/test_runtime/test_dspark_disagg_launch.py index 016998598..d1c702875 100644 --- a/tests/test_runtime/test_dspark_disagg_launch.py +++ b/tests/test_runtime/test_dspark_disagg_launch.py @@ -24,12 +24,21 @@ @unittest.skipUnless(CUDA, "DSpark disaggregated optimizer gate requires CUDA") class TestDSparkDisaggregatedLaunch(unittest.TestCase): def test_synthetic_server_features_train_through_canonical_consumer(self): + self._check_backend("fsdp") + + def test_fsdp2_trains_through_the_same_launcher(self): + self._check_backend("fsdp2") + + def _check_backend(self, training_backend): from tests.test_runtime import _fixtures as fx fx.build_single_rank_distributed(port="29579") + from torch.distributed.fsdp import FSDPModule from torch.distributed.fsdp import FullyShardedDataParallel as FSDP + expected_wrapper = FSDPModule if training_backend == "fsdp2" else FSDP + from specforge.launch import build_disagg_online_consumer from specforge.optimizer import BF16Optimizer from specforge.runtime.data_plane.mooncake_store import MooncakeFeatureStore @@ -67,6 +76,7 @@ def test_synthetic_server_features_train_through_canonical_consumer(self): feature_store=consumer_store, channel=channel, draft_model=model, + training_backend=training_backend, optimizer_factory=lambda module: BF16Optimizer( module, lr=1e-3, @@ -87,7 +97,7 @@ def test_synthetic_server_features_train_through_canonical_consumer(self): strategy = trainer.core.strategy self.assertIsInstance(strategy, DSparkTrainStrategy) - self.assertIsInstance(strategy.trainable_module(), FSDP) + self.assertIsInstance(strategy.trainable_module(), expected_wrapper) self.assertEqual( {cls.__name__ for cls in trainer.backend.auto_wrap_block_classes}, {"Qwen3DFlashDecoderLayer"}, diff --git a/tests/test_runtime/test_dspark_offline_launch.py b/tests/test_runtime/test_dspark_offline_launch.py index 3b008ccae..6caa1a67a 100644 --- a/tests/test_runtime/test_dspark_offline_launch.py +++ b/tests/test_runtime/test_dspark_offline_launch.py @@ -16,12 +16,21 @@ @unittest.skipUnless(CUDA, "DSpark offline launcher requires CUDA") class TestDSparkOfflineLaunch(unittest.TestCase): def test_dspark_trains_from_precomputed_target_features(self): + self._check_backend("fsdp") + + def test_fsdp2_trains_through_the_same_launcher(self): + self._check_backend("fsdp2") + + def _check_backend(self, training_backend): from tests.test_runtime import _fixtures as fx fx.build_single_rank_distributed(port="29580") + from torch.distributed.fsdp import FSDPModule from torch.distributed.fsdp import FullyShardedDataParallel as FSDP + expected_wrapper = FSDPModule if training_backend == "fsdp2" else FSDP + from specforge.launch import build_offline_runtime from specforge.optimizer import BF16Optimizer from specforge.training.strategies.base import DSparkTrainStrategy @@ -47,6 +56,7 @@ def test_dspark_trains_from_precomputed_target_features(self): algorithm=ALGORITHM, hidden_states_path=feature_dir, draft_model=model, + training_backend=training_backend, target_head=None, optimizer_factory=lambda module: BF16Optimizer( module, @@ -67,7 +77,7 @@ def test_dspark_trains_from_precomputed_target_features(self): strategy = trainer.core.strategy self.assertIsInstance(strategy, DSparkTrainStrategy) module = strategy.trainable_module() - self.assertIsInstance(module, FSDP) + self.assertIsInstance(module, expected_wrapper) self.assertEqual(trainer.fit(), 2) self.assertTrue(all(torch.isfinite(p).all() for p in module.parameters())) diff --git a/tests/test_runtime/test_equiv_4rank.py b/tests/test_runtime/test_equiv_4rank.py index 04a8640e0..7137b8dd2 100644 --- a/tests/test_runtime/test_equiv_4rank.py +++ b/tests/test_runtime/test_equiv_4rank.py @@ -55,7 +55,9 @@ def _build_model(workdir: str, attention_backend: str): ).cuda() -def _worker(rank: int, world_size: int, port: int, workdir: str) -> None: +def _worker( + rank: int, world_size: int, port: int, workdir: str, training_backend: str +) -> None: fx.init_rank_distributed( rank, world_size, @@ -146,6 +148,7 @@ def optimizer_factory(module): draft_model=usp_model, target_head=target_head, optimizer_factory=optimizer_factory, + training_backend=training_backend, run_id="usp-parity", output_dir=os.path.join(workdir, "output"), ttt_length=3, @@ -195,6 +198,12 @@ def optimizer_factory(module): ) class TestEquiv4Rank(unittest.TestCase): def test_dp2_sp2_trainer_loss_matches_full_sequence_reference(self): + self._check_backend("fsdp") + + def test_fsdp2_dp2_sp2_matches_full_sequence_reference(self): + self._check_backend("fsdp2") + + def _check_backend(self, training_backend): import torch.multiprocessing as mp from tests.utils import get_available_port @@ -209,7 +218,7 @@ def test_dp2_sp2_trainer_loss_matches_full_sequence_reference(self): mp.spawn( _worker, nprocs=WORLD_SIZE, - args=(WORLD_SIZE, get_available_port(), work), + args=(WORLD_SIZE, get_available_port(), work, training_backend), join=True, ) diff --git a/tests/test_runtime/test_export.py b/tests/test_runtime/test_export.py index 0f9d5083d..4d1825919 100644 --- a/tests/test_runtime/test_export.py +++ b/tests/test_runtime/test_export.py @@ -304,6 +304,8 @@ def test_cli_mapping_override_with_external_embedding(self): @unittest.skipUnless(CUDA, "export round-trip requires CUDA") class TestExporters(unittest.TestCase): + training_backend = "fsdp" + @classmethod def setUpClass(cls): torch.manual_seed(0) @@ -315,7 +317,7 @@ def setUpClass(cls): from specforge.modeling.auto import AutoDraftModel, AutoDraftModelConfig from specforge.modeling.target.target_head import TargetHead from specforge.optimizer import BF16Optimizer - from specforge.training.backend import FSDPTrainingBackend, ParallelConfig + from specforge.training.backend import ParallelConfig, create_training_backend from specforge.training.controller import TrainerController, TrainerCore from specforge.training.strategies.base import Eagle3TrainStrategy @@ -349,13 +351,17 @@ def setUpClass(cls): ) ) - opt = BF16Optimizer( - dm, lr=1e-3, max_grad_norm=0.5, warmup_ratio=0.0, total_steps=10 + backend = create_training_backend( + cls.training_backend, + ParallelConfig.from_distributed(), + optimizer_factory=lambda draft: BF16Optimizer( + draft, lr=1e-3, max_grad_norm=0.5, warmup_ratio=0.0, total_steps=10 + ), + ) + wrapped = backend.prepare_model( + model, wrap=cls.training_backend == "fsdp2", optimizer_target=dm ) - backend = FSDPTrainingBackend(ParallelConfig.from_distributed()) - backend.prepare_model(model, wrap=False) - backend.set_optimizer(opt) - strategy = Eagle3TrainStrategy(model, target_head=head) + strategy = Eagle3TrainStrategy(wrapped, target_head=head) ctrl = TrainerController( TrainerCore(strategy, backend), run_id="exp", @@ -464,5 +470,9 @@ def test_missing_required_serving_key_fails_loudly(self): _serving_state({"fc.weight": torch.zeros(1)}, {}) +class TestFSDP2Exporters(TestExporters): + training_backend = "fsdp2" + + if __name__ == "__main__": unittest.main(verbosity=2) diff --git a/tests/test_runtime/test_fsdp2_backend.py b/tests/test_runtime/test_fsdp2_backend.py new file mode 100644 index 000000000..e156d4067 --- /dev/null +++ b/tests/test_runtime/test_fsdp2_backend.py @@ -0,0 +1,290 @@ +"""FSDP2 numerical, local-master, and checkpoint gates on real process groups.""" + +import copy +import os +import sys +import tempfile +import unittest +from unittest import mock + +import torch +import torch.distributed as dist +import torch.nn as nn +from torch.distributed.tensor import DTensor + +from specforge.optimizer import BF16Optimizer +from specforge.training.backend import ParallelConfig, create_training_backend + + +class TinyBlock(nn.Module): + def __init__(self): + super().__init__() + self.linear = nn.Linear(7, 7) + + def forward(self, x): + return x + self.linear(x).tanh() + + +class TinyDraft(nn.Module): + _no_split_modules = ["TinyBlock"] + + def __init__(self): + super().__init__() + self.layers = nn.Sequential(TinyBlock(), TinyBlock()) + self.head = nn.Linear(7, 5) + # FSDP2's dim-0 sharding leaves an empty local shard on rank one. + self.scale = nn.Parameter(torch.ones(1)) + + def forward(self, x): + return self.layers(x) + + def project(self, x): + return self.head(x) * self.scale + + +class TinyComposite(nn.Module): + def __init__(self): + super().__init__() + self.draft_model = TinyDraft() + self.embed_tokens = nn.Embedding(11, 7).requires_grad_(False) + self.lm_head = nn.Linear(7, 5, bias=False).requires_grad_(False) + + def forward(self, x): + x = x + self.embed_tokens(torch.arange(x.shape[0], device=x.device) % 3) + hidden = self.draft_model(x) + # Real DFlash-family objectives also call heads outside draft.forward. + return self.draft_model.project(hidden) + self.lm_head(hidden) + + +def _optimizer(model, *, offload=False): + return BF16Optimizer( + model, + lr=1e-2, + weight_decay=0.01, + total_steps=8, + warmup_ratio=0, + max_grad_norm=0.1, + offload_master=offload, + ) + + +def _assert_plain_tensors(value): + if isinstance(value, torch.Tensor): + assert not isinstance(value, DTensor), "checkpoint leaked a DTensor" + elif isinstance(value, dict): + for item in value.values(): + _assert_plain_tensors(item) + elif isinstance(value, (list, tuple)): + for item in value: + _assert_plain_tensors(item) + + +def _numeric_worker(rank, world_size, port, workdir): + from tests.test_runtime import _fixtures as fx + + fx.init_rank_distributed(rank, world_size, port=str(port)) + try: + for dtype in (torch.float32, torch.bfloat16): + for sharding in ("SHARD_GRAD_OP", "FULL_SHARD", "NO_SHARD"): + for offload in (False, True): + torch.manual_seed(21) + template = TinyComposite().cuda().to(dtype) + reference_model = copy.deepcopy(template) + reference = create_training_backend( + "fsdp", ParallelConfig(), optimizer_factory=_optimizer + ) + reference.prepare_model( + reference_model, + wrap=False, + optimizer_target=reference_model.draft_model, + ) + pc = ParallelConfig.from_distributed( + sharding_strategy=sharding, param_dtype=dtype + ) + factory = lambda m: _optimizer(m, offload=offload) + backends = [] + for name in ("fsdp", "fsdp2"): + model = copy.deepcopy(template) + backend = create_training_backend( + name, pc, optimizer_factory=factory + ) + backend.prepare_model(model, optimizer_target=model.draft_model) + if name == "fsdp2" and sharding != "NO_SHARD": + assert isinstance(model.draft_model.scale, DTensor) + assert not isinstance(model.lm_head.weight, DTensor) + assert not isinstance(model.embed_tokens.weight, DTensor) + assert all( + not isinstance(master, DTensor) + for master in backend.optimizer.fp32_params + ) + backends.append(backend) + + def update(backend, *, full_batch=False, step=0): + for micro in range(2): + x = torch.arange( + world_size * 21, device="cuda", dtype=torch.float32 + ).reshape(world_size * 3, 7) + x = (x / 31 + micro * 0.2 + step * 0.3).to(dtype) + if not full_batch: + x = x[rank * 3 : (rank + 1) * 3] + loss = backend.module(x).float().square().mean() / 2 + backend.backward(loss, is_boundary=micro == 1) + backend.scale_gradients(torch.tensor(0.75, device="cuda")) + return backend.step( + loss_denominator=torch.tensor(11.0, device="cuda") + ) + + expected_norm = update(reference, full_batch=True) + tolerance = 0.03 if dtype == torch.bfloat16 else 2e-5 + for backend in backends: + actual_norm = update(backend) + assert actual_norm > 0.1, "the clipping gate must be exercised" + torch.testing.assert_close( + actual_norm, expected_norm, rtol=tolerance, atol=tolerance + ) + state = backend.state_dict() + _assert_plain_tensors(state) + if rank == 0: + for key, actual in state["model"].items(): + if sharding != "NO_SHARD" and key.startswith( + "draft_model." + ): + assert actual.device.type == "cpu" + expected = reference_model.state_dict()[key].cpu() + torch.testing.assert_close( + actual.cpu(), + expected, + rtol=tolerance, + atol=2e-3 if dtype == torch.bfloat16 else 2e-6, + ) + torch.save( + state["model"], os.path.join(workdir, "model.pt") + ) + else: + # FSDP1 may retain ignored, replicated target tables + # in nonzero-rank state dicts; no draft may leak. + assert not any( + key.startswith("draft_model.") for key in state["model"] + ), (backend.name, list(state["model"])) + del state["model"] + rank_path = os.path.join(workdir, f"rank{rank}.pt") + torch.save(state, rank_path) + dist.barrier() + restored = torch.load( + rank_path, map_location="cpu", weights_only=False + ) + restored["model"] = torch.load( + os.path.join(workdir, "model.pt"), + map_location="cpu", + weights_only=True, + ) + resumed_model = copy.deepcopy(template) + # Also exercise a CPU-master placement change on resume. + resumed = create_training_backend( + backend.name, + pc, + optimizer_factory=lambda m: _optimizer( + m, offload=not offload + ), + ) + resumed.prepare_model( + resumed_model, optimizer_target=resumed_model.draft_model + ) + resumed.load_state_dict(restored) + torch.testing.assert_close( + torch.get_rng_state(), restored["rng"]["torch"] + ) + torch.testing.assert_close( + update(resumed, step=1), + update(backend, step=1), + rtol=2e-5, + atol=2e-6, + ) + actual = resumed.state_dict()["model"] + expected = backend.state_dict()["model"] + for key in actual: + torch.testing.assert_close( + actual[key], expected[key], rtol=0, atol=2e-6 + ) + for a, b in zip( + resumed.optimizer.fp32_params, backend.optimizer.fp32_params + ): + torch.testing.assert_close( + a.cpu(), b.cpu(), rtol=2e-5, atol=2e-7 + ) + assert ( + resumed.optimizer.scheduler.last_epoch + == backend.optimizer.scheduler.last_epoch + ) + dist.barrier() + finally: + from specforge.distributed import destroy_distributed + + destroy_distributed(abort=sys.exc_info()[0] is not None) + + +class TestBackendSelection(unittest.TestCase): + def test_default_and_explicit_sharding(self): + from specforge.config.schema import TrainingConfig + + self.assertEqual(TrainingConfig().backend, "fsdp") + self.assertEqual(TrainingConfig(backend="fsdp2").backend, "fsdp2") + with self.assertRaises(ValueError): + TrainingConfig(backend="unknown") + with mock.patch.dict(os.environ, {"FSDP_SHARDING": "NO_SHARD"}): + self.assertEqual( + ParallelConfig.from_distributed().sharding_strategy, "NO_SHARD" + ) + self.assertEqual( + ParallelConfig.from_distributed( + sharding_strategy="FULL_SHARD" + ).sharding_strategy, + "FULL_SHARD", + ) + + def test_checkpoint_backend_mismatch_fails_before_loading_weights(self): + def build(name): + model = nn.Linear(7, 5) + backend = create_training_backend( + name, ParallelConfig(), optimizer_factory=_optimizer + ) + backend.prepare_model(model, wrap=False) + return backend + + for source, target in (("fsdp", "fsdp2"), ("fsdp2", "fsdp")): + state = build(source).state_dict() + backend = build(target) + before = copy.deepcopy(backend.module.state_dict()) + with self.assertRaisesRegex(ValueError, "optimizer checkpoint backend"): + backend.load_state_dict(state) + for name, param in backend.module.state_dict().items(): + torch.testing.assert_close(param, before[name], rtol=0, atol=0) + + legacy = build("fsdp").state_dict() + del legacy["metadata"] + with mock.patch("specforge.optimizer.print_on_rank0"): + build("fsdp").load_state_dict(legacy) + with self.assertRaisesRegex(ValueError, "optimizer checkpoint backend"): + build("fsdp2").load_state_dict(legacy) + + def test_unknown_backend_rejected(self): + with self.assertRaisesRegex(ValueError, "unsupported training backend"): + create_training_backend("unknown", ParallelConfig()) + + +@unittest.skipUnless(torch.cuda.device_count() >= 2, "requires two CUDA devices") +class TestFSDP2Distributed(unittest.TestCase): + def test_dense_equivalence_and_resume(self): + from tests.utils import get_available_port + + with tempfile.TemporaryDirectory(prefix="fsdp2_backend_") as workdir: + torch.multiprocessing.spawn( + _numeric_worker, + args=(2, get_available_port(), workdir), + nprocs=2, + join=True, + ) + + +if __name__ == "__main__": + unittest.main(verbosity=2) diff --git a/tests/test_runtime/test_offline_launch_fsdp.py b/tests/test_runtime/test_offline_launch_fsdp.py index 5bb22bc00..f5e0000b8 100644 --- a/tests/test_runtime/test_offline_launch_fsdp.py +++ b/tests/test_runtime/test_offline_launch_fsdp.py @@ -24,13 +24,22 @@ @unittest.skipUnless(CUDA, "launcher FSDP path requires CUDA") class TestOfflineLaunchFSDP(unittest.TestCase): def test_fsdp_in_forward_path_and_optimizer_step_semantics(self): + self._check_backend("fsdp") + + def test_fsdp2_trains_through_the_same_launcher(self): + self._check_backend("fsdp2") + + def _check_backend(self, training_backend): torch.manual_seed(0) from tests.test_runtime import _fixtures as fx fx.build_single_rank_distributed(port="29566") + from torch.distributed.fsdp import FSDPModule from torch.distributed.fsdp import FullyShardedDataParallel as FSDP + expected_wrapper = FSDPModule if training_backend == "fsdp2" else FSDP + from specforge.launch import build_offline_runtime from specforge.optimizer import BF16Optimizer @@ -54,6 +63,7 @@ def optimizer_factory(draft_module): draft_model=eagle3_model, target_head=target_head, optimizer_factory=optimizer_factory, + training_backend=training_backend, run_id="launch", output_dir=os.path.join(workdir, "out"), ttt_length=TTT, @@ -67,7 +77,7 @@ def optimizer_factory(draft_module): # Issue 1: the strategy runs forward through the FSDP-wrapped module module = trainer.core.strategy.trainable_module() self.assertIsInstance( - module, FSDP, "strategy must hold the FSDP-wrapped module" + module, expected_wrapper, "strategy must hold the FSDP-wrapped module" ) self.assertIsNotNone(trainer.core.backend.optimizer) From 11841f77c980cbc38308a3eb8d54e934c8036e53 Mon Sep 17 00:00:00 2001 From: "Ethan (Yusheng) Su" Date: Fri, 2 Oct 2026 10:57:20 +0900 Subject: [PATCH 02/11] Fix FSDP2 parameter retention and root resharding --- docs/sections/basic_usage/training.md | 5 +- specforge/training/fsdp2.py | 12 +- tests/test_runtime/test_fsdp2_backend.py | 167 +++++++++++++++++++++++ 3 files changed, 182 insertions(+), 2 deletions(-) diff --git a/docs/sections/basic_usage/training.md b/docs/sections/basic_usage/training.md index 3931a8920..8221f9266 100644 --- a/docs/sections/basic_usage/training.md +++ b/docs/sections/basic_usage/training.md @@ -498,7 +498,10 @@ training: Both backends retain BF16 compute, FP32 optimizer masters, gradient accumulation, and `training.optimizer_cpu_offload`. `SHARD_GRAD_OP` keeps parameters gathered -between forward and backward; `FULL_SHARD` reshards them after forward. Both +through backward and between gradient-accumulation micro-steps, resharding at +the optimizer boundary. This trades higher live memory between micro-steps for +fewer all-gathers. `FULL_SHARD` reshards child blocks after forward but keeps the +root gathered for backward; all units reshard after each backward. Both backends use DDP for `NO_SHARD`. Configured sharding takes precedence over the legacy `FSDP_SHARDING` environment fallback used by direct Python builders. diff --git a/specforge/training/fsdp2.py b/specforge/training/fsdp2.py index 2adf1daba..8616134e6 100644 --- a/specforge/training/fsdp2.py +++ b/specforge/training/fsdp2.py @@ -34,7 +34,6 @@ def _shard_model(self, model, block_classes, ignored_frozen_modules): setattr(module, name, buffer.float()) kwargs = dict( mesh=mesh, - reshard_after_forward=pc.sharding_strategy == "FULL_SHARD", ignored_params=ignored_params, ) # Bottom-up application preserves the FSDP1 block/root boundaries. @@ -44,6 +43,7 @@ def _shard_model(self, model, block_classes, ignored_frozen_modules): if module is not model and type(module) in block_classes: fully_shard( module, + reshard_after_forward=pc.sharding_strategy == "FULL_SHARD", mp_policy=MixedPrecisionPolicy( param_dtype=pc.param_dtype, cast_forward_inputs=False ), @@ -51,6 +51,9 @@ def _shard_model(self, model, block_classes, ignored_frozen_modules): ) fully_shard( model, + # Like FSDP1, reuse the root's full parameters in backward even + # under FULL_SHARD. Child blocks still reshard after forward. + reshard_after_forward=False, mp_policy=MixedPrecisionPolicy(param_dtype=pc.param_dtype), **kwargs, ) @@ -60,10 +63,17 @@ def backward(self, loss: torch.Tensor, *, is_boundary: bool = True) -> None: if self._wrapper_kind != "fsdp2": return super().backward(loss, is_boundary=is_boundary) self.module.set_requires_gradient_sync(is_boundary) + # FSDP1 SHARD_GRAD_OP retains parameters across no_sync micro-steps. + # Retain them until the optimizer boundary to avoid re-gathering on + # every forward; FULL_SHARD still releases them after each backward. + self.module.set_reshard_after_backward( + is_boundary or self.parallel_config.sharding_strategy != "SHARD_GRAD_OP" + ) try: loss.backward() finally: self.module.set_requires_gradient_sync(True) + self.module.set_reshard_after_backward(True) def _sharded_model_state_dict(self) -> dict: from torch.distributed.checkpoint.state_dict import ( diff --git a/tests/test_runtime/test_fsdp2_backend.py b/tests/test_runtime/test_fsdp2_backend.py index e156d4067..eb5d9d8bf 100644 --- a/tests/test_runtime/test_fsdp2_backend.py +++ b/tests/test_runtime/test_fsdp2_backend.py @@ -5,6 +5,7 @@ import sys import tempfile import unittest +from contextlib import nullcontext from unittest import mock import torch @@ -223,6 +224,129 @@ def update(backend, *, full_batch=False, step=0): destroy_distributed(abort=sys.exc_info()[0] is not None) +def _collective_worker(rank, world_size, port): + from tests.test_runtime import _fixtures as fx + + fx.init_rank_distributed(rank, world_size, port=str(port)) + try: + for sharding in ("SHARD_GRAD_OP", "FULL_SHARD"): + for accumulation in (1, 4): + torch.manual_seed(42) + template = TinyComposite().cuda() + reference_model = copy.deepcopy(template) + reference = create_training_backend( + "fsdp", ParallelConfig(), optimizer_factory=_optimizer + ) + reference.prepare_model( + reference_model, + wrap=False, + optimizer_target=reference_model.draft_model, + ) + pc = ParallelConfig.from_distributed( + sharding_strategy=sharding, param_dtype=torch.float32 + ) + backends = [] + for name in ("fsdp", "fsdp2"): + model = copy.deepcopy(template) + backend = create_training_backend( + name, pc, optimizer_factory=_optimizer + ) + backend.prepare_model(model, optimizer_target=model.draft_model) + backends.append(backend) + + def accumulate(backend, step, *, full_batch=False): + for micro in range(accumulation): + x = torch.arange( + world_size * 21, device="cuda", dtype=torch.float32 + ).reshape(world_size * 3, 7) + x = x / 31 + micro * 0.1 + step * 0.2 + if not full_batch: + x = x[rank * 3 : (rank + 1) * 3] + loss = backend.module(x).square().mean() / accumulation + backend.backward(loss, is_boundary=micro == accumulation - 1) + + # Warm up lazy FSDP state, then measure two consecutive windows. + # No checkpoint/unshard operation may reset parameter residency + # between them and hide a missing boundary reshard. + for step in range(4): + accumulate(reference, step, full_batch=True) + expected_grads = { + name: param.grad.detach().clone() + for name, param in reference_model.named_parameters() + if param.requires_grad + } + expected_norm = reference.step() + counts = [] + for backend in backends: + torch.cuda.synchronize() + context = ( + torch.profiler.profile( + activities=[torch.profiler.ProfilerActivity.CPU] + ) + if step >= 2 + else nullcontext() + ) + with context as profile: + accumulate(backend, step) + torch.cuda.synchronize() + if step >= 2: + counts.append( + sum( + event.count + for event in profile.key_averages() + if event.key == "c10d::_allgather_base_" + ) + ) + if backend.name == "fsdp2": + # Gather gradients outside the measured region; + # full_tensor does not change parameter residency. + for name, param in backend.module.named_parameters(): + if param.requires_grad: + assert param.grad is not None, name + actual = param.grad.full_tensor() + torch.testing.assert_close( + actual, + expected_grads[name], + rtol=2e-5, + atol=2e-6, + ) + torch.testing.assert_close( + backend.step(), expected_norm, rtol=2e-5, atol=2e-6 + ) + if step >= 2: + blocks = len(template.draft_model.layers) + # SHARD_GRAD_OP gathers each unit once per window. + # FULL_SHARD re-gathers each child in backward, while + # the root reuses its forward buffer until backward ends. + expected = ( + blocks + 1 + if sharding == "SHARD_GRAD_OP" + else (2 * blocks + 1) * accumulation + ) + assert counts == [expected, expected], ( + sharding, + accumulation, + step, + counts, + expected, + ) + for backend in backends: + state = backend.state_dict()["model"] + if rank == 0: + for name, expected in reference_model.state_dict().items(): + torch.testing.assert_close( + state[name].cpu(), + expected.cpu(), + rtol=2e-5, + atol=2e-6, + ) + dist.barrier() + finally: + from specforge.distributed import destroy_distributed + + destroy_distributed(abort=sys.exc_info()[0] is not None) + + class TestBackendSelection(unittest.TestCase): def test_default_and_explicit_sharding(self): from specforge.config.schema import TrainingConfig @@ -271,9 +395,52 @@ def test_unknown_backend_rejected(self): with self.assertRaisesRegex(ValueError, "unsupported training backend"): create_training_backend("unknown", ParallelConfig()) + def test_backward_failure_restores_sync_and_reshard(self): + for sharding in ("SHARD_GRAD_OP", "FULL_SHARD"): + for boundary in (False, True): + with self.subTest(sharding=sharding, boundary=boundary): + backend = create_training_backend( + "fsdp2", ParallelConfig(sharding_strategy=sharding) + ) + backend._wrapper_kind = "fsdp2" + backend.module = mock.Mock( + spec=[ + "set_requires_gradient_sync", + "set_reshard_after_backward", + ] + ) + loss = mock.Mock() + loss.backward.side_effect = RuntimeError("backward failed") + with self.assertRaisesRegex(RuntimeError, "backward failed"): + backend.backward(loss, is_boundary=boundary) + loss.backward.assert_called_once_with() + self.assertEqual( + backend.module.set_requires_gradient_sync.call_args_list, + [mock.call(boundary), mock.call(True)], + ) + self.assertEqual( + backend.module.set_reshard_after_backward.call_args_list, + [ + mock.call( + boundary if sharding == "SHARD_GRAD_OP" else True + ), + mock.call(True), + ], + ) + @unittest.skipUnless(torch.cuda.device_count() >= 2, "requires two CUDA devices") class TestFSDP2Distributed(unittest.TestCase): + def test_collective_counts_and_accumulated_updates(self): + from tests.utils import get_available_port + + torch.multiprocessing.spawn( + _collective_worker, + args=(2, get_available_port()), + nprocs=2, + join=True, + ) + def test_dense_equivalence_and_resume(self): from tests.utils import get_available_port From 80fe135655e47faa67ac30d0c55e4e6e13a91af7 Mon Sep 17 00:00:00 2001 From: Kan Wu Date: Wed, 2 Sep 2026 23:11:44 +0000 Subject: [PATCH 03/11] feat: MoE FFN skeleton for DFlash-family drafters (contracts, config, seams) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Lays out specforge/modeling/draft/moe as one configurable MoE layer with swappable components, and wires every seam a target-family implementation needs — with no routing math yet: - config.py: MoEConfig, moe_preset registry, resolve_moe_config(). Architecture keys use the target checkpoints' native HF names at the draft-JSON top level; training-only knobs live under dflash_config as moe_*. - router/balance/experts/shared.py: component contracts + named registries (score functions, routers, balance controllers, experts backends, shared experts). BalanceController pins the deferred-update timing that keeps activation checkpointing correct. - layer.py: MoELayer composes gate/experts/shared_experts (official attribute names); build_ffn() is the dense/MoE switch — dense drafts get the kernel provider's MLP verbatim. - hooks.py: apply_pending_balance_updates / collect_moe_aux_loss / collect_moe_metrics over any module tree. - state_dict.py: to/from_checkpoint_state_dict boundary (module layout vs official file naming), applied in the FSDP backend, warm start, materialize_draft and both exporters. FSDP full-state-dict hooks need the module's own FQNs, so the rename cannot live in state_dict(). - init.py: WarmStartPlan / select_target_experts. - dflash.py: decoder layers build their FFN through build_ffn; _init_weights reaches MoE bare Parameters; the model forward applies pending balance updates in training. DFlash/DSpark strategies add moe/* load metrics. - DESIGN.md + customization docs; tests pin the contracts with stub components. Co-Authored-By: Claude Fable 5.1 --- .../advanced_features/customization.md | 32 ++ specforge/export/checkpoint_io.py | 7 +- specforge/export/to_hf.py | 3 +- specforge/export/to_sglang.py | 7 +- specforge/modeling/draft/dflash.py | 18 +- specforge/modeling/draft/moe/DESIGN.md | 101 ++++ specforge/modeling/draft/moe/__init__.py | 121 +++++ specforge/modeling/draft/moe/_registry.py | 49 ++ specforge/modeling/draft/moe/balance.py | 71 +++ specforge/modeling/draft/moe/config.py | 208 ++++++++ specforge/modeling/draft/moe/experts.py | 53 ++ specforge/modeling/draft/moe/hooks.py | 63 +++ specforge/modeling/draft/moe/init.py | 62 +++ specforge/modeling/draft/moe/layer.py | 92 ++++ specforge/modeling/draft/moe/router.py | 90 ++++ specforge/modeling/draft/moe/shared.py | 46 ++ specforge/modeling/draft/moe/state_dict.py | 56 +++ specforge/training/backend.py | 14 +- specforge/training/model_loading.py | 5 +- specforge/training/strategies/base.py | 12 + tests/test_modeling/test_moe.py | 464 ++++++++++++++++++ 21 files changed, 1566 insertions(+), 8 deletions(-) create mode 100644 specforge/modeling/draft/moe/DESIGN.md create mode 100644 specforge/modeling/draft/moe/__init__.py create mode 100644 specforge/modeling/draft/moe/_registry.py create mode 100644 specforge/modeling/draft/moe/balance.py create mode 100644 specforge/modeling/draft/moe/config.py create mode 100644 specforge/modeling/draft/moe/experts.py create mode 100644 specforge/modeling/draft/moe/hooks.py create mode 100644 specforge/modeling/draft/moe/init.py create mode 100644 specforge/modeling/draft/moe/layer.py create mode 100644 specforge/modeling/draft/moe/router.py create mode 100644 specforge/modeling/draft/moe/shared.py create mode 100644 specforge/modeling/draft/moe/state_dict.py create mode 100644 tests/test_modeling/test_moe.py diff --git a/docs/sections/advanced_features/customization.md b/docs/sections/advanced_features/customization.md index c7b58f870..38a5583b0 100644 --- a/docs/sections/advanced_features/customization.md +++ b/docs/sections/advanced_features/customization.md @@ -128,6 +128,38 @@ drafts currently implements the GQA/MHA layout only, so plan benchmarks accordingly. DFlash2 otherwise follows the same mode selection; its convolution and selector do not change the attention projection contract. +## MoE FFN for DFlash-family drafts + +Any DFlash-family draft (DFlash, DFlash2, DSpark) swaps its dense MLP for a +sparse MoE FFN when the draft JSON sets `n_routed_experts > 0`. The MoE is one +configurable layer (`specforge/modeling/draft/moe/`, see its `DESIGN.md`): +a `moe_preset` names a target family's routing recipe, and the architecture +keys use the target checkpoints' native HF names so they can be copied from +the target's `config.json`: + +```json +{ + "moe_preset": "", + "n_routed_experts": 64, + "num_experts_per_tok": 6, + "moe_intermediate_size": 2048, + "n_shared_experts": 1, + "dflash_config": {"moe_bias_update_rate": 0.001, "moe_dispatch": "grouped_mm"} +} +``` + +Top-level keys override the preset (for ablations: `scoring_func`, +`norm_topk_prob`, `routed_scaling_factor`, `balance`, `shared_expert_gate`, +`swiglu_limit`, ...). Training-only knobs live under `dflash_config` with an +`moe_` prefix and never change the checkpoint. Checkpoints, warm starts and +exports keep the official per-expert naming (`experts.{i}.w{1,2,3}.weight`), +so an exported drafter loads into SGLang unchanged. Dense drafts are +unaffected: with no `n_routed_experts` the kernel provider's MLP is used as-is. + +A new target family is a preset registration plus whichever components it +needs (score function, balance controller, experts backend, shared expert); +each registers by name from its own module. + ## Draft architectures Draft classes register through `@register_draft`. The key defaults to the diff --git a/specforge/export/checkpoint_io.py b/specforge/export/checkpoint_io.py index 87aad6929..97c9347aa 100644 --- a/specforge/export/checkpoint_io.py +++ b/specforge/export/checkpoint_io.py @@ -159,7 +159,12 @@ def materialize_draft( draft_config = AutoDraftModelConfig.from_file(draft_config_path) model = AutoDraftModel.from_config(draft_config, torch_dtype=torch.bfloat16) - missing, unexpected = model.load_state_dict(state["draft_state_dict"], strict=False) + from specforge.modeling.draft.moe import from_checkpoint_state_dict + + # Files use the official naming; modules may use a native MoE layout. + missing, unexpected = model.load_state_dict( + from_checkpoint_state_dict(state["draft_state_dict"]), strict=False + ) if unexpected: raise ValueError( f"checkpoint carries weights the {type(model).__name__} architecture " diff --git a/specforge/export/to_hf.py b/specforge/export/to_hf.py index 928408b14..d224e5c0c 100644 --- a/specforge/export/to_hf.py +++ b/specforge/export/to_hf.py @@ -31,6 +31,7 @@ materialize_draft, resolve_training_state, ) +from specforge.modeling.draft.moe import to_checkpoint_state_dict def _load_embedding_tensor(source: str, key: str) -> torch.Tensor: @@ -87,7 +88,7 @@ def export_to_hf( model = materialize_draft( state, draft_config_path, vocab_mapping_path=vocab_mapping_path ) - full_state = dict(model.state_dict()) + full_state = dict(to_checkpoint_state_dict(model.state_dict())) owns_embedding = hasattr(model, "embed_tokens") if owns_embedding and "embed_tokens.weight" not in state["draft_state_dict"]: if not embedding_source: diff --git a/specforge/export/to_sglang.py b/specforge/export/to_sglang.py index b08b856ec..88a4f2449 100644 --- a/specforge/export/to_sglang.py +++ b/specforge/export/to_sglang.py @@ -28,6 +28,7 @@ materialize_draft, resolve_training_state, ) +from specforge.modeling.draft.moe import to_checkpoint_state_dict #: per-architecture trainer-key -> serving-key renames ({} = identity). WEIGHT_MAPS: Dict[str, Dict[str, str]] = { @@ -82,7 +83,11 @@ def export_to_sglang( weight_map = WEIGHT_MAPS.get(type(model).__name__, {}) # the model's state dict includes any refreshed t2d/d2t buffers; drop the # embeddings exactly as the trainer-side checkpoint filter does. - full = {k: v for k, v in model.state_dict().items() if "embed" not in k.lower()} + full = { + k: v + for k, v in to_checkpoint_state_dict(model.state_dict()).items() + if "embed" not in k.lower() + } model.save_pretrained(output_dir, state_dict=_serving_state(full, weight_map)) apply_legacy_rope_scaling(output_dir) return output_dir diff --git a/specforge/modeling/draft/dflash.py b/specforge/modeling/draft/dflash.py index ea50ec122..68fe9f6ce 100644 --- a/specforge/modeling/draft/dflash.py +++ b/specforge/modeling/draft/dflash.py @@ -21,6 +21,7 @@ from typing_extensions import Tuple, Unpack from .dflash_kernels import DEFAULT_DFLASH_KERNELS, DFlashKernels +from .moe import MoELayer, apply_pending_balance_updates, build_ffn from .flex_attention_backend import flex_attention_backend from .registry import register_draft @@ -558,7 +559,9 @@ def __init__( layer_idx=layer_idx, kernels=kernels, ) - self.mlp = kernels.make_mlp(config) + # Dense MLP from the kernel provider, or an MoELayer when the draft + # JSON sets n_routed_experts > 0 (see modeling/draft/moe). + self.mlp = build_ffn(config, dense=kernels.make_mlp) self.input_layernorm = kernels.make_rms_norm( config.hidden_size, config.rms_norm_eps ) @@ -733,6 +736,14 @@ def _build_decoder_layer( return self.decoder_layer_class(config, layer_idx, kernels) + def _init_weights(self, module: nn.Module) -> None: + if isinstance(module, MoELayer): + # MoE components hold bare Parameters (gate weight, stacked + # experts) that the inherited Qwen3 init never visits. + module.reset_parameters(std=self.config.initializer_range) + return + super()._init_weights(module) + def _init_draft_head(self, config, dflash_config: dict) -> None: del config, dflash_config @@ -796,6 +807,11 @@ def forward( **kwargs, ) -> CausalLMOutputWithPast: hidden_states = noise_embedding + if self.training: + # Consume the PREVIOUS forward's routing statistics before any + # routing this step, outside every activation-checkpoint region + # (see modeling/draft/moe/balance.py for why the timing matters). + apply_pending_balance_updates(self) target_hidden = self.hidden_norm(self.fc(target_hidden)) position_embeddings = self.rotary_emb(hidden_states, position_ids) for layer_type, layer in zip(self.layer_types, self.layers): diff --git a/specforge/modeling/draft/moe/DESIGN.md b/specforge/modeling/draft/moe/DESIGN.md new file mode 100644 index 000000000..8242d4cc4 --- /dev/null +++ b/specforge/modeling/draft/moe/DESIGN.md @@ -0,0 +1,101 @@ +# MoE FFN Design (`specforge.modeling.draft.moe`) + +Design note for the sparse-MoE FFN that any DFlash-family draft (DFlash, +DFlash2, DSpark) can opt into. The training plane's picture is in +[`../../../training/DESIGN.md`](../../../training/DESIGN.md). + +## Responsibility + +Owns the FFN of a decoder layer when the draft JSON sets +`n_routed_experts > 0`, and nothing else: attention, heads, losses and the +trainer loop are unchanged. The package is **one configurable layer**, not one +block per target family. This is the Megatron-Core split (one `MoELayer`, +router/balancing/experts/dispatcher/shared-expert as orthogonal components +picked by config), chosen over the transformers pattern of a copied +`XxxSparseMoeBlock` per model because the family differences are a handful of +small functions while the expensive parts are shared: + +| differs per target family | shared by every family | +| -------------------------------- | ----------------------------------------- | +| score function (softmax, sigmoid, sqrtsoftplus) | token dispatch (sorted segments, grouped GEMM) | +| balancing policy (aux loss, aux-loss-free bias, none) | deferred balance-update timing vs activation checkpointing | +| combine-weight renorm and scale | FSDP-friendly stacked expert layout | +| shared-expert gate (none, sigmoid) | checkpoint naming boundary | +| SwiGLU clamp | load metrics, warm-start plans | + +## Why match the target's MoE + +The drafter's MoE should be the *target's* MoE, expressed as a preset: + +- **Warm start.** Same expert shape means the draft experts can be seeded from + a subset of the target's, which a dense drafter cannot do. +- **Serving.** The drafter runs inside SGLang's draft model; matching the + target's routing reuses its fused MoE kernels and weight naming. +- **Latency.** At small batch, top-k of narrow experts reads about the same + bytes as the dense MLP, so MoE buys parameters at roughly constant step + cost. That is the hypothesis the dense-vs-MoE ablation tests. + +## Layout + +``` +config.py MoEConfig + preset registry. Architecture keys use the target + checkpoints' native HF names at the draft JSON top level; + training-only knobs live under dflash_config as moe_*. +router.py Router contract (x -> RoutingResult) + score-function registry. +balance.py BalanceController contract; owns selection-bias buffers, stashes + counts in forward, applies updates from the model forward. +experts.py RoutedExperts contract (weights + dispatch), MoEConfig.dispatch knob. +shared.py SharedExpert contract, gate variant via MoEConfig.shared_expert_gate. +layer.py MoELayer = gate + experts + shared_experts; build_ffn() is the + dense/MoE switch used by the DFlash decoder layer. +hooks.py apply_pending_balance_updates / collect_moe_aux_loss / + collect_moe_metrics over any module tree. +state_dict.py to/from_checkpoint_state_dict: module layout <-> official names. +init.py WarmStartPlan: which target experts seed which draft experts. +``` + +Implementations (a target-family preset, its score function, controller, +experts backend, shared expert, converter) live in their own module and +register into these registries at import time. This package holds contracts +and composition only. + +## Contracts that matter + +**Attribute names are the checkpoint contract.** `MoELayer` exposes `gate`, +`experts`, `shared_experts` (the DeepSeek-family names SGLang loads). A +component whose native parameter layout differs from the official file naming +registers a `state_dict` converter pair; both directions are idempotent and +no-ops on dense models. + +**Naming is converted at the boundary, not in `state_dict()`.** FSDP's +full-state-dict hooks index the gathered dict by the module's own parameter +FQNs, so a rename inside `state_dict()` breaks under `use_orig_params`. Every +file read/write goes through `to_checkpoint_state_dict` / +`from_checkpoint_state_dict`: `FSDPTrainingBackend` save/load, warm start, +`materialize_draft`, and the HF and SGLang exporters. + +**Balance updates are deferred.** `BalanceController.observe` only stashes +(overwrite, never accumulate) so an activation-checkpoint recompute leaves +identical state. `apply_pending_update` runs from the *model* forward before +any routing, outside checkpoint regions, and may run collectives. Mutating +selection state inside a layer forward would make the recompute route +differently and raise `CheckpointError`. + +**Metrics ride the existing scalar channel.** `collect_moe_metrics` yields +`moe/load_max_ratio`, `moe/load_min_ratio`, `moe/experts_unused_frac` plus +controller metrics; the DFlash/DSpark strategies add them to `StepOutput.metrics`, +and the trainer DP-averages and logs them like any other scalar. + +**Aux losses are collected, not yet consumed.** `collect_moe_aux_loss` sums +scaled layer losses; wiring it into an objective is done with the first preset +whose balancing policy emits one (aux-loss-free policies do not). + +## Extension points + +- New target family: `register_moe_preset("", scoring_func=..., ...)` + plus any missing component registrations. The draft JSON then sets + `moe_preset` and the per-run sizes. +- New dispatch: a `RoutedExperts` subclass under `register_experts_backend`, + or a new `MoEConfig.dispatch` value handled inside an existing backend. +- Ablation knobs: any `ARCHITECTURE_KEYS` entry at the draft JSON top level + overrides its preset default; `dflash_config.moe_*` overrides training knobs. diff --git a/specforge/modeling/draft/moe/__init__.py b/specforge/modeling/draft/moe/__init__.py new file mode 100644 index 000000000..a102c5efd --- /dev/null +++ b/specforge/modeling/draft/moe/__init__.py @@ -0,0 +1,121 @@ +# coding=utf-8 +# Copyright 2024 The SpecForge 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 +"""Sparse MoE FFN for DFlash-family drafts. + +One configurable layer, not one block per target family. A draft JSON selects a +``moe_preset`` (the routing recipe of a target family) and may override any +architecture knob for ablations; the expensive parts (dispatch, checkpoint +naming, FSDP layout, deferred balance updates) are shared. See ``DESIGN.md``. + +Modules: + +- :mod:`.config` ``MoEConfig`` + the preset registry; resolves a draft JSON +- :mod:`.router` ``Router`` contract (scores -> top-k) + score functions +- :mod:`.balance` ``BalanceController`` contract: balancing as a policy +- :mod:`.experts` ``RoutedExperts`` contract: expert weights + dispatch +- :mod:`.shared` ``SharedExpert`` contract +- :mod:`.layer` ``MoELayer`` composition; ``build_ffn`` is the dense/MoE switch +- :mod:`.hooks` model-level plumbing: balance updates, aux loss, metrics +- :mod:`.state_dict` module layout <-> official checkpoint naming boundary +- :mod:`.init` warm-start plans from a target model's experts + +Implementations register into the registries from their own modules; this +package defines contracts and holds no routing math itself. +""" + +from .balance import ( + BALANCE_CONTROLLERS, + BalanceController, + build_balance_controller, + register_balance_controller, +) +from .config import ( + MOE_PRESETS, + MoEConfig, + available_moe_presets, + is_moe_config, + register_moe_preset, + resolve_moe_config, +) +from .experts import ( + EXPERTS_BACKENDS, + RoutedExperts, + build_routed_experts, + register_experts_backend, +) +from .hooks import ( + apply_pending_balance_updates, + collect_moe_aux_loss, + collect_moe_metrics, + iter_moe_layers, +) +from .init import WarmStartPlan, plan_warm_start, select_target_experts +from .layer import MoELayer, build_ffn +from .router import ( + ROUTERS, + SCORE_FUNCTIONS, + Router, + RoutingResult, + build_router, + get_score_function, + register_router, + register_score_function, +) +from .shared import ( + SHARED_EXPERTS, + SharedExpert, + build_shared_expert, + register_shared_expert, +) +from .state_dict import ( + from_checkpoint_state_dict, + register_state_dict_converter, + to_checkpoint_state_dict, +) + +__all__ = [ + "BALANCE_CONTROLLERS", + "BalanceController", + "EXPERTS_BACKENDS", + "MOE_PRESETS", + "MoEConfig", + "MoELayer", + "ROUTERS", + "RoutedExperts", + "Router", + "RoutingResult", + "SCORE_FUNCTIONS", + "SHARED_EXPERTS", + "SharedExpert", + "WarmStartPlan", + "apply_pending_balance_updates", + "available_moe_presets", + "build_balance_controller", + "build_ffn", + "build_routed_experts", + "build_router", + "build_shared_expert", + "collect_moe_aux_loss", + "collect_moe_metrics", + "from_checkpoint_state_dict", + "get_score_function", + "is_moe_config", + "iter_moe_layers", + "plan_warm_start", + "register_balance_controller", + "register_experts_backend", + "register_moe_preset", + "register_router", + "register_score_function", + "register_shared_expert", + "register_state_dict_converter", + "resolve_moe_config", + "select_target_experts", + "to_checkpoint_state_dict", +] diff --git a/specforge/modeling/draft/moe/_registry.py b/specforge/modeling/draft/moe/_registry.py new file mode 100644 index 000000000..e1ef85e83 --- /dev/null +++ b/specforge/modeling/draft/moe/_registry.py @@ -0,0 +1,49 @@ +# coding=utf-8 +"""Tiny named-registry helper shared by the MoE component registries.""" + +from __future__ import annotations + +from typing import Callable, Dict, Generic, Optional, TypeVar + +T = TypeVar("T") + + +class Registry(Generic[T]): + """Name -> implementation map with a decorator form and helpful errors.""" + + def __init__(self, kind: str) -> None: + self.kind = kind + self._items: Dict[str, T] = {} + + def register(self, name: str, item: Optional[T] = None) -> Callable[[T], T] | T: + def _register(obj: T) -> T: + if name in self._items and self._items[name] is not obj: + raise ValueError(f"{self.kind} {name!r} is already registered") + self._items[name] = obj + return obj + + return _register if item is None else _register(item) + + def unregister(self, name: str) -> None: + self._items.pop(name, None) + + def get(self, name: str) -> T: + try: + return self._items[name] + except KeyError: + available = ", ".join(sorted(self._items)) or "" + raise KeyError( + f"unknown {self.kind} {name!r}; available: {available}" + ) from None + + def names(self) -> list[str]: + return sorted(self._items) + + def __contains__(self, name: object) -> bool: + return name in self._items + + def __getitem__(self, name: str) -> T: + return self.get(name) + + def __len__(self) -> int: + return len(self._items) diff --git a/specforge/modeling/draft/moe/balance.py b/specforge/modeling/draft/moe/balance.py new file mode 100644 index 000000000..5fbc962a6 --- /dev/null +++ b/specforge/modeling/draft/moe/balance.py @@ -0,0 +1,71 @@ +# coding=utf-8 +"""Load balancing as a swappable policy. + +The controller sees routing outcomes and may (a) shift scores used for expert +*selection* (never the combine weights), (b) emit an auxiliary loss, and (c) +report load metrics. It is an ``nn.Module`` so implementations can own +buffers that travel with the checkpoint (e.g. a selection bias). + +Two timing rules every implementation must respect: + +- :meth:`observe` is called from the layer forward and must only *stash* + (overwrite, never accumulate): an activation-checkpoint recompute re-runs the + forward and must leave identical state behind. +- :meth:`apply_pending_update` is called by the *model* before the next + forward, outside any checkpoint region, and may mutate selection state and + run collectives. Mutating selection state inside the forward would make the + recompute route differently and break checkpointing. +""" + +from __future__ import annotations + +from typing import Dict, Optional, Type, Union + +import torch +from torch import nn + +from ._registry import Registry +from .config import MoEConfig + +MetricValue = Union[torch.Tensor, float] + + +class BalanceController(nn.Module): + """No balancing (registered as ``"none"``); the base for real policies.""" + + def __init__(self, cfg: MoEConfig, n_experts: int) -> None: + super().__init__() + self.cfg = cfg + self.n_experts = n_experts + + def adjust_selection_scores(self, scores: torch.Tensor) -> torch.Tensor: + """Scores used to pick experts; combine weights still use the raw ones.""" + return scores + + def observe(self, counts: torch.Tensor) -> None: + """Stash this forward's per-expert token counts (training only).""" + + def apply_pending_update(self) -> None: + """Consume the stash; called by the model outside checkpoint regions.""" + + def aux_loss(self) -> Optional[torch.Tensor]: + """Scaled auxiliary loss for the last forward, or ``None``.""" + return None + + def metrics(self) -> Dict[str, MetricValue]: + """Scalar diagnostics (rank-local; the trainer DP-averages them).""" + return {} + + +BALANCE_CONTROLLERS: Registry[Type[BalanceController]] = Registry( + "MoE balance controller" +) +BALANCE_CONTROLLERS.register("none", BalanceController) + + +def register_balance_controller(name: str): + return BALANCE_CONTROLLERS.register(name) + + +def build_balance_controller(cfg: MoEConfig, n_experts: int) -> BalanceController: + return BALANCE_CONTROLLERS.get(cfg.balance)(cfg, n_experts) diff --git a/specforge/modeling/draft/moe/config.py b/specforge/modeling/draft/moe/config.py new file mode 100644 index 000000000..d2a55a802 --- /dev/null +++ b/specforge/modeling/draft/moe/config.py @@ -0,0 +1,208 @@ +# coding=utf-8 +"""MoE architecture config: the knobs a draft JSON can set, and their presets. + +Two kinds of keys, deliberately kept apart: + +- **Architecture keys** live at the top level of the draft JSON under the + target checkpoints' native HF names (``n_routed_experts``, + ``moe_intermediate_size``, ``scoring_func``, ...). They determine which + weights a checkpoint carries and how serving must route, so a draft JSON can + be assembled by copying them from the target's ``config.json``. A + ``moe_preset`` supplies the defaults for one target family; explicit keys + override the preset for ablations. +- **Training-only keys** live under ``dflash_config`` with an ``moe_`` prefix + (``moe_bias_update_rate``, ``moe_aux_loss_coeff``, ``moe_dispatch``). They + never change the checkpoint and are invisible to serving. +""" + +from __future__ import annotations + +from dataclasses import dataclass, fields +from typing import Any, Dict, Mapping, Optional + +from ._registry import Registry + +_MISSING = object() + +#: Draft-JSON top-level keys that describe the MoE architecture. +ARCHITECTURE_KEYS = ( + "n_routed_experts", + "num_experts_per_tok", + "moe_intermediate_size", + "n_shared_experts", + "shared_expert_intermediate_size", + "scoring_func", + "norm_topk_prob", + "routed_scaling_factor", + "n_group", + "topk_group", + "swiglu_limit", + "router", + "balance", + "shared_expert", + "shared_expert_gate", + "experts_backend", +) + +#: ``dflash_config`` keys (training-only) -> MoEConfig field. +TRAINING_KEYS = { + "moe_bias_update_rate": "bias_update_rate", + "moe_aux_loss_coeff": "aux_loss_coeff", + "moe_dispatch": "dispatch", +} + + +@dataclass(frozen=True) +class MoEConfig: + """Resolved MoE configuration for one draft model (all layers share it).""" + + preset: str + n_routed_experts: int + num_experts_per_tok: int + moe_intermediate_size: int + n_shared_experts: int = 1 + shared_expert_intermediate_size: Optional[int] = None + # Routing recipe. ``scoring_func`` names a registered score function; + # ``router``/``balance`` name registered component implementations. + scoring_func: str = "softmax" + norm_topk_prob: bool = True + routed_scaling_factor: float = 1.0 + n_group: int = 1 + topk_group: int = 1 + router: str = "topk" + balance: str = "none" + # Expert MLPs. + swiglu_limit: float = 0.0 + experts_backend: str = "grouped" + shared_expert: str = "swiglu" + shared_expert_gate: str = "none" + # Training-only (never part of the checkpoint). + bias_update_rate: float = 0.0 + aux_loss_coeff: float = 0.0 + dispatch: str = "sorted_loop" + + def __post_init__(self) -> None: + if self.n_routed_experts <= 0: + raise ValueError("n_routed_experts must be positive for an MoE FFN") + if not 0 < self.num_experts_per_tok <= self.n_routed_experts: + raise ValueError( + "num_experts_per_tok must be in [1, n_routed_experts], got " + f"{self.num_experts_per_tok} with {self.n_routed_experts} experts" + ) + if self.moe_intermediate_size <= 0: + raise ValueError("moe_intermediate_size must be positive") + if self.n_shared_experts not in (0, 1): + raise ValueError( + "n_shared_experts must be 0 or 1 (one shared expert of " + "shared_expert_intermediate_size width), got " + f"{self.n_shared_experts}" + ) + if self.shared_expert_intermediate_size is None: + object.__setattr__( + self, "shared_expert_intermediate_size", self.moe_intermediate_size + ) + if self.shared_expert_intermediate_size <= 0: + raise ValueError("shared_expert_intermediate_size must be positive") + if self.n_group <= 0 or self.n_routed_experts % self.n_group: + raise ValueError( + f"n_group={self.n_group} must divide n_routed_experts=" + f"{self.n_routed_experts}" + ) + if not 0 < self.topk_group <= self.n_group: + raise ValueError( + f"topk_group={self.topk_group} must be in [1, n_group={self.n_group}]" + ) + if self.swiglu_limit < 0: + raise ValueError("swiglu_limit must be >= 0 (0 disables the clamp)") + if self.bias_update_rate < 0 or self.aux_loss_coeff < 0: + raise ValueError("bias_update_rate and aux_loss_coeff must be >= 0") + + @property + def group_limited(self) -> bool: + """Whether routing restricts top-k to ``topk_group`` of ``n_group``.""" + return self.topk_group < self.n_group + + def as_dict(self) -> Dict[str, Any]: + return {f.name: getattr(self, f.name) for f in fields(self)} + + +#: preset name -> architecture defaults (a subset of ARCHITECTURE_KEYS). +MOE_PRESETS: Registry[Dict[str, Any]] = Registry("MoE preset") + +_FIELD_NAMES = {f.name for f in fields(MoEConfig)} + + +def register_moe_preset(name: str, **defaults: Any) -> Dict[str, Any]: + """Register the architecture defaults of one target family. + + A preset may set any ``MoEConfig`` field except the training-only ones and + the per-run sizes (``n_routed_experts``, ``num_experts_per_tok``, + ``moe_intermediate_size``), which the draft JSON must state explicitly. + """ + forbidden = set(TRAINING_KEYS.values()) | { + "preset", + "n_routed_experts", + "num_experts_per_tok", + "moe_intermediate_size", + } + bad = sorted(set(defaults) - _FIELD_NAMES) + if bad: + raise ValueError(f"preset {name!r} sets unknown MoEConfig fields: {bad}") + bad = sorted(set(defaults) & forbidden) + if bad: + raise ValueError(f"preset {name!r} may not set per-run/training fields: {bad}") + MOE_PRESETS.register(name, dict(defaults)) + return defaults + + +def available_moe_presets() -> list[str]: + return MOE_PRESETS.names() + + +def _get(config: Any, key: str, default: Any = _MISSING) -> Any: + if isinstance(config, Mapping): + return config.get(key, default) + return getattr(config, key, default) + + +def is_moe_config(config: Any) -> bool: + """True when a draft config asks for an MoE FFN (``n_routed_experts > 0``).""" + value = _get(config, "n_routed_experts", 0) + return int(value or 0) > 0 + + +def resolve_moe_config(config: Any) -> Optional[MoEConfig]: + """Resolve a draft config (HF ``PretrainedConfig`` or dict) to ``MoEConfig``. + + Returns ``None`` for dense drafts. For MoE drafts, ``moe_preset`` is + required: it names the target family's routing recipe and is the only way + the architecture keys get validated defaults. + """ + if not is_moe_config(config): + return None + preset = _get(config, "moe_preset", None) + if not preset: + raise ValueError( + "n_routed_experts > 0 requires moe_preset in the draft config; " + f"available presets: {available_moe_presets() or ''}" + ) + values: Dict[str, Any] = dict(MOE_PRESETS.get(preset)) + for key in ARCHITECTURE_KEYS: + explicit = _get(config, key, _MISSING) + if explicit is not _MISSING and explicit is not None: + values[key] = explicit + dflash_config = _get(config, "dflash_config", None) or {} + unknown = sorted( + key + for key in dflash_config + if key.startswith("moe_") and key not in TRAINING_KEYS + ) + if unknown: + raise ValueError( + f"unknown MoE training keys in dflash_config: {unknown}; " + f"known: {sorted(TRAINING_KEYS)}" + ) + for json_key, field_name in TRAINING_KEYS.items(): + if json_key in dflash_config: + values[field_name] = dflash_config[json_key] + return MoEConfig(preset=preset, **values) diff --git a/specforge/modeling/draft/moe/experts.py b/specforge/modeling/draft/moe/experts.py new file mode 100644 index 000000000..c5d0f3f94 --- /dev/null +++ b/specforge/modeling/draft/moe/experts.py @@ -0,0 +1,53 @@ +# coding=utf-8 +"""Routed experts contract: the expert weights and how tokens reach them. + +An implementation owns the parameters of all ``n_routed_experts`` experts and +turns a :class:`RoutingResult` into the combined routed output. Weight layout +is the implementation's choice (per-expert modules, stacked ``[E, out, in]`` +tensors, ...); it must register a :mod:`.state_dict` converter if its native +layout differs from the official checkpoint naming +(``experts.{i}.w{1,2,3}.weight``). ``MoEConfig.dispatch`` is the +implementation's execution knob (e.g. sorted-segment loop vs grouped GEMM). +""" + +from __future__ import annotations + +import abc +from typing import Type + +import torch +from torch import nn + +from ._registry import Registry +from .config import MoEConfig +from .router import RoutingResult + + +class RoutedExperts(nn.Module, abc.ABC): + def __init__(self, cfg: MoEConfig, hidden_size: int) -> None: + super().__init__() + self.cfg = cfg + self.hidden_size = hidden_size + self.n_experts = cfg.n_routed_experts + self.intermediate_size = cfg.moe_intermediate_size + + @abc.abstractmethod + def forward(self, x: torch.Tensor, routing: RoutingResult) -> torch.Tensor: + """``x`` is ``[T, hidden]``; return the combined routed output ``[T, hidden]`` + in ``x.dtype``.""" + + @abc.abstractmethod + def reset_parameters(self, std: float) -> None: + """Initialize expert weights with the draft's ``initializer_range`` so an + MoE FFN starts from the same distribution as the dense MLP it replaces.""" + + +EXPERTS_BACKENDS: Registry[Type[RoutedExperts]] = Registry("MoE experts backend") + + +def register_experts_backend(name: str): + return EXPERTS_BACKENDS.register(name) + + +def build_routed_experts(cfg: MoEConfig, hidden_size: int) -> RoutedExperts: + return EXPERTS_BACKENDS.get(cfg.experts_backend)(cfg, hidden_size) diff --git a/specforge/modeling/draft/moe/hooks.py b/specforge/modeling/draft/moe/hooks.py new file mode 100644 index 000000000..2867d370a --- /dev/null +++ b/specforge/modeling/draft/moe/hooks.py @@ -0,0 +1,63 @@ +# coding=utf-8 +"""Model-level plumbing for MoE layers. + +A draft model with MoE FFNs needs three things from its trainer loop, all +expressed here as functions over any ``nn.Module`` tree so DFlash, DFlash2 and +DSpark share them: + +- :func:`apply_pending_balance_updates` at the top of the model forward (in + training), outside activation-checkpoint regions; +- :func:`collect_moe_aux_loss` to add to the objective when a balance policy + emits one; +- :func:`collect_moe_metrics` for per-step diagnostics (``moe/...``). +""" + +from __future__ import annotations + +from typing import Dict, Iterator, Optional + +import torch +from torch import nn + +from .balance import MetricValue +from .layer import MoELayer + + +def iter_moe_layers(module: nn.Module) -> Iterator[MoELayer]: + for sub in module.modules(): + if isinstance(sub, MoELayer): + yield sub + + +def apply_pending_balance_updates(module: nn.Module) -> None: + for layer in iter_moe_layers(module): + layer.apply_pending_balance_update() + + +def collect_moe_aux_loss(module: nn.Module) -> Optional[torch.Tensor]: + """Sum of the layers' (already scaled) auxiliary losses, or ``None``.""" + total: Optional[torch.Tensor] = None + for layer in iter_moe_layers(module): + loss = layer.aux_loss() + if loss is None: + continue + total = loss if total is None else total + loss + return total + + +def collect_moe_metrics( + module: nn.Module, prefix: str = "moe/" +) -> Dict[str, MetricValue]: + """Layer-averaged scalar diagnostics; ``{}`` for dense models.""" + sums: Dict[str, MetricValue] = {} + n = 0 + for layer in iter_moe_layers(module): + n += 1 + for key, value in layer.metrics().items(): + sums[key] = value if key not in sums else sums[key] + value + if n == 0: + return {} + return { + f"{prefix}{key}": (value / n if isinstance(value, torch.Tensor) else value / n) + for key, value in sums.items() + } diff --git a/specforge/modeling/draft/moe/init.py b/specforge/modeling/draft/moe/init.py new file mode 100644 index 000000000..4f1b57e4b --- /dev/null +++ b/specforge/modeling/draft/moe/init.py @@ -0,0 +1,62 @@ +# coding=utf-8 +"""Warm-start plans: which target experts seed which draft experts. + +A drafter whose MoE matches the target's expert shape can inherit expert +weights instead of training them from scratch. This module holds the +implementation-independent part: choosing the mapping. Applying a plan needs +the target checkpoint's (possibly quantized) weights and the experts backend's +native layout, and is provided alongside each target-family preset. +""" + +from __future__ import annotations + +from dataclasses import dataclass +from typing import Tuple + +from .config import MoEConfig + + +@dataclass(frozen=True) +class WarmStartPlan: + """``target_expert_ids[i]`` is the target expert that seeds draft expert ``i``.""" + + target_expert_ids: Tuple[int, ...] + copy_shared_expert: bool = True + copy_gate_rows: bool = True + + @property + def n_draft_experts(self) -> int: + return len(self.target_expert_ids) + + +def select_target_experts( + n_target: int, n_draft: int, strategy: str = "strided" +) -> Tuple[int, ...]: + """Pick ``n_draft`` distinct target experts. + + ``"strided"`` spreads picks evenly over the target's expert ids (the + default: no assumption about which experts matter for the draft's data); + ``"first"`` takes the leading ids. + """ + if n_draft <= 0 or n_target <= 0: + raise ValueError("expert counts must be positive") + if n_draft > n_target: + raise ValueError( + f"cannot seed {n_draft} draft experts from {n_target} target experts" + ) + if strategy == "first": + return tuple(range(n_draft)) + if strategy == "strided": + return tuple((i * n_target) // n_draft for i in range(n_draft)) + raise ValueError(f"unknown warm-start selection strategy {strategy!r}") + + +def plan_warm_start( + cfg: MoEConfig, n_target_experts: int, strategy: str = "strided" +) -> WarmStartPlan: + return WarmStartPlan( + target_expert_ids=select_target_experts( + n_target_experts, cfg.n_routed_experts, strategy + ), + copy_shared_expert=bool(cfg.n_shared_experts), + ) diff --git a/specforge/modeling/draft/moe/layer.py b/specforge/modeling/draft/moe/layer.py new file mode 100644 index 000000000..bc0954299 --- /dev/null +++ b/specforge/modeling/draft/moe/layer.py @@ -0,0 +1,92 @@ +# coding=utf-8 +"""``MoELayer``: the FFN that composes router, experts and shared expert. + +Attribute names follow the official DeepSeek-style checkpoint layout +(``gate``, ``experts``, ``shared_experts``) so that per-implementation +converters only need to handle their own internals. +""" + +from __future__ import annotations + +from typing import Callable, Dict, Optional + +import torch +from torch import nn + +from .balance import MetricValue, build_balance_controller +from .config import MoEConfig, resolve_moe_config +from .experts import build_routed_experts +from .router import RoutingResult, build_router +from .shared import build_shared_expert + + +class MoELayer(nn.Module): + """Routed FFN: ``y = experts(x, gate(x)) + shared_experts(x)``.""" + + def __init__(self, cfg: MoEConfig, hidden_size: int) -> None: + super().__init__() + self.cfg = cfg + self.hidden_size = hidden_size + balance = build_balance_controller(cfg, cfg.n_routed_experts) + self.gate = build_router(cfg, hidden_size, balance) + self.experts = build_routed_experts(cfg, hidden_size) + self.shared_experts: Optional[nn.Module] = ( + build_shared_expert(cfg, hidden_size) if cfg.n_shared_experts else None + ) + # Detached per-expert counts of the last training forward, for metrics. + self.last_counts: Optional[torch.Tensor] = None + + @property + def balance(self): + return self.gate.balance + + def forward(self, x: torch.Tensor) -> torch.Tensor: + shape = x.shape + x = x.reshape(-1, self.hidden_size) + routing: RoutingResult = self.gate(x) + if self.training: + self.last_counts = routing.counts.detach() + self.balance.observe(routing.counts) + y = self.experts(x, routing) + if self.shared_experts is not None: + y = y + self.shared_experts(x) + return y.view(shape) + + # -- model-level hooks (see hooks.py) --------------------------------- + def apply_pending_balance_update(self) -> None: + self.balance.apply_pending_update() + + def aux_loss(self) -> Optional[torch.Tensor]: + return self.balance.aux_loss() + + def metrics(self) -> Dict[str, MetricValue]: + out: Dict[str, MetricValue] = {} + counts = self.last_counts + if counts is not None and counts.numel(): + load = counts.float() + mean = load.mean().clamp_min(1e-9) + out["load_max_ratio"] = load.max() / mean + out["load_min_ratio"] = load.min() / mean + out["experts_unused_frac"] = (load == 0).float().mean() + out.update(self.balance.metrics()) + return out + + def reset_parameters(self, std: float) -> None: + """Initialize bare Parameters the HF ``_init_weights`` pass cannot see.""" + self.gate.reset_parameters(std) + self.experts.reset_parameters(std) + if self.shared_experts is not None: + self.shared_experts.reset_parameters(std) + + +def build_ffn(config, dense: Callable[[object], nn.Module]) -> nn.Module: + """The dense/MoE switch for a decoder layer's FFN. + + ``config`` is the draft's HF config; ``dense`` builds the dense MLP (the + kernel provider's factory) and is used verbatim when the config is dense, + so dense drafts are byte-for-byte unaffected by this package. + """ + moe_cfg = resolve_moe_config(config) + if moe_cfg is None: + return dense(config) + return MoELayer(moe_cfg, int(config.hidden_size)) diff --git a/specforge/modeling/draft/moe/router.py b/specforge/modeling/draft/moe/router.py new file mode 100644 index 000000000..b495a383f --- /dev/null +++ b/specforge/modeling/draft/moe/router.py @@ -0,0 +1,90 @@ +# coding=utf-8 +"""Router contract: hidden states -> per-token expert choices + combine weights. + +A router owns the gate projection and composes a :class:`BalanceController` +(``self.balance``) that may shift scores for *selection only*. Concrete routers +register by name (``MoEConfig.router``); score functions register separately +(``MoEConfig.scoring_func``) so one top-k router serves softmax, sigmoid and +sqrtsoftplus families. +""" + +from __future__ import annotations + +import abc +from dataclasses import dataclass +from typing import Callable, Type + +import torch +from torch import nn + +from ._registry import Registry +from .balance import BalanceController +from .config import MoEConfig + + +@dataclass +class RoutingResult: + """Routing decision for a flat batch of ``T`` tokens. + + ``weights`` are the final combine weights (normalized and scaled as the + recipe dictates) in fp32; ``indices`` the chosen experts; ``counts`` the + per-expert token counts on device (no host sync), which dispatch and + balancing both consume. + """ + + weights: torch.Tensor # [T, k] fp32 + indices: torch.Tensor # [T, k] long + counts: torch.Tensor # [E] long + + @property + def topk(self) -> int: + return int(self.indices.shape[-1]) + + +#: name -> f(logits [T, E] fp32) -> scores [T, E] fp32 +SCORE_FUNCTIONS: Registry[Callable[[torch.Tensor], torch.Tensor]] = Registry( + "MoE score function" +) + + +def register_score_function(name: str): + return SCORE_FUNCTIONS.register(name) + + +def get_score_function(name: str) -> Callable[[torch.Tensor], torch.Tensor]: + return SCORE_FUNCTIONS.get(name) + + +class Router(nn.Module, abc.ABC): + """Base router. Subclasses implement :meth:`forward` and :meth:`reset_parameters`.""" + + def __init__( + self, cfg: MoEConfig, hidden_size: int, balance: BalanceController + ) -> None: + super().__init__() + self.cfg = cfg + self.hidden_size = hidden_size + self.n_experts = cfg.n_routed_experts + self.topk = cfg.num_experts_per_tok + self.balance = balance + + @abc.abstractmethod + def forward(self, x: torch.Tensor) -> RoutingResult: + """Route a flat ``[T, hidden]`` batch.""" + + @abc.abstractmethod + def reset_parameters(self, std: float) -> None: + """Initialize the gate weights (called from the model's ``_init_weights``).""" + + +ROUTERS: Registry[Type[Router]] = Registry("MoE router") + + +def register_router(name: str): + return ROUTERS.register(name) + + +def build_router( + cfg: MoEConfig, hidden_size: int, balance: BalanceController +) -> Router: + return ROUTERS.get(cfg.router)(cfg, hidden_size, balance) diff --git a/specforge/modeling/draft/moe/shared.py b/specforge/modeling/draft/moe/shared.py new file mode 100644 index 000000000..b655b577a --- /dev/null +++ b/specforge/modeling/draft/moe/shared.py @@ -0,0 +1,46 @@ +# coding=utf-8 +"""Shared expert contract: an always-on FFN added to the routed output. + +Families differ in the gate on it (DeepSeek: none; Qwen: a per-token sigmoid +gate), which is ``MoEConfig.shared_expert_gate``, and in its width +(``shared_expert_intermediate_size``). Implementations register by name +(``MoEConfig.shared_expert``). Checkpoint naming follows the official +``shared_experts.w{1,2,3}.weight`` layout unless a converter says otherwise. +""" + +from __future__ import annotations + +import abc +from typing import Type + +import torch +from torch import nn + +from ._registry import Registry +from .config import MoEConfig + + +class SharedExpert(nn.Module, abc.ABC): + def __init__(self, cfg: MoEConfig, hidden_size: int) -> None: + super().__init__() + self.cfg = cfg + self.hidden_size = hidden_size + self.intermediate_size = cfg.shared_expert_intermediate_size + + @abc.abstractmethod + def forward(self, x: torch.Tensor) -> torch.Tensor: + """``x`` is ``[T, hidden]``; return ``[T, hidden]`` in ``x.dtype``.""" + + def reset_parameters(self, std: float) -> None: + """Hook for bare Parameters; ``nn.Linear`` children are covered by HF init.""" + + +SHARED_EXPERTS: Registry[Type[SharedExpert]] = Registry("MoE shared expert") + + +def register_shared_expert(name: str): + return SHARED_EXPERTS.register(name) + + +def build_shared_expert(cfg: MoEConfig, hidden_size: int) -> SharedExpert: + return SHARED_EXPERTS.get(cfg.shared_expert)(cfg, hidden_size) diff --git a/specforge/modeling/draft/moe/state_dict.py b/specforge/modeling/draft/moe/state_dict.py new file mode 100644 index 000000000..020a43694 --- /dev/null +++ b/specforge/modeling/draft/moe/state_dict.py @@ -0,0 +1,56 @@ +# coding=utf-8 +"""Checkpoint-naming boundary for MoE modules. + +Checkpoint FILES (trainer checkpoints, warm-start sources, exports) use the +official per-expert naming so SGLang and HF loaders read them unchanged. +Modules may use a different native layout (e.g. stacked ``[E, out, in]`` +expert tensors, which FSDP ``use_orig_params`` and grouped GEMMs want). + +FSDP's full-state-dict hooks index the gathered dict by the module's own +parameter FQNs, so the rename cannot live inside ``state_dict()``; it lives at +the save/load boundary instead. Every place that reads a model's state for a +file, or loads a file into a model, goes through :func:`to_checkpoint_state_dict` +/ :func:`from_checkpoint_state_dict`: the training backend, warm start, and +the HF/SGLang exporters. Implementations with a native layout register a +converter pair; both directions must be no-ops on dicts already in the other +form, and on dense models. +""" + +from __future__ import annotations + +from typing import Callable, Dict, List, Tuple + +Converter = Callable[[Dict[str, object]], Dict[str, object]] + +_CONVERTERS: List[Tuple[str, Converter, Converter]] = [] + + +def register_state_dict_converter( + name: str, *, to_checkpoint: Converter, from_checkpoint: Converter +) -> None: + for existing, _, _ in _CONVERTERS: + if existing == name: + raise ValueError(f"state-dict converter {name!r} is already registered") + _CONVERTERS.append((name, to_checkpoint, from_checkpoint)) + + +def unregister_state_dict_converter(name: str) -> None: + _CONVERTERS[:] = [entry for entry in _CONVERTERS if entry[0] != name] + + +def registered_state_dict_converters() -> List[str]: + return [name for name, _, _ in _CONVERTERS] + + +def to_checkpoint_state_dict(state: Dict[str, object]) -> Dict[str, object]: + """Module-native naming -> official checkpoint naming (identity for dense).""" + for _, to_checkpoint, _ in _CONVERTERS: + state = to_checkpoint(state) + return state + + +def from_checkpoint_state_dict(state: Dict[str, object]) -> Dict[str, object]: + """Official checkpoint naming -> module-native naming (identity for dense).""" + for _, _, from_checkpoint in reversed(_CONVERTERS): + state = from_checkpoint(state) + return state diff --git a/specforge/training/backend.py b/specforge/training/backend.py index f4d673bb4..e55f2fad3 100644 --- a/specforge/training/backend.py +++ b/specforge/training/backend.py @@ -411,22 +411,30 @@ def _full_state_ctx(self, state_dict_config=None): ) def _module_state_dict(self) -> dict: + # Checkpoint FILES use the official parameter naming; modules may use + # a different native layout (MoE experts). Convert at this boundary: + # FSDP's full-state-dict hooks need the module's own FQNs. + from specforge.modeling.draft.moe import to_checkpoint_state_dict + if self._wrapper_kind == "ddp": if dist.is_initialized() and dist.get_rank() != 0: return {} - return self.module.module.state_dict() + return to_checkpoint_state_dict(self.module.module.state_dict()) if self._wrapper_kind != "fsdp": - return self.module.state_dict() + return to_checkpoint_state_dict(self.module.state_dict()) from torch.distributed.fsdp import FullStateDictConfig # gather to rank0 CPU only — materializing the full model on every # rank's GPU is wasted memory when only rank0 writes it. cfg = FullStateDictConfig(offload_to_cpu=True, rank0_only=True) with self._full_state_ctx(cfg): - return self.module.state_dict() + return to_checkpoint_state_dict(self.module.state_dict()) def _load_module_state_dict(self, model_state: dict) -> None: + from specforge.modeling.draft.moe import from_checkpoint_state_dict + # every rank loads the full state dict read from the shared file. + model_state = from_checkpoint_state_dict(model_state) if self._wrapper_kind == "ddp": self.module.module.load_state_dict(model_state) return diff --git a/specforge/training/model_loading.py b/specforge/training/model_loading.py index a4f8d8616..6ad8ee6d9 100644 --- a/specforge/training/model_loading.py +++ b/specforge/training/model_loading.py @@ -426,7 +426,10 @@ def warm_start_draft_model( if not state: raise ValueError(f"warm-start checkpoint {source!r} contains no draft weights") try: - result = model.load_state_dict(state, strict=False) + # Files use the official naming; modules may use a native MoE layout. + from specforge.modeling.draft.moe import from_checkpoint_state_dict + + result = model.load_state_dict(from_checkpoint_state_dict(state), strict=False) except RuntimeError as exc: raise ValueError( f"warm-start checkpoint {source!r} has incompatible draft tensor " diff --git a/specforge/training/strategies/base.py b/specforge/training/strategies/base.py index f592ccdc0..175545b54 100644 --- a/specforge/training/strategies/base.py +++ b/specforge/training/strategies/base.py @@ -56,6 +56,16 @@ class StepContext: collect_detailed_metrics: bool = True +def _moe_metrics(model_wrapper: nn.Module) -> Dict[str, Any]: + """``moe/...`` load diagnostics of the wrapped draft; ``{}`` for dense drafts.""" + from specforge.modeling.draft.moe import collect_moe_metrics + + draft_model = getattr(model_wrapper, "draft_model", None) + if draft_model is None: + return {} + return collect_moe_metrics(draft_model) + + def linear_lambda_base( global_step: int, total_steps: int, @@ -546,6 +556,7 @@ def forward_loss( metrics["accuracy_denom"] = model_metrics["accuracy_denom"] if "selector_loss_alpha" in model_metrics: metrics["selector_loss_alpha"] = model_metrics["selector_loss_alpha"] + metrics.update(_moe_metrics(self.dflash_model)) return StepOutput( loss=loss, metrics=metrics, @@ -616,6 +627,7 @@ def forward_loss( ): if name in model_metrics: metrics[name] = model_metrics[name] + metrics.update(_moe_metrics(self.dspark_model)) return StepOutput( loss=loss, metrics=metrics, diff --git a/tests/test_modeling/test_moe.py b/tests/test_modeling/test_moe.py new file mode 100644 index 000000000..01ef7dede --- /dev/null +++ b/tests/test_modeling/test_moe.py @@ -0,0 +1,464 @@ +# coding=utf-8 +"""MoE FFN skeleton: config resolution, component contracts, composition, and +the model/trainer/checkpoint seams — exercised with stub components so the +contracts are pinned independently of any target-family implementation.""" + +import unittest +from types import SimpleNamespace + +import torch +from torch import nn +from transformers import Qwen3Config +from transformers.models.qwen3.modeling_qwen3 import Qwen3MLP + +from specforge.modeling.draft.dflash import DFlashDraftModel +from specforge.modeling.draft.moe import ( + BALANCE_CONTROLLERS, + EXPERTS_BACKENDS, + MOE_PRESETS, + ROUTERS, + SCORE_FUNCTIONS, + SHARED_EXPERTS, + BalanceController, + MoEConfig, + MoELayer, + RoutedExperts, + Router, + RoutingResult, + SharedExpert, + apply_pending_balance_updates, + build_ffn, + collect_moe_aux_loss, + collect_moe_metrics, + from_checkpoint_state_dict, + iter_moe_layers, + plan_warm_start, + register_balance_controller, + register_experts_backend, + register_moe_preset, + register_router, + register_score_function, + register_shared_expert, + register_state_dict_converter, + resolve_moe_config, + select_target_experts, + to_checkpoint_state_dict, +) +from specforge.modeling.draft.moe.state_dict import unregister_state_dict_converter +from specforge.training.strategies.base import _moe_metrics + +PRESET = "_test_family" + + +# --- stub components: a minimal but real top-k reference for the contracts --- + + +class _StubBalance(BalanceController): + def __init__(self, cfg, n_experts): + super().__init__(cfg, n_experts) + self.register_buffer("bias", torch.zeros(n_experts)) + self.observed = [] + self.applied = 0 + + def adjust_selection_scores(self, scores): + return scores + self.bias + + def observe(self, counts): + self.observed.append(counts.detach().clone()) + + def apply_pending_update(self): + self.applied += 1 + + def aux_loss(self): + if self.cfg.aux_loss_coeff <= 0: + return None + return torch.tensor(self.cfg.aux_loss_coeff) + + def metrics(self): + return {"stub_metric": 1.0} + + +class _StubRouter(Router): + def __init__(self, cfg, hidden_size, balance): + super().__init__(cfg, hidden_size, balance) + self.weight = nn.Parameter(torch.empty(self.n_experts, hidden_size)) + self.score_fn = SCORE_FUNCTIONS.get(cfg.scoring_func) + + def reset_parameters(self, std): + nn.init.normal_(self.weight, std=std) + + def forward(self, x): + scores = self.score_fn(x.float() @ self.weight.float().t()) + indices = self.balance.adjust_selection_scores(scores).topk(self.topk).indices + weights = scores.gather(1, indices) + if self.cfg.norm_topk_prob: + weights = weights / weights.sum(-1, keepdim=True) + weights = weights * self.cfg.routed_scaling_factor + counts = torch.zeros( + self.n_experts, dtype=torch.long, device=x.device + ).scatter_add_(0, indices.flatten(), torch.ones_like(indices.flatten())) + return RoutingResult(weights=weights, indices=indices, counts=counts) + + +class _StubExperts(RoutedExperts): + def __init__(self, cfg, hidden_size): + super().__init__(cfg, hidden_size) + self.w = nn.Parameter(torch.empty(self.n_experts, hidden_size, hidden_size)) + + def reset_parameters(self, std): + nn.init.normal_(self.w, std=std) + + def forward(self, x, routing): + # dense gather: fine for tiny tests, pins the [T, k] -> [T, H] contract + per_choice = torch.einsum( + "td,tkdo->tko", x.float(), self.w[routing.indices].float() + ) + return (routing.weights.unsqueeze(-1) * per_choice).sum(1).to(x.dtype) + + +class _StubShared(SharedExpert): + def __init__(self, cfg, hidden_size): + super().__init__(cfg, hidden_size) + self.proj = nn.Linear(hidden_size, hidden_size, bias=False) + + def forward(self, x): + return self.proj(x) + + +def setUpModule(): + register_score_function("_test_softplus")(torch.nn.functional.softplus) + register_router("_test_router")(_StubRouter) + register_balance_controller("_test_balance")(_StubBalance) + register_experts_backend("_test_experts")(_StubExperts) + register_shared_expert("_test_shared")(_StubShared) + register_moe_preset( + PRESET, + scoring_func="_test_softplus", + router="_test_router", + balance="_test_balance", + experts_backend="_test_experts", + shared_expert="_test_shared", + routed_scaling_factor=1.5, + ) + + +def tearDownModule(): + SCORE_FUNCTIONS.unregister("_test_softplus") + ROUTERS.unregister("_test_router") + BALANCE_CONTROLLERS.unregister("_test_balance") + EXPERTS_BACKENDS.unregister("_test_experts") + SHARED_EXPERTS.unregister("_test_shared") + MOE_PRESETS.unregister(PRESET) + + +def _moe_json(**overrides): + payload = dict( + hidden_size=16, + moe_preset=PRESET, + n_routed_experts=8, + num_experts_per_tok=2, + moe_intermediate_size=8, + dflash_config={}, + ) + payload.update(overrides) + return payload + + +def _dflash_config(moe=True, **overrides): + fields = dict( + architectures=["DFlashDraftModel"], + block_size=2, + hidden_size=16, + intermediate_size=32, + num_attention_heads=2, + num_key_value_heads=1, + num_hidden_layers=2, + num_target_layers=6, + head_dim=8, + max_position_embeddings=64, + vocab_size=32, + layer_types=["full_attention", "full_attention"], + dflash_config={"attention_mode": "gqa"}, + ) + if moe: + fields.update(_moe_json(hidden_size=16)) + fields["dflash_config"] = { + "attention_mode": "gqa", + "moe_bias_update_rate": 0.005, + } + fields.update(overrides) + config = Qwen3Config(**fields) + config._attn_implementation = "sdpa" + return config + + +class TestMoEConfig(unittest.TestCase): + def test_dense_config_resolves_to_none(self): + self.assertIsNone(resolve_moe_config({"hidden_size": 16})) + self.assertIsNone(resolve_moe_config(SimpleNamespace(n_routed_experts=0))) + + def test_preset_is_required_for_moe(self): + with self.assertRaisesRegex(ValueError, "moe_preset"): + resolve_moe_config(_moe_json(moe_preset=None)) + with self.assertRaisesRegex(KeyError, "available"): + resolve_moe_config(_moe_json(moe_preset="no-such-family")) + + def test_preset_defaults_and_explicit_overrides(self): + cfg = resolve_moe_config(_moe_json()) + self.assertEqual(cfg.preset, PRESET) + self.assertEqual(cfg.scoring_func, "_test_softplus") + self.assertEqual(cfg.routed_scaling_factor, 1.5) + self.assertEqual( + cfg.shared_expert_intermediate_size, 8 + ) # defaults to moe width + cfg = resolve_moe_config( + _moe_json(routed_scaling_factor=2.0, shared_expert_intermediate_size=4) + ) + self.assertEqual(cfg.routed_scaling_factor, 2.0) + self.assertEqual(cfg.shared_expert_intermediate_size, 4) + # Works on attribute-style configs (HF PretrainedConfig) as well. + self.assertEqual(resolve_moe_config(_dflash_config()).n_routed_experts, 8) + + def test_training_knobs_come_from_dflash_config(self): + cfg = resolve_moe_config( + _moe_json( + dflash_config={ + "moe_bias_update_rate": 0.01, + "moe_dispatch": "grouped_mm", + } + ) + ) + self.assertEqual(cfg.bias_update_rate, 0.01) + self.assertEqual(cfg.dispatch, "grouped_mm") + with self.assertRaisesRegex(ValueError, "unknown MoE training keys"): + resolve_moe_config(_moe_json(dflash_config={"moe_bias_udpate_rate": 1})) + + def test_validation(self): + with self.assertRaisesRegex(ValueError, "num_experts_per_tok"): + resolve_moe_config(_moe_json(num_experts_per_tok=9)) + with self.assertRaisesRegex(ValueError, "n_shared_experts"): + resolve_moe_config(_moe_json(n_shared_experts=2)) + with self.assertRaisesRegex(ValueError, "n_group"): + resolve_moe_config(_moe_json(n_group=3)) + cfg = resolve_moe_config(_moe_json(n_group=4, topk_group=2)) + self.assertTrue(cfg.group_limited) + self.assertFalse(resolve_moe_config(_moe_json()).group_limited) + + def test_presets_cannot_set_per_run_or_training_fields(self): + with self.assertRaisesRegex(ValueError, "per-run/training"): + register_moe_preset("_bad", n_routed_experts=4) + with self.assertRaisesRegex(ValueError, "unknown MoEConfig fields"): + register_moe_preset("_bad", nonsense=1) + self.assertNotIn("_bad", MOE_PRESETS) + + def test_registry_errors_name_the_kind_and_choices(self): + with self.assertRaisesRegex(KeyError, "MoE router.*available"): + ROUTERS.get("missing") + with self.assertRaisesRegex(ValueError, "already registered"): + register_router("_test_router")(_StubRouter.__mro__[1]) + self.assertIn("none", BALANCE_CONTROLLERS) + + +class TestMoELayerComposition(unittest.TestCase): + def _layer(self, **overrides): + torch.manual_seed(0) + cfg = resolve_moe_config(_moe_json(**overrides)) + layer = MoELayer(cfg, 16) + layer.reset_parameters(std=0.02) + return layer + + def test_dense_config_uses_the_dense_factory_verbatim(self): + sentinel = nn.Identity() + self.assertIs( + build_ffn(SimpleNamespace(hidden_size=16), lambda c: sentinel), sentinel + ) + + def test_moe_config_builds_official_attribute_layout(self): + layer = build_ffn(SimpleNamespace(**_moe_json()), lambda c: nn.Identity()) + self.assertIsInstance(layer, MoELayer) + names = {name for name, _ in layer.named_children()} + self.assertEqual(names, {"gate", "experts", "shared_experts"}) + self.assertIsInstance(layer.gate.balance, _StubBalance) + self.assertIs(layer.balance, layer.gate.balance) + + def test_forward_preserves_shape_and_adds_shared_expert(self): + layer = self._layer() + x = torch.randn(2, 3, 16) + y = layer(x) + self.assertEqual(y.shape, x.shape) + routed_only = self._layer(n_shared_experts=0) + self.assertIsNone(routed_only.shared_experts) + routed_only.load_state_dict( + {k: v for k, v in layer.state_dict().items() if "shared" not in k} + ) + shared = layer.shared_experts(x.reshape(-1, 16)).view_as(x) + self.assertTrue(torch.allclose(y, routed_only(x) + shared, atol=1e-5)) + + def test_training_observes_counts_and_eval_does_not(self): + layer = self._layer().train() + x = torch.randn(5, 16) + layer(x) + self.assertEqual(len(layer.balance.observed), 1) + self.assertEqual(int(layer.balance.observed[0].sum()), 5 * 2) + self.assertTrue(torch.equal(layer.last_counts, layer.balance.observed[0])) + layer.eval() + layer(x) + self.assertEqual(len(layer.balance.observed), 1) + + def test_model_hooks_delegate_to_the_controller(self): + layer = self._layer(dflash_config={"moe_aux_loss_coeff": 0.25}).train() + layer(torch.randn(4, 16)) + layer.apply_pending_balance_update() + self.assertEqual(layer.balance.applied, 1) + self.assertAlmostEqual(float(layer.aux_loss()), 0.25) + metrics = layer.metrics() + self.assertEqual( + set(metrics), + {"load_max_ratio", "load_min_ratio", "experts_unused_frac", "stub_metric"}, + ) + self.assertGreaterEqual(float(metrics["load_max_ratio"]), 1.0) + self.assertLessEqual(float(metrics["load_min_ratio"]), 1.0) + + def test_selection_bias_changes_choice_not_weights(self): + layer = self._layer().eval() + x = torch.randn(1, 16) + before = layer.gate(x) + layer.balance.bias[:] = -1e3 + favored = int((before.indices[0, 0] + 1) % 8) + layer.balance.bias[favored] = 0.0 + after = layer.gate(x) + self.assertIn(favored, after.indices[0].tolist()) + # combine weights come from raw scores: still normalized and scaled + self.assertAlmostEqual(float(after.weights.detach().sum()), 1.5, places=5) + + +class TestHooks(unittest.TestCase): + def _tree(self, **overrides): + cfg = resolve_moe_config(_moe_json(**overrides)) + a, b = MoELayer(cfg, 16), MoELayer(cfg, 16) + for layer in (a, b): + layer.reset_parameters(std=0.02) + return nn.Sequential(nn.Linear(16, 16), a, nn.Linear(16, 16), b), (a, b) + + def test_iteration_updates_and_aggregation(self): + tree, (a, b) = self._tree(dflash_config={"moe_aux_loss_coeff": 0.5}) + self.assertEqual(list(iter_moe_layers(tree)), [a, b]) + tree.train()(torch.randn(3, 16)) + apply_pending_balance_updates(tree) + self.assertEqual((a.balance.applied, b.balance.applied), (1, 1)) + self.assertAlmostEqual(float(collect_moe_aux_loss(tree)), 1.0) + metrics = collect_moe_metrics(tree) + self.assertTrue(all(key.startswith("moe/") for key in metrics)) + self.assertAlmostEqual(float(metrics["moe/stub_metric"]), 1.0) + + def test_dense_trees_are_inert(self): + dense = nn.Sequential(nn.Linear(16, 16)) + apply_pending_balance_updates(dense) + self.assertIsNone(collect_moe_aux_loss(dense)) + self.assertEqual(collect_moe_metrics(dense), {}) + self.assertEqual(_moe_metrics(SimpleNamespace()), {}) + + def test_strategy_metrics_read_the_wrapped_draft(self): + tree, _ = self._tree() + tree.train()(torch.randn(3, 16)) + metrics = _moe_metrics(SimpleNamespace(draft_model=tree)) + self.assertIn("moe/load_max_ratio", metrics) + + +class TestStateDictBoundary(unittest.TestCase): + def test_dense_state_passes_through_unchanged(self): + state = {"a": torch.zeros(1), "layers.0.mlp.gate_proj.weight": torch.ones(1)} + for convert in (to_checkpoint_state_dict, from_checkpoint_state_dict): + out = convert(dict(state)) + self.assertEqual(set(out), set(state)) + for key in state: + self.assertIs(out[key], state[key]) + + def test_registered_converters_apply_in_both_directions(self): + def to_ckpt(state): + return {k.replace("native.", "official."): v for k, v in state.items()} + + def from_ckpt(state): + return {k.replace("official.", "native."): v for k, v in state.items()} + + register_state_dict_converter( + "_test", to_checkpoint=to_ckpt, from_checkpoint=from_ckpt + ) + try: + with self.assertRaisesRegex(ValueError, "already registered"): + register_state_dict_converter( + "_test", to_checkpoint=to_ckpt, from_checkpoint=from_ckpt + ) + official = to_checkpoint_state_dict({"native.w": 1, "other": 2}) + self.assertEqual(official, {"official.w": 1, "other": 2}) + self.assertEqual( + from_checkpoint_state_dict(official), {"native.w": 1, "other": 2} + ) + finally: + unregister_state_dict_converter("_test") + self.assertEqual(to_checkpoint_state_dict({"native.w": 1}), {"native.w": 1}) + + +class TestWarmStartPlan(unittest.TestCase): + def test_selection_strategies(self): + self.assertEqual(select_target_experts(256, 4), (0, 64, 128, 192)) + self.assertEqual(select_target_experts(8, 3, "strided"), (0, 2, 5)) + self.assertEqual(select_target_experts(8, 3, "first"), (0, 1, 2)) + self.assertEqual(len(set(select_target_experts(256, 64))), 64) + with self.assertRaises(ValueError): + select_target_experts(4, 8) + with self.assertRaises(ValueError): + select_target_experts(8, 2, "random") + + def test_plan_follows_the_moe_config(self): + cfg = resolve_moe_config(_moe_json(n_shared_experts=0)) + plan = plan_warm_start(cfg, n_target_experts=64) + self.assertEqual(plan.n_draft_experts, 8) + self.assertFalse(plan.copy_shared_expert) + self.assertTrue(plan.copy_gate_rows) + + +class TestDFlashWiring(unittest.TestCase): + def _forward(self, model): + return model( + position_ids=torch.arange(6).unsqueeze(0), + noise_embedding=torch.randn(1, 2, 16), + target_hidden=torch.randn(1, 4, 2 * 16), + ) + + def test_dense_draft_is_unchanged(self): + model = DFlashDraftModel(_dflash_config(moe=False)) + for layer in model.layers: + self.assertIsInstance(layer.mlp, Qwen3MLP) + self.assertEqual(list(iter_moe_layers(model)), []) + + def test_moe_draft_layers_init_and_apply_balance_updates(self): + model = DFlashDraftModel(_dflash_config()) + layers = list(iter_moe_layers(model)) + self.assertEqual(len(layers), 2) + for layer in layers: + self.assertIs(layer, [m for m in model.layers if m.mlp is layer][0].mlp) + self.assertEqual(layer.cfg.bias_update_rate, 0.005) + # _init_weights reached the bare Parameters (no uninitialized memory) + self.assertTrue(torch.isfinite(layer.gate.weight).all()) + self.assertGreater(float(layer.gate.weight.abs().sum()), 0.0) + self.assertTrue(torch.isfinite(layer.experts.w).all()) + model.train() + self._forward(model) + self._forward(model) + self.assertEqual([layer.balance.applied for layer in layers], [2, 2]) + self.assertEqual([len(layer.balance.observed) for layer in layers], [2, 2]) + model.eval() + self._forward(model) + self.assertEqual([layer.balance.applied for layer in layers], [2, 2]) + + def test_state_dict_names_follow_the_official_layout(self): + model = DFlashDraftModel(_dflash_config()) + keys = set(model.state_dict()) + self.assertIn("layers.0.mlp.gate.weight", keys) + self.assertIn("layers.0.mlp.shared_experts.proj.weight", keys) + self.assertTrue(any(k.startswith("layers.0.mlp.experts.") for k in keys)) + + +if __name__ == "__main__": + unittest.main(verbosity=2) From d83bc8f6d36fb43cfd1190ad11b9561a49b6ed5d Mon Sep 17 00:00:00 2001 From: Kan Wu Date: Thu, 3 Sep 2026 07:27:00 +0000 Subject: [PATCH 04/11] moe: balance controllers observe the full RoutingResult RoutingResult carries the differentiable pre-selection scores (optional), and BalanceController.observe() receives the whole result instead of bare counts, so a policy can build an auxiliary balance loss for the same forward. Co-Authored-By: Claude Fable 5.1 --- specforge/modeling/draft/moe/balance.py | 16 ++++++++++++---- specforge/modeling/draft/moe/layer.py | 2 +- specforge/modeling/draft/moe/router.py | 7 +++++-- tests/test_modeling/test_moe.py | 10 +++++++--- 4 files changed, 25 insertions(+), 10 deletions(-) diff --git a/specforge/modeling/draft/moe/balance.py b/specforge/modeling/draft/moe/balance.py index 5fbc962a6..3f45aa9d5 100644 --- a/specforge/modeling/draft/moe/balance.py +++ b/specforge/modeling/draft/moe/balance.py @@ -10,7 +10,8 @@ - :meth:`observe` is called from the layer forward and must only *stash* (overwrite, never accumulate): an activation-checkpoint recompute re-runs the - forward and must leave identical state behind. + forward and must leave identical state behind. An auxiliary loss built here + is consumed by the trainer right after the same forward. - :meth:`apply_pending_update` is called by the *model* before the next forward, outside any checkpoint region, and may mutate selection state and run collectives. Mutating selection state inside the forward would make the @@ -19,7 +20,7 @@ from __future__ import annotations -from typing import Dict, Optional, Type, Union +from typing import TYPE_CHECKING, Dict, Optional, Type, Union import torch from torch import nn @@ -27,6 +28,9 @@ from ._registry import Registry from .config import MoEConfig +if TYPE_CHECKING: # pragma: no cover - import cycle with router.py + from .router import RoutingResult + MetricValue = Union[torch.Tensor, float] @@ -42,8 +46,12 @@ def adjust_selection_scores(self, scores: torch.Tensor) -> torch.Tensor: """Scores used to pick experts; combine weights still use the raw ones.""" return scores - def observe(self, counts: torch.Tensor) -> None: - """Stash this forward's per-expert token counts (training only).""" + def observe(self, routing: "RoutingResult") -> None: + """Stash this forward's routing outcome (training only). + + ``routing.counts`` feeds load statistics; ``routing.scores`` (when the + router provides it) lets a policy build a differentiable auxiliary + loss to return from :meth:`aux_loss` for the same forward.""" def apply_pending_update(self) -> None: """Consume the stash; called by the model outside checkpoint regions.""" diff --git a/specforge/modeling/draft/moe/layer.py b/specforge/modeling/draft/moe/layer.py index bc0954299..612ca8c1b 100644 --- a/specforge/modeling/draft/moe/layer.py +++ b/specforge/modeling/draft/moe/layer.py @@ -46,7 +46,7 @@ def forward(self, x: torch.Tensor) -> torch.Tensor: routing: RoutingResult = self.gate(x) if self.training: self.last_counts = routing.counts.detach() - self.balance.observe(routing.counts) + self.balance.observe(routing) y = self.experts(x, routing) if self.shared_experts is not None: y = y + self.shared_experts(x) diff --git a/specforge/modeling/draft/moe/router.py b/specforge/modeling/draft/moe/router.py index b495a383f..757bcc441 100644 --- a/specforge/modeling/draft/moe/router.py +++ b/specforge/modeling/draft/moe/router.py @@ -12,7 +12,7 @@ import abc from dataclasses import dataclass -from typing import Callable, Type +from typing import Callable, Optional, Type import torch from torch import nn @@ -29,12 +29,15 @@ class RoutingResult: ``weights`` are the final combine weights (normalized and scaled as the recipe dictates) in fp32; ``indices`` the chosen experts; ``counts`` the per-expert token counts on device (no host sync), which dispatch and - balancing both consume. + balancing both consume. ``scores`` are the full pre-selection affinities + (differentiable, fp32) for balance policies that need a gradient signal, + e.g. an auxiliary balance loss; routers may leave it ``None``. """ weights: torch.Tensor # [T, k] fp32 indices: torch.Tensor # [T, k] long counts: torch.Tensor # [E] long + scores: Optional[torch.Tensor] = None # [T, E] fp32, differentiable @property def topk(self) -> int: diff --git a/tests/test_modeling/test_moe.py b/tests/test_modeling/test_moe.py index 01ef7dede..e8745c7bb 100644 --- a/tests/test_modeling/test_moe.py +++ b/tests/test_modeling/test_moe.py @@ -63,8 +63,9 @@ def __init__(self, cfg, n_experts): def adjust_selection_scores(self, scores): return scores + self.bias - def observe(self, counts): - self.observed.append(counts.detach().clone()) + def observe(self, routing): + self.observed.append(routing.counts.detach().clone()) + self.saw_scores = routing.scores is not None def apply_pending_update(self): self.applied += 1 @@ -97,7 +98,9 @@ def forward(self, x): counts = torch.zeros( self.n_experts, dtype=torch.long, device=x.device ).scatter_add_(0, indices.flatten(), torch.ones_like(indices.flatten())) - return RoutingResult(weights=weights, indices=indices, counts=counts) + return RoutingResult( + weights=weights, indices=indices, counts=counts, scores=scores + ) class _StubExperts(RoutedExperts): @@ -301,6 +304,7 @@ def test_training_observes_counts_and_eval_does_not(self): self.assertEqual(len(layer.balance.observed), 1) self.assertEqual(int(layer.balance.observed[0].sum()), 5 * 2) self.assertTrue(torch.equal(layer.last_counts, layer.balance.observed[0])) + self.assertTrue(layer.balance.saw_scores) layer.eval() layer(x) self.assertEqual(len(layer.balance.observed), 1) From cfc53a403e2e645171f0c8267949ab9c600c4c80 Mon Sep 17 00:00:00 2001 From: Kan Wu Date: Wed, 2 Sep 2026 23:20:42 +0000 Subject: [PATCH 05/11] feat: DeepSeek-V4 MoE FFN for DFlash-family drafters + DSV4-Flash ablation arm MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Fills the MoE skeleton with the DeepSeek-V4 recipe and ships the MoE arm of the dense-vs-MoE DSpark ablation for DeepSeek-V4-Flash. - topk_router.py: "topk" router; softmax / sigmoid / sqrtsoftplus score functions; optional group-limited selection (DeepSeek top-2 group scores). Routing math in fp32; combine weights renormalized then scaled. - noaux_tc.py: aux-loss-free balancing — fp32 selection bias (kept fp32 through module dtype casts) moved by a sign controller on all-reduced expert loads, applied from the model forward outside checkpoint regions. Stored in checkpoints as gate.bias (the DeepSeek native key SGLang maps onto e_score_correction_bias) via a state-dict converter. - grouped_experts.py: experts as stacked [E, out, in] w1/w2/w3 (FSDP tracks 3 tensors, grouped GEMMs read them directly); sorted-segment dispatch with a torch._grouped_mm path (no host sync) or the portable per-expert loop (dflash_config.moe_dispatch). Converter keeps files in the official experts.{i}.w{1,2,3}.weight naming. Experts and gate init with the draft's initializer_range, like the dense MLP they replace. - swiglu_shared.py: one ungated SwiGLU shared expert with the V4 clamp. - presets.py: "deepseek_v4" = sqrtsoftplus + noaux_tc + renorm x1.5 + one shared expert + swiglu_limit 10. - init.py: apply_warm_start() seeds a layer from a target layer's dequantized tensors (selected experts, gate rows/bias, shared expert) through the checkpoint-naming boundary. Reading/dequantizing DSV4-Flash's fp4 experts is left to the target tooling. - configs/deepseek-v4-flash-dspark-moe.json: the dense DSpark config plus moe_preset deepseek_v4, 64 routed + 1 shared experts, top-6, width 2048 (activated FFN width ~= the dense 12288). The disaggregated recipe mirrors the dense one exactly except the draft JSON and run/store names. - Tests: recipe/config resolution, router weights and group-limited routing, deferred bias updates and fp32 buffer, dense per-token reference forward, grouped-GEMM vs loop parity (CUDA), official-naming round trips at layer and model level, warm start, DFlash training/balancing integration. Co-Authored-By: Claude Fable 5.1 --- configs/deepseek-v4-flash-dspark-moe.json | 60 +++ .../deepseek-v4-flash-dspark-disaggregated.md | 24 + .../advanced_features/customization.md | 12 +- examples/configs/README.md | 5 + ...eek-v4-flash-dspark-moe-disaggregated.yaml | 97 ++++ specforge/modeling/draft/moe/DESIGN.md | 19 +- specforge/modeling/draft/moe/__init__.py | 22 +- .../modeling/draft/moe/grouped_experts.py | 154 +++++++ specforge/modeling/draft/moe/init.py | 52 ++- specforge/modeling/draft/moe/noaux_tc.py | 102 +++++ specforge/modeling/draft/moe/presets.py | 27 ++ specforge/modeling/draft/moe/swiglu_shared.py | 30 ++ specforge/modeling/draft/moe/topk_router.py | 85 ++++ tests/test_modeling/test_moe_deepseek_v4.py | 414 ++++++++++++++++++ .../test_runtime/test_package_architecture.py | 1 + 15 files changed, 1089 insertions(+), 15 deletions(-) create mode 100644 configs/deepseek-v4-flash-dspark-moe.json create mode 100644 examples/configs/online/disaggregated/external/deepseek-v4-flash-dspark-moe-disaggregated.yaml create mode 100644 specforge/modeling/draft/moe/grouped_experts.py create mode 100644 specforge/modeling/draft/moe/noaux_tc.py create mode 100644 specforge/modeling/draft/moe/presets.py create mode 100644 specforge/modeling/draft/moe/swiglu_shared.py create mode 100644 specforge/modeling/draft/moe/topk_router.py create mode 100644 tests/test_modeling/test_moe_deepseek_v4.py diff --git a/configs/deepseek-v4-flash-dspark-moe.json b/configs/deepseek-v4-flash-dspark-moe.json new file mode 100644 index 000000000..4e7f9410a --- /dev/null +++ b/configs/deepseek-v4-flash-dspark-moe.json @@ -0,0 +1,60 @@ +{ + "architectures": ["DSparkDraftModel"], + "attention_bias": false, + "attention_dropout": 0.0, + "auto_map": {"AutoModel": "dspark.DSparkDraftModel"}, + "block_size": 7, + "bos_token_id": 0, + "dflash_config": { + "attention_mode": "gqa", + "confidence_head_alpha": 1.0, + "confidence_head_with_markov": true, + "enable_confidence_head": true, + "markov_head_type": "vanilla", + "markov_rank": 256, + "mask_token_id": 128799, + "moe_bias_update_rate": 0.001, + "moe_dispatch": "grouped_mm", + "projector_type": "dspark", + "target_layer_ids": [1, 11, 21, 31, 41] + }, + "dtype": "bfloat16", + "eos_token_id": 1, + "head_dim": 128, + "hidden_act": "silu", + "hidden_size": 4096, + "initializer_range": 0.02, + "intermediate_size": 12288, + "layer_types": [ + "full_attention", + "full_attention", + "full_attention", + "full_attention", + "full_attention" + ], + "max_position_embeddings": 1048576, + "max_window_layers": 5, + "model_type": "qwen3", + "moe_intermediate_size": 2048, + "moe_preset": "deepseek_v4", + "n_routed_experts": 64, + "n_shared_experts": 1, + "num_attention_heads": 32, + "num_experts_per_tok": 6, + "num_hidden_layers": 5, + "num_key_value_heads": 8, + "num_target_layers": 43, + "pad_token_id": 1, + "rms_norm_eps": 1e-06, + "rope_parameters": { + "factor": 16.0, + "original_max_position_embeddings": 65536, + "rope_theta": 10000.0, + "rope_type": "yarn" + }, + "sliding_window": null, + "tie_word_embeddings": false, + "use_cache": true, + "use_sliding_window": false, + "vocab_size": 129280 +} diff --git a/docs/recipes/deepseek-v4-flash-dspark-disaggregated.md b/docs/recipes/deepseek-v4-flash-dspark-disaggregated.md index e62b4e013..8876e4485 100644 --- a/docs/recipes/deepseek-v4-flash-dspark-disaggregated.md +++ b/docs/recipes/deepseek-v4-flash-dspark-disaggregated.md @@ -103,6 +103,30 @@ MI355X node over the full two epochs (1,885 optimizer steps): 6.32 s per 128-sample step on average, 6.1-6.4 s at steady state (about 3.1 s waiting for capture and 3.0 s of trainer compute), 3 h 18 min end to end. +## MoE-FFN arm (ablation) + +`examples/configs/online/disaggregated/external/deepseek-v4-flash-dspark-moe-disaggregated.yaml` +is the same recipe with `configs/deepseek-v4-flash-dspark-moe.json`: the +five-layer GQA decoder keeps its attention, and each layer's dense MLP becomes +the target's MoE (`moe_preset: deepseek_v4`: sqrt-softplus scores, +aux-loss-free top-k with the sign-controlled selection bias, combine weights +renormalized and scaled by 1.5, one ungated shared expert, SwiGLU clamped at +10). Sizes are per run: 64 routed experts, top-6, width 2048, so the activated +FFN width (6 x 2048 + 2048 shared) matches the dense 12288 at ~10x the FFN +parameters. Run it against the dense recipe with identical hparams; the two +YAMLs differ only in the draft JSON and run names. The capture servers are +shared by both arms unchanged. + +Training-only knobs live under the draft JSON's `dflash_config`: +`moe_bias_update_rate` (0.001, the balancing controller's step) and +`moe_dispatch` (`grouped_mm` runs the experts as grouped GEMMs with no host +sync; `sorted_loop` is the portable fallback). The trainer logs `moe/*` load +metrics (max/min load ratios, unused-expert fraction, balancing-bias +magnitude) alongside the usual scalars. Checkpoints keep the official +per-expert naming (`layers.N.mlp.experts.{i}.w{1,2,3}.weight`, +`layers.N.mlp.gate.bias`, `layers.N.mlp.shared_experts.w{1,2,3}.weight`), so +exports load into SGLang's DeepSeek-V4 MoE unchanged. + ## Fresh attempts Delete the run's `outputs/` directory and, whenever a capture server was diff --git a/docs/sections/advanced_features/customization.md b/docs/sections/advanced_features/customization.md index 38a5583b0..a2ae5cbc1 100644 --- a/docs/sections/advanced_features/customization.md +++ b/docs/sections/advanced_features/customization.md @@ -139,7 +139,7 @@ the target's `config.json`: ```json { - "moe_preset": "", + "moe_preset": "deepseek_v4", "n_routed_experts": 64, "num_experts_per_tok": 6, "moe_intermediate_size": 2048, @@ -156,9 +156,13 @@ exports keep the official per-expert naming (`experts.{i}.w{1,2,3}.weight`), so an exported drafter loads into SGLang unchanged. Dense drafts are unaffected: with no `n_routed_experts` the kernel provider's MLP is used as-is. -A new target family is a preset registration plus whichever components it -needs (score function, balance controller, experts backend, shared expert); -each registers by name from its own module. +`deepseek_v4` is the checked-in preset (DeepSeek-V4 routing: +`sqrtsoftplus` scores, aux-loss-free `noaux_tc` balancing, combine weights +renormalized and scaled by 1.5, one ungated shared expert, SwiGLU clamp 10); +`configs/deepseek-v4-flash-dspark-moe.json` uses it. A new target family is a +preset registration plus whichever components it needs (score function, +balance controller, experts backend, shared expert); each registers by name +from its own module. ## Draft architectures diff --git a/examples/configs/README.md b/examples/configs/README.md index 905c4eb10..18b6f47a7 100644 --- a/examples/configs/README.md +++ b/examples/configs/README.md @@ -88,6 +88,11 @@ servers. Its [runbook](../../docs/recipes/deepseek-v4-flash-dspark-disaggregated.md) covers the v0.5.18 SGLang capture patch and the bundled `deepseek-v4` chat template (the checkpoint ships no Jinja template). +`deepseek-v4-flash-dspark-moe-disaggregated.yaml` is the MoE-FFN arm of the +drafter-architecture ablation: the same recipe with +`configs/deepseek-v4-flash-dspark-moe.json`, whose `moe_preset: deepseek_v4` +swaps the dense MLP for the target's routing (64 routed + 1 shared experts, +top-6, width 2048); see the runbook's MoE section. `qwen3.8-27b-dflash2-disaggregated.yaml` (external services, two nodes) and its managed-local siblings `qwen3.8-27b-dflash2-4server-dp4-disaggregated.yaml` diff --git a/examples/configs/online/disaggregated/external/deepseek-v4-flash-dspark-moe-disaggregated.yaml b/examples/configs/online/disaggregated/external/deepseek-v4-flash-dspark-moe-disaggregated.yaml new file mode 100644 index 000000000..615db8755 --- /dev/null +++ b/examples/configs/online/disaggregated/external/deepseek-v4-flash-dspark-moe-disaggregated.yaml @@ -0,0 +1,97 @@ +# MoE-FFN arm of the DSpark drafter-architecture ablation for +# DeepSeek-V4-Flash. Identical to deepseek-v4-flash-dspark-disaggregated.yaml +# (the dense arm) except the draft JSON and the run/store names, so the A/B +# diff is the FFN: 64 routed experts + 1 shared, top-6, moe_intermediate 2048 +# (activated width 6x2048 + 2048 ~= the dense 12288) with DeepSeek-V4 routing. +# If a from-scratch MoE run destabilizes at the shared hparams, lower the LR +# and clip on BOTH arms rather than on this file alone. +model: + target_model_path: deepseek-ai/DeepSeek-V4-Flash-0731 + draft_model_config: configs/deepseek-v4-flash-dspark-moe.json + target_backend: sglang + trust_remote_code: true + # DeepSeek-V4 checkpoints use the native inference weight layout. + embedding_key: embed.weight + lm_head_key: head.weight + mask_token_id: 128799 + torch_dtype: bfloat16 + sglang_mem_fraction_static: 0.85 + # data.max_length plus headroom: capture rejects inputs at exactly context length. + sglang_context_length: 8704 + sglang_max_running_requests: 8 + # Routed experts are fp4; the default MoE path cannot run them on B200. + sglang_moe_runner_backend: flashinfer_mxfp4 + +data: + train_data_path: ./cache/dataset/sharegpt_train.jsonl + max_length: 8192 + chat_template: deepseek-v4 + cache_dir: cache + build_dataset_num_proc: 64 + # Async prefetch: overlaps per-sample Mooncake feature fetches with compute. + dataloader_num_workers: 8 + +training: + strategy: dspark + num_epochs: 2 + # 4 ranks x 32 microbatches -> global batch 128. + batch_size: 1 + accumulation_steps: 32 + learning_rate: 0.0006 + lr_scheduler: constant + warmup_ratio: 0 + max_grad_norm: 1 + attention_backend: flex_attention + num_anchors: 512 + loss_decay_gamma: 4.0 + objective_chunk_blocks: 128 + dspark_ce_loss_alpha: 0.1 + dspark_l1_loss_alpha: 0.9 + dspark_confidence_head_alpha: 1.0 + save_interval: 128 + log_interval: 10 + dist_timeout: 30 + seed: 42 + prompt_seed: 1 + +tracking: + report_to: wandb + wandb_project: specforge + wandb_name: deepseek-v4-flash-dspark-moe-disaggregated + wandb_dir: outputs/deepseek-v4-flash-dspark-moe-disaggregated/wandb + +runtime: + producer_lease: 8 + producer_concurrency: 8 + # Keep two 128-sample optimizer quanta in flight (~0.4 GiB features/sample); + # the low watermark is one full quantum so complete windows always dispatch. + in_flight_high_watermark: 256 + in_flight_low_watermark: 128 + resident_high_watermark_bytes: 137438953472 + resident_low_watermark_bytes: 103079215104 + feature_store_max_resident_bytes: 171798691840 + +run_id: deepseek-v4-flash-dspark-moe-disaggregated +output_dir: outputs/deepseek-v4-flash-dspark-moe-disaggregated + +deployment: + mode: disaggregated + trainer: + nnodes: 1 + nproc_per_node: 4 + disaggregated: + control_dir: outputs/deepseek-v4-flash-dspark-moe-disaggregated/control + consumer_state_dir: outputs/deepseek-v4-flash-dspark-moe-disaggregated/consumer-state + backend: mooncake + store_id: deepseek-v4-flash-dspark-moe-disaggregated + # Two TP2 servers out-produce one TP4 server (TP prefill scaling is sublinear). + server_urls: + - http://127.0.0.1:30000 + - http://127.0.0.1:30001 + mooncake_metadata_server: http://127.0.0.1:35880/metadata + mooncake_master_server_addr: 127.0.0.1:35551 + mooncake_local_hostname: 127.0.0.1 + mooncake_protocol: tcp + client_buffer_size: 1073741824 + idle_timeout_s: 7200 + peer_wait_timeout_s: 7200 diff --git a/specforge/modeling/draft/moe/DESIGN.md b/specforge/modeling/draft/moe/DESIGN.md index 8242d4cc4..15bd32538 100644 --- a/specforge/modeling/draft/moe/DESIGN.md +++ b/specforge/modeling/draft/moe/DESIGN.md @@ -54,10 +54,21 @@ state_dict.py to/from_checkpoint_state_dict: module layout <-> official names. init.py WarmStartPlan: which target experts seed which draft experts. ``` -Implementations (a target-family preset, its score function, controller, -experts backend, shared expert, converter) live in their own module and -register into these registries at import time. This package holds contracts -and composition only. +Implementations register into these registries at import time (imported at +the bottom of `__init__.py`): + +``` +topk_router.py "topk" router; score functions softmax / sigmoid / sqrtsoftplus; + optional group-limited selection (DeepSeek top-2 group scores). +noaux_tc.py "noaux_tc" controller: fp32 selection bias + sign controller on + all-reduced loads; converter gate.balance.bias <-> gate.bias. +grouped_experts.py "grouped" experts: stacked [E, out, in] w1/w2/w3, sorted-segment + loop or torch._grouped_mm dispatch; converter experts.w1 <-> + experts.{i}.w1.weight. +swiglu_shared.py "swiglu" ungated shared expert (shared_experts.w1/w2/w3). +presets.py "deepseek_v4": sqrtsoftplus + noaux_tc + renorm x1.5 + one + ungated shared expert + SwiGLU clamp 10. +``` ## Contracts that matter diff --git a/specforge/modeling/draft/moe/__init__.py b/specforge/modeling/draft/moe/__init__.py index a102c5efd..3cd8ad691 100644 --- a/specforge/modeling/draft/moe/__init__.py +++ b/specforge/modeling/draft/moe/__init__.py @@ -25,10 +25,20 @@ - :mod:`.state_dict` module layout <-> official checkpoint naming boundary - :mod:`.init` warm-start plans from a target model's experts -Implementations register into the registries from their own modules; this -package defines contracts and holds no routing math itself. +Implementations register into the registries from their own modules +(:mod:`.topk_router`, :mod:`.noaux_tc`, :mod:`.grouped_experts`, +:mod:`.swiglu_shared`) and presets in :mod:`.presets`; the modules above +define contracts and hold no routing math themselves. """ +# Implementations and presets register at import time. +from . import ( # noqa: E402,F401 isort: skip + grouped_experts, + noaux_tc, + presets, + swiglu_shared, + topk_router, +) from .balance import ( BALANCE_CONTROLLERS, BalanceController, @@ -55,7 +65,12 @@ collect_moe_metrics, iter_moe_layers, ) -from .init import WarmStartPlan, plan_warm_start, select_target_experts +from .init import ( + WarmStartPlan, + apply_warm_start, + plan_warm_start, + select_target_experts, +) from .layer import MoELayer, build_ffn from .router import ( ROUTERS, @@ -95,6 +110,7 @@ "SharedExpert", "WarmStartPlan", "apply_pending_balance_updates", + "apply_warm_start", "available_moe_presets", "build_balance_controller", "build_ffn", diff --git a/specforge/modeling/draft/moe/grouped_experts.py b/specforge/modeling/draft/moe/grouped_experts.py new file mode 100644 index 000000000..7a511a60e --- /dev/null +++ b/specforge/modeling/draft/moe/grouped_experts.py @@ -0,0 +1,154 @@ +# coding=utf-8 +"""Routed experts as three stacked parameters with sorted-segment dispatch. + +Weights live as ``w1``/``w2``/``w3`` of shape ``[E, out, in]``: grouped GEMMs +read them directly (a per-call ``torch.stack`` of hundreds of expert weights +would allocate a transient multi-GiB tensor) and FSDP ``use_orig_params`` +tracks 3 tensors instead of ``3*E``. Checkpoint FILES keep the official +per-expert naming (``experts.{i}.w{1,2,3}.weight``) through the converter +registered below. + +Dispatch (``MoEConfig.dispatch``): + +- ``"sorted_loop"``: one stable argsort turns routing into contiguous + per-expert segments, then one small GEMM per active expert. A per-expert + ``torch.where`` loop scales launch and autograd overhead with the number of + ACTIVE experts (~2x step time once the balancer spreads load). +- ``"grouped_mm"``: the same segments through ``torch._grouped_mm`` with + on-device offsets (no host sync). Used on CUDA when available; falls back to + the loop elsewhere. Same math up to bf16 rounding. +""" + +from __future__ import annotations + +import re + +import torch +import torch.nn.functional as F +from torch import nn + +from .config import MoEConfig +from .experts import RoutedExperts, register_experts_backend +from .router import RoutingResult +from .state_dict import register_state_dict_converter + +DISPATCH_MODES = ("sorted_loop", "grouped_mm") + + +def swiglu_clamped(gate: torch.Tensor, up: torch.Tensor, limit: float) -> torch.Tensor: + """SwiGLU in fp32 with the DeepSeek-V4 activation clamp (``limit`` 0 = off).""" + gate = gate.float() + up = up.float() + if limit > 0: + up = torch.clamp(up, min=-limit, max=limit) + gate = torch.clamp(gate, max=limit) + return F.silu(gate) * up + + +@register_experts_backend("grouped") +class GroupedExperts(RoutedExperts): + _WEIGHT_NAMES = ("w1", "w2", "w3") + + def __init__(self, cfg: MoEConfig, hidden_size: int) -> None: + super().__init__(cfg, hidden_size) + if cfg.dispatch not in DISPATCH_MODES: + raise ValueError( + f"unknown MoE dispatch {cfg.dispatch!r}; choose from {DISPATCH_MODES}" + ) + self.grouped_mm = cfg.dispatch == "grouped_mm" and hasattr(torch, "_grouped_mm") + self.swiglu_limit = float(cfg.swiglu_limit) + e, d, i = self.n_experts, hidden_size, self.intermediate_size + self.w1 = nn.Parameter(torch.empty(e, i, d)) + self.w2 = nn.Parameter(torch.empty(e, d, i)) + self.w3 = nn.Parameter(torch.empty(e, i, d)) + + def reset_parameters(self, std: float) -> None: + if self.w1.device.type == "meta": + return + for name in self._WEIGHT_NAMES: + nn.init.normal_(getattr(self, name), mean=0.0, std=std) + + def forward(self, x: torch.Tensor, routing: RoutingResult) -> torch.Tensor: + flat_expert = routing.indices.flatten() # [T*k] + order = flat_expert.argsort(stable=True) + token_of = order // routing.topk # routed token index per sorted slot + x_sorted = x.index_select(0, token_of) + w_sorted = routing.weights.reshape(-1, 1).index_select(0, order).float() + counts = routing.counts + + if self.grouped_mm and x.is_cuda: + offs = counts.cumsum(0).to(torch.int32) + gate = torch._grouped_mm(x_sorted, self.w1.transpose(-1, -2), offs=offs) + up = torch._grouped_mm(x_sorted, self.w3.transpose(-1, -2), offs=offs) + h = w_sorted * swiglu_clamped(gate, up, self.swiglu_limit) + y_routed = torch._grouped_mm( + h.to(x.dtype), self.w2.transpose(-1, -2), offs=offs + ) + else: + counts_list = counts.tolist() # one host sync per MoE forward + parts = [] + offset = 0 + for i, n in enumerate(counts_list): + if n == 0: + continue + seg = x_sorted[offset : offset + n] + h = w_sorted[offset : offset + n] * swiglu_clamped( + F.linear(seg, self.w1[i]), + F.linear(seg, self.w3[i]), + self.swiglu_limit, + ) + parts.append(F.linear(h.to(seg.dtype), self.w2[i])) + offset += n + if not parts: + return torch.zeros_like(x) + y_routed = torch.cat(parts, dim=0) + + y = torch.zeros(x.shape, dtype=torch.float32, device=x.device) + y = y.index_add(0, token_of, y_routed.float()) + return y.to(x.dtype) + + +_STACKED_KEY = re.compile(r"^(?P(?:.*\.)?experts)\.(?Pw[123])$") +_PER_EXPERT_KEY = re.compile( + r"^(?P(?:.*\.)?experts)\.(?P\d+)\.(?Pw[123])\.weight$" +) + + +def unstack_grouped_expert_state_dict(state: dict) -> dict: + """``experts.w1`` [E, out, in] -> ``experts.{i}.w1.weight``; no-op otherwise.""" + out = {} + for key, value in state.items(): + m = _STACKED_KEY.match(key) + if m is None or not isinstance(value, torch.Tensor) or value.dim() != 3: + out[key] = value + continue + for i in range(value.shape[0]): + out[f"{m['base']}.{i}.{m['w']}.weight"] = value[i] + return out + + +def stack_grouped_expert_state_dict(state: dict) -> dict: + """Inverse of :func:`unstack_grouped_expert_state_dict`.""" + groups: dict = {} + out = {} + for key, value in state.items(): + m = _PER_EXPERT_KEY.match(key) + if m is None: + out[key] = value + continue + groups.setdefault((m["base"], m["w"]), {})[int(m["idx"])] = value + for (base, w), members in groups.items(): + n = max(members) + 1 + if sorted(members) != list(range(n)): + raise KeyError( + f"{base}.*.{w}.weight is missing expert indices: have {sorted(members)}" + ) + out[f"{base}.{w}"] = torch.stack([members[i] for i in range(n)], dim=0) + return out + + +register_state_dict_converter( + "grouped_experts", + to_checkpoint=unstack_grouped_expert_state_dict, + from_checkpoint=stack_grouped_expert_state_dict, +) diff --git a/specforge/modeling/draft/moe/init.py b/specforge/modeling/draft/moe/init.py index 4f1b57e4b..67f448280 100644 --- a/specforge/modeling/draft/moe/init.py +++ b/specforge/modeling/draft/moe/init.py @@ -3,17 +3,22 @@ A drafter whose MoE matches the target's expert shape can inherit expert weights instead of training them from scratch. This module holds the -implementation-independent part: choosing the mapping. Applying a plan needs -the target checkpoint's (possibly quantized) weights and the experts backend's -native layout, and is provided alongside each target-family preset. +mapping (:func:`plan_warm_start`) and applying it to one ``MoELayer`` from a +target layer's *dequantized* tensors in official naming +(:func:`apply_warm_start`). Reading and dequantizing the target checkpoint is +target-specific and lives with the target's tooling. """ from __future__ import annotations from dataclasses import dataclass -from typing import Tuple +from typing import List, Mapping, Tuple + +import torch from .config import MoEConfig +from .layer import MoELayer +from .state_dict import from_checkpoint_state_dict @dataclass(frozen=True) @@ -60,3 +65,42 @@ def plan_warm_start( ), copy_shared_expert=bool(cfg.n_shared_experts), ) + + +_EXPERT_WEIGHTS = ("w1", "w2", "w3") + + +def apply_warm_start( + layer: MoELayer, plan: WarmStartPlan, source: Mapping[str, torch.Tensor] +) -> List[str]: + """Seed ``layer`` from one target MoE layer. + + ``source`` holds the target layer's tensors in official naming, relative + to the layer: ``experts.{j}.w{1,2,3}.weight``, ``gate.weight`` ``[E_t, H]``, + optionally ``gate.bias`` ``[E_t]`` and ``shared_experts.w{1,2,3}.weight``. + Returns the (module-native) keys that were loaded. + """ + if plan.n_draft_experts != layer.cfg.n_routed_experts: + raise ValueError( + f"plan seeds {plan.n_draft_experts} experts but the layer has " + f"{layer.cfg.n_routed_experts}" + ) + official = {} + for i, j in enumerate(plan.target_expert_ids): + for w in _EXPERT_WEIGHTS: + official[f"experts.{i}.{w}.weight"] = source[f"experts.{j}.{w}.weight"] + if plan.copy_gate_rows: + rows = torch.as_tensor(plan.target_expert_ids, dtype=torch.long) + official["gate.weight"] = source["gate.weight"][rows] + if "gate.bias" in source: + official["gate.bias"] = source["gate.bias"][rows] + if plan.copy_shared_expert and layer.shared_experts is not None: + for w in _EXPERT_WEIGHTS: + official[f"shared_experts.{w}.weight"] = source[ + f"shared_experts.{w}.weight" + ] + native = from_checkpoint_state_dict(official) + result = layer.load_state_dict(native, strict=False) + if result.unexpected_keys: + raise KeyError(f"warm start produced unexpected keys: {result.unexpected_keys}") + return sorted(native) diff --git a/specforge/modeling/draft/moe/noaux_tc.py b/specforge/modeling/draft/moe/noaux_tc.py new file mode 100644 index 000000000..14182ccf3 --- /dev/null +++ b/specforge/modeling/draft/moe/noaux_tc.py @@ -0,0 +1,102 @@ +# coding=utf-8 +"""Aux-loss-free balancing (DeepSeek-V3/V4 ``noaux_tc``). + +A per-expert fp32 bias shifts the scores used for *selection*; combine weights +still use the raw scores. A sign controller moves the bias against the +all-reduced expert load, so it is updated by the trainer loop, not gradients. + +Checkpoint naming: the bias is stored as ``.gate.bias`` (the DeepSeek +native key SGLang maps onto ``e_score_correction_bias``); the module keeps it +at ``gate.balance.bias``, converted at the state-dict boundary. +""" + +from __future__ import annotations + +import re +from typing import Dict, Optional + +import torch + +from .balance import BalanceController, MetricValue, register_balance_controller +from .config import MoEConfig +from .state_dict import register_state_dict_converter + + +@register_balance_controller("noaux_tc") +class NoAuxTCController(BalanceController): + def __init__(self, cfg: MoEConfig, n_experts: int) -> None: + super().__init__(cfg, n_experts) + self.update_rate = float(cfg.bias_update_rate) + self.register_buffer("bias", torch.zeros(n_experts, dtype=torch.float32)) + self._pending_counts: Optional[torch.Tensor] = None + self.last_load: Optional[torch.Tensor] = None + + def _apply(self, fn, recurse=True): + module = super()._apply(fn, recurse) + # Sign-controller steps (~1e-3) vanish under bf16 rounding once the + # bias grows; keep the buffer fp32 through module-wide dtype casts. + if module.bias.dtype != torch.float32: + module.bias.data = module.bias.data.float() + return module + + def adjust_selection_scores(self, scores: torch.Tensor) -> torch.Tensor: + return scores + self.bias + + def observe(self, counts: torch.Tensor) -> None: + # Overwrite, never accumulate: a checkpoint recompute re-runs the + # forward and must leave identical state behind. + self._pending_counts = counts + + def apply_pending_update(self) -> None: + import torch.distributed as dist + + counts = self._pending_counts + self._pending_counts = None + if counts is None or self.update_rate <= 0: + return + load = counts.float() + if dist.is_available() and dist.is_initialized(): + dist.all_reduce(load) + self.last_load = load + error = load.mean() - load + with torch.no_grad(): + self.bias += self.update_rate * torch.sign(error) + + def metrics(self) -> Dict[str, MetricValue]: + out: Dict[str, MetricValue] = {"bias_abs_max": self.bias.abs().max()} + if self.last_load is not None: + mean = self.last_load.mean().clamp_min(1e-9) + out["global_load_max_ratio"] = self.last_load.max() / mean + out["global_load_min_ratio"] = self.last_load.min() / mean + return out + + +_NATIVE_BIAS = re.compile(r"^(?P(?:.*\.)?)gate\.balance\.bias$") +_OFFICIAL_BIAS = re.compile(r"^(?P(?:.*\.)?)gate\.bias$") + + +def _to_checkpoint(state: dict) -> dict: + return { + (f"{m['base']}gate.bias" if (m := _NATIVE_BIAS.match(k)) else k): v + for k, v in state.items() + } + + +def _is_moe_layer(state: dict, base: str) -> bool: + return f"{base}experts.w1" in state or f"{base}experts.0.w1.weight" in state + + +def _from_checkpoint(state: dict) -> dict: + out = {} + for key, value in state.items(): + m = _OFFICIAL_BIAS.match(key) + # Only an MoE layer's gate: a dense module named ``gate`` keeps its bias. + if m is not None and _is_moe_layer(state, m["base"]): + key = f"{m['base']}gate.balance.bias" + out[key] = value + return out + + +register_state_dict_converter( + "noaux_tc_bias", to_checkpoint=_to_checkpoint, from_checkpoint=_from_checkpoint +) diff --git a/specforge/modeling/draft/moe/presets.py b/specforge/modeling/draft/moe/presets.py new file mode 100644 index 000000000..3d0bc25d3 --- /dev/null +++ b/specforge/modeling/draft/moe/presets.py @@ -0,0 +1,27 @@ +# coding=utf-8 +"""Target-family MoE presets. + +A preset is the routing recipe of one target family; the draft JSON adds the +per-run sizes (``n_routed_experts``, ``num_experts_per_tok``, +``moe_intermediate_size``) and may override any key for ablations. +""" + +from .config import register_moe_preset + +# DeepSeek-V4 (e.g. DeepSeek-V4-Flash): sqrt(softplus) scores, aux-loss-free +# top-k with the sign-controlled selection bias, renormalized combine weights +# scaled by 1.5, one ungated shared expert, SwiGLU clamped at 10. The target's +# n_group == topk_group, so group-limited routing is off by default. +register_moe_preset( + "deepseek_v4", + scoring_func="sqrtsoftplus", + norm_topk_prob=True, + routed_scaling_factor=1.5, + n_shared_experts=1, + swiglu_limit=10.0, + router="topk", + balance="noaux_tc", + experts_backend="grouped", + shared_expert="swiglu", + shared_expert_gate="none", +) diff --git a/specforge/modeling/draft/moe/swiglu_shared.py b/specforge/modeling/draft/moe/swiglu_shared.py new file mode 100644 index 000000000..e2a06061c --- /dev/null +++ b/specforge/modeling/draft/moe/swiglu_shared.py @@ -0,0 +1,30 @@ +# coding=utf-8 +"""Ungated SwiGLU shared expert (DeepSeek layout ``shared_experts.w{1,2,3}``).""" + +from __future__ import annotations + +import torch +from torch import nn + +from .config import MoEConfig +from .grouped_experts import swiglu_clamped +from .shared import SharedExpert, register_shared_expert + + +@register_shared_expert("swiglu") +class SwiGLUSharedExpert(SharedExpert): + def __init__(self, cfg: MoEConfig, hidden_size: int) -> None: + super().__init__(cfg, hidden_size) + if cfg.shared_expert_gate != "none": + raise ValueError( + "shared_expert='swiglu' is ungated; shared_expert_gate=" + f"{cfg.shared_expert_gate!r} needs a gated shared-expert implementation" + ) + self.swiglu_limit = float(cfg.swiglu_limit) + self.w1 = nn.Linear(hidden_size, self.intermediate_size, bias=False) + self.w2 = nn.Linear(self.intermediate_size, hidden_size, bias=False) + self.w3 = nn.Linear(hidden_size, self.intermediate_size, bias=False) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + h = swiglu_clamped(self.w1(x), self.w3(x), self.swiglu_limit) + return self.w2(h.to(x.dtype)) diff --git a/specforge/modeling/draft/moe/topk_router.py b/specforge/modeling/draft/moe/topk_router.py new file mode 100644 index 000000000..0888dd006 --- /dev/null +++ b/specforge/modeling/draft/moe/topk_router.py @@ -0,0 +1,85 @@ +# coding=utf-8 +"""Top-k router with pluggable score functions and optional group-limited routing.""" + +from __future__ import annotations + +import torch +import torch.nn.functional as F +from torch import nn + +from .balance import BalanceController +from .config import MoEConfig +from .router import ( + Router, + RoutingResult, + get_score_function, + register_router, + register_score_function, +) + + +@register_score_function("softmax") +def _softmax(logits: torch.Tensor) -> torch.Tensor: + return logits.softmax(dim=-1) + + +@register_score_function("sigmoid") +def _sigmoid(logits: torch.Tensor) -> torch.Tensor: + return torch.sigmoid(logits) + + +@register_score_function("sqrtsoftplus") +def _sqrtsoftplus(logits: torch.Tensor) -> torch.Tensor: + """DeepSeek-V4 scoring: ``sqrt(softplus(logits))``.""" + return F.softplus(logits).sqrt() + + +def group_limited_mask( + selection: torch.Tensor, n_group: int, topk_group: int +) -> torch.Tensor: + """Keep only the ``topk_group`` groups with the highest top-2 score sums + (DeepSeek ``noaux_tc`` group scoring); other groups become ``-inf``.""" + tokens, n_experts = selection.shape + grouped = selection.view(tokens, n_group, n_experts // n_group) + group_scores = grouped.topk(min(2, grouped.shape[-1]), dim=-1).values.sum(-1) + keep = group_scores.topk(topk_group, dim=-1).indices + mask = torch.zeros_like(group_scores, dtype=torch.bool).scatter_(1, keep, True) + return grouped.masked_fill(~mask.unsqueeze(-1), float("-inf")).view( + tokens, n_experts + ) + + +@register_router("topk") +class TopKRouter(Router): + """``scores = f(x W^T)``; pick top-k on balance-adjusted scores; combine + with the raw scores (optionally renormalized, then scaled).""" + + def __init__( + self, cfg: MoEConfig, hidden_size: int, balance: BalanceController + ) -> None: + super().__init__(cfg, hidden_size, balance) + self.weight = nn.Parameter(torch.empty(self.n_experts, hidden_size)) + self.score_fn = get_score_function(cfg.scoring_func) + + def reset_parameters(self, std: float) -> None: + nn.init.normal_(self.weight, mean=0.0, std=std) + + def forward(self, x: torch.Tensor) -> RoutingResult: + # Routing math in fp32 regardless of the model dtype. + scores = self.score_fn(F.linear(x.float(), self.weight.float())) + selection = self.balance.adjust_selection_scores(scores) + if self.cfg.group_limited: + selection = group_limited_mask( + selection, self.cfg.n_group, self.cfg.topk_group + ) + indices = selection.topk(self.topk, dim=-1).indices + weights = scores.gather(1, indices) + if self.cfg.norm_topk_prob: + weights = weights / (weights.sum(dim=-1, keepdim=True) + 1e-20) + weights = weights * self.cfg.routed_scaling_factor + flat = indices.flatten() + # scatter_add instead of bincount: CUDA bincount hides a device sync. + counts = torch.zeros( + self.n_experts, dtype=torch.long, device=x.device + ).scatter_add_(0, flat, torch.ones_like(flat)) + return RoutingResult(weights=weights, indices=indices, counts=counts) diff --git a/tests/test_modeling/test_moe_deepseek_v4.py b/tests/test_modeling/test_moe_deepseek_v4.py new file mode 100644 index 000000000..5085d5b4e --- /dev/null +++ b/tests/test_modeling/test_moe_deepseek_v4.py @@ -0,0 +1,414 @@ +# coding=utf-8 +"""DeepSeek-V4 MoE preset: routing, aux-loss-free balancing, grouped experts, +official checkpoint naming, warm start, and DFlash integration.""" + +import json +import unittest +from pathlib import Path + +import torch +from torch import nn +from transformers import Qwen3Config + +from specforge.modeling.draft.dflash import DFlashDraftModel +from specforge.modeling.draft.moe import ( + MOE_PRESETS, + MoELayer, + apply_pending_balance_updates, + apply_warm_start, + collect_moe_metrics, + from_checkpoint_state_dict, + get_score_function, + iter_moe_layers, + plan_warm_start, + resolve_moe_config, + to_checkpoint_state_dict, +) +from specforge.modeling.draft.moe.grouped_experts import ( + GroupedExperts, + stack_grouped_expert_state_dict, + swiglu_clamped, + unstack_grouped_expert_state_dict, +) +from specforge.modeling.draft.moe.noaux_tc import NoAuxTCController +from specforge.modeling.draft.moe.swiglu_shared import SwiGLUSharedExpert +from specforge.modeling.draft.moe.topk_router import TopKRouter, group_limited_mask + +REPO_ROOT = Path(__file__).resolve().parents[2] +CUDA = torch.cuda.is_available() + + +def _json(**overrides): + payload = dict( + moe_preset="deepseek_v4", + n_routed_experts=8, + num_experts_per_tok=2, + moe_intermediate_size=16, + dflash_config={"moe_bias_update_rate": 1e-3}, + ) + payload.update(overrides) + return payload + + +def _layer(**overrides) -> MoELayer: + torch.manual_seed(0) + layer = MoELayer(resolve_moe_config(_json(**overrides)), 32) + layer.reset_parameters(std=0.05) + for p in layer.shared_experts.parameters(): + nn.init.normal_(p, std=0.05) + return layer + + +def _dflash_config(**overrides): + fields = dict( + architectures=["DFlashDraftModel"], + block_size=2, + hidden_size=32, + intermediate_size=64, + num_attention_heads=2, + num_key_value_heads=1, + num_hidden_layers=2, + num_target_layers=6, + head_dim=16, + max_position_embeddings=64, + vocab_size=32, + layer_types=["full_attention", "full_attention"], + initializer_range=0.02, + **_json(dflash_config={"attention_mode": "gqa", "moe_bias_update_rate": 0.005}), + ) + fields.update(overrides) + config = Qwen3Config(**fields) + config._attn_implementation = "sdpa" + return config + + +def _reference_forward(layer: MoELayer, x: torch.Tensor) -> torch.Tensor: + """Per-token dense reference of the routed + shared FFN.""" + routing = layer.gate(x) + e = layer.experts + out = torch.zeros_like(x, dtype=torch.float32) + for t in range(x.shape[0]): + for k in range(routing.topk): + i = int(routing.indices[t, k]) + h = swiglu_clamped(x[t] @ e.w1[i].t(), x[t] @ e.w3[i].t(), e.swiglu_limit) + out[t] += routing.weights[t, k] * (h.to(x.dtype) @ e.w2[i].t()).float() + return (out + layer.shared_experts(x).float()).to(x.dtype) + + +class TestPresetAndConfig(unittest.TestCase): + def test_preset_matches_deepseek_v4_recipe(self): + self.assertIn("deepseek_v4", MOE_PRESETS) + cfg = resolve_moe_config(_json()) + self.assertEqual(cfg.scoring_func, "sqrtsoftplus") + self.assertEqual(cfg.balance, "noaux_tc") + self.assertEqual(cfg.routed_scaling_factor, 1.5) + self.assertTrue(cfg.norm_topk_prob) + self.assertEqual(cfg.swiglu_limit, 10.0) + self.assertEqual(cfg.n_shared_experts, 1) + self.assertFalse(cfg.group_limited) + + def test_checked_in_draft_config_resolves(self): + payload = json.loads( + (REPO_ROOT / "configs" / "deepseek-v4-flash-dspark-moe.json").read_text() + ) + cfg = resolve_moe_config(payload) + self.assertEqual( + (cfg.n_routed_experts, cfg.num_experts_per_tok, cfg.moe_intermediate_size), + (64, 6, 2048), + ) + self.assertEqual(cfg.dispatch, "grouped_mm") + self.assertEqual(cfg.bias_update_rate, 1e-3) + dense = json.loads( + (REPO_ROOT / "configs" / "deepseek-v4-flash-dspark.json").read_text() + ) + moe_only = { + "moe_preset", + "n_routed_experts", + "num_experts_per_tok", + "moe_intermediate_size", + "n_shared_experts", + } + self.assertEqual(set(payload) - set(dense), moe_only) + for key in dense: + if key != "dflash_config": + self.assertEqual(payload[key], dense[key], key) + + def test_score_functions(self): + logits = torch.tensor([[0.0, 2.0, -3.0]]) + self.assertTrue( + torch.allclose( + get_score_function("sqrtsoftplus")(logits), + torch.nn.functional.softplus(logits).sqrt(), + ) + ) + self.assertAlmostEqual( + float(get_score_function("softmax")(logits).sum()), 1.0, places=5 + ) + self.assertTrue( + torch.allclose(get_score_function("sigmoid")(logits), logits.sigmoid()) + ) + + +class TestTopKRouter(unittest.TestCase): + def test_weights_are_renormalized_and_scaled(self): + layer = _layer() + routing = layer.gate(torch.randn(5, 32)) + self.assertIsInstance(layer.gate, TopKRouter) + self.assertTrue(torch.allclose(routing.weights.sum(-1), torch.full((5,), 1.5))) + self.assertEqual(int(routing.counts.sum()), 10) + for row in routing.indices.tolist(): + self.assertEqual(len(set(row)), 2) + + def test_group_limited_routing_stays_within_selected_groups(self): + selection = torch.randn(64, 16) + masked = group_limited_mask(selection, n_group=4, topk_group=1) + finite = torch.isfinite(masked).view(64, 4, 4) + self.assertTrue((finite.all(-1).sum(-1) == 1).all()) + layer = _layer( + n_routed_experts=16, n_group=4, topk_group=1, num_experts_per_tok=3 + ) + routing = layer.gate(torch.randn(10, 32)) + groups = routing.indices // 4 + self.assertTrue((groups == groups[:, :1]).all()) + + +class TestNoAuxTC(unittest.TestCase): + def test_deferred_bias_update_semantics(self): + layer = _layer().train() + ctrl = layer.balance + self.assertIsInstance(ctrl, NoAuxTCController) + layer(torch.randn(6, 32)) + self.assertIsNotNone(ctrl._pending_counts) + before = ctrl.bias.clone() + layer.apply_pending_balance_update() + self.assertIsNone(ctrl._pending_counts) + self.assertFalse(torch.equal(before, ctrl.bias)) + self.assertTrue(((ctrl.bias - before).abs() <= 1e-3 + 1e-7).all()) + after = ctrl.bias.clone() + layer.apply_pending_balance_update() # nothing pending: no-op + self.assertTrue(torch.equal(after, ctrl.bias)) + layer.eval() + layer(torch.randn(6, 32)) + self.assertIsNone(ctrl._pending_counts) + + def test_bias_stays_fp32_and_moves_selection_only(self): + layer = _layer().to(torch.bfloat16).eval() + self.assertEqual(layer.balance.bias.dtype, torch.float32) + x = torch.randn(4, 32, dtype=torch.bfloat16) + self.assertEqual(layer(x).dtype, torch.bfloat16) + layer.balance.bias[:] = -100.0 + layer.balance.bias[3] = 0.0 + routing = layer.gate(x) + self.assertTrue((routing.indices == 3).any(dim=-1).all()) + self.assertTrue(torch.allclose(routing.weights.sum(-1), torch.full((4,), 1.5))) + + def test_metrics_include_bias_and_global_load(self): + layer = _layer().train() + layer(torch.randn(6, 32)) + layer.apply_pending_balance_update() + metrics = collect_moe_metrics(nn.Sequential(layer)) + for key in ( + "moe/load_max_ratio", + "moe/bias_abs_max", + "moe/global_load_max_ratio", + ): + self.assertIn(key, metrics) + + +class TestGroupedExperts(unittest.TestCase): + def test_layout_init_and_dense_reference(self): + layer = _layer().eval() + e = layer.experts + self.assertIsInstance(e, GroupedExperts) + self.assertIsInstance(layer.shared_experts, SwiGLUSharedExpert) + self.assertEqual(tuple(e.w1.shape), (8, 16, 32)) + self.assertEqual(tuple(e.w2.shape), (8, 32, 16)) + self.assertAlmostEqual(float(e.w1.std()), 0.05, delta=0.01) + x = torch.randn(7, 32) + self.assertTrue( + torch.allclose(layer(x), _reference_forward(layer, x), atol=1e-5) + ) + + def test_swiglu_clamp(self): + gate = torch.tensor([50.0, -50.0]) + up = torch.tensor([50.0, -50.0]) + clamped = swiglu_clamped(gate, up, 10.0) + # gate clamps to max 10, up to [-10, 10]: silu(10)*10 and silu(-50)*-10 + expected = torch.nn.functional.silu(torch.tensor([10.0, -50.0])) * torch.tensor( + [10.0, -10.0] + ) + self.assertTrue(torch.allclose(clamped, expected, atol=1e-6)) + self.assertGreater(float(swiglu_clamped(gate, up, 0.0)[0]), 100.0) + + def test_unknown_dispatch_is_rejected(self): + with self.assertRaisesRegex(ValueError, "dispatch"): + _layer(dflash_config={"moe_dispatch": "magic"}) + + @unittest.skipUnless( + CUDA and hasattr(torch, "_grouped_mm"), "needs CUDA grouped GEMM" + ) + def test_grouped_mm_matches_sorted_loop(self): + torch.manual_seed(3) + layer = _layer(n_routed_experts=16).to("cuda", torch.bfloat16).train() + x = (torch.randn(6, 32, device="cuda") * 0.5).to(torch.bfloat16) + results = {} + for grouped in (False, True): + layer.experts.grouped_mm = grouped + layer.zero_grad(set_to_none=True) + xg = x.clone().requires_grad_(True) + y = layer(xg) + y.float().square().sum().backward() + results[grouped] = ( + y.detach().clone(), + xg.grad.clone(), + { + n: p.grad.clone() + for n, p in layer.named_parameters() + if p.grad is not None + }, + ) + (y0, dx0, g0), (y1, dx1, g1) = results[False], results[True] + self.assertTrue(torch.allclose(y0.float(), y1.float(), rtol=2e-2, atol=2e-2)) + self.assertTrue(torch.allclose(dx0.float(), dx1.float(), rtol=2e-2, atol=2e-2)) + self.assertLessEqual(set(g0), set(g1)) + for name in g1: + if name in g0: + self.assertTrue( + torch.allclose( + g0[name].float(), g1[name].float(), rtol=2e-2, atol=2e-2 + ), + name, + ) + else: + self.assertEqual(int(torch.count_nonzero(g1[name])), 0, name) + + +class TestCheckpointNaming(unittest.TestCase): + def test_layer_roundtrip_through_official_naming(self): + layer = _layer() + native = layer.state_dict() + self.assertIn("experts.w1", native) + self.assertIn("gate.balance.bias", native) + official = to_checkpoint_state_dict(native) + self.assertIn("experts.0.w1.weight", official) + self.assertIn("gate.bias", official) + self.assertIn("shared_experts.w1.weight", official) + self.assertFalse( + any("balance" in k or k.endswith("experts.w1") for k in official) + ) + fresh = _layer(dflash_config={"moe_bias_update_rate": 0.0}) + fresh.load_state_dict(from_checkpoint_state_dict(official), strict=True) + self.assertTrue(torch.equal(fresh.experts.w2, layer.experts.w2)) + # both directions are idempotent + self.assertEqual(set(to_checkpoint_state_dict(official)), set(official)) + self.assertEqual(set(from_checkpoint_state_dict(native)), set(native)) + + def test_dense_gate_bias_is_left_alone(self): + state = { + "head.gate.bias": torch.zeros(1), + "head.gate.weight": torch.zeros(1, 1), + } + self.assertEqual(set(from_checkpoint_state_dict(state)), set(state)) + + def test_stack_rejects_missing_expert_indices(self): + official = unstack_grouped_expert_state_dict(_layer().state_dict()) + del official["experts.3.w2.weight"] + with self.assertRaises(KeyError): + stack_grouped_expert_state_dict(official) + + +class TestWarmStart(unittest.TestCase): + def test_apply_plan_copies_selected_experts_gate_rows_and_shared(self): + layer = _layer() + n_target = 16 + source = { + "gate.weight": torch.randn(n_target, 32), + "gate.bias": torch.randn(n_target), + } + for j in range(n_target): + source[f"experts.{j}.w1.weight"] = torch.randn(16, 32) + source[f"experts.{j}.w2.weight"] = torch.randn(32, 16) + source[f"experts.{j}.w3.weight"] = torch.randn(16, 32) + for w, shape in (("w1", (16, 32)), ("w2", (32, 16)), ("w3", (16, 32))): + source[f"shared_experts.{w}.weight"] = torch.randn(*shape) + plan = plan_warm_start(layer.cfg, n_target_experts=n_target) + self.assertEqual(plan.target_expert_ids, (0, 2, 4, 6, 8, 10, 12, 14)) + loaded = apply_warm_start(layer, plan, source) + self.assertIn("experts.w1", loaded) + for i, j in enumerate(plan.target_expert_ids): + self.assertTrue( + torch.equal(layer.experts.w1[i], source[f"experts.{j}.w1.weight"]) + ) + self.assertTrue(torch.equal(layer.gate.weight[i], source["gate.weight"][j])) + self.assertEqual( + float(layer.balance.bias[i]), float(source["gate.bias"][j]) + ) + self.assertTrue( + torch.equal( + layer.shared_experts.w2.weight, source["shared_experts.w2.weight"] + ) + ) + with self.assertRaises(ValueError): + apply_warm_start(_layer(n_routed_experts=4), plan, source) + + +class TestDFlashIntegration(unittest.TestCase): + def _forward(self, model): + return model( + position_ids=torch.arange(6).unsqueeze(0), + noise_embedding=torch.randn(1, 2, 32), + target_hidden=torch.randn(1, 4, 2 * 32), + ) + + def test_layers_train_and_balance_through_the_model(self): + model = DFlashDraftModel(_dflash_config()) + layers = list(iter_moe_layers(model)) + self.assertEqual(len(layers), 2) + for layer in layers: + self.assertIsInstance(layer.experts, GroupedExperts) + self.assertAlmostEqual(float(layer.experts.w1.std()), 0.02, delta=0.005) + self.assertAlmostEqual(float(layer.gate.weight.std()), 0.02, delta=0.005) + self.assertTrue(torch.equal(layer.balance.bias, torch.zeros(8))) + model.train() + out = self._forward(model) + out.float().square().mean().backward() + for layer in layers: + self.assertIsNotNone(layer.experts.w2.grad) + self.assertIsNotNone(layer.gate.weight.grad) + self.assertTrue(torch.equal(layer.balance.bias, torch.zeros(8))) # deferred + self._forward(model) # applies the pending update before routing + self.assertTrue(any(layer.balance.bias.abs().sum() > 0 for layer in layers)) + + def test_model_checkpoint_uses_official_naming_and_reloads(self): + model = DFlashDraftModel(_dflash_config()) + official = to_checkpoint_state_dict(model.state_dict()) + self.assertIn("layers.0.mlp.experts.0.w1.weight", official) + self.assertIn("layers.0.mlp.gate.bias", official) + self.assertIn("layers.1.mlp.shared_experts.w3.weight", official) + self.assertFalse( + any(".balance." in k or k.endswith(".experts.w1") for k in official) + ) + fresh = DFlashDraftModel(_dflash_config()) + fresh.load_state_dict(from_checkpoint_state_dict(official), strict=True) + self.assertTrue( + torch.equal(fresh.layers[1].mlp.experts.w1, model.layers[1].mlp.experts.w1) + ) + + def test_dense_config_is_unaffected(self): + config = _dflash_config() + for key in ( + "moe_preset", + "n_routed_experts", + "num_experts_per_tok", + "moe_intermediate_size", + "n_shared_experts", + ): + if hasattr(config, key): + delattr(config, key) + model = DFlashDraftModel(config) + self.assertEqual(list(iter_moe_layers(model)), []) + apply_pending_balance_updates(model) + + +if __name__ == "__main__": + unittest.main(verbosity=2) diff --git a/tests/test_runtime/test_package_architecture.py b/tests/test_runtime/test_package_architecture.py index 9ede76de6..b0cfd681c 100644 --- a/tests/test_runtime/test_package_architecture.py +++ b/tests/test_runtime/test_package_architecture.py @@ -711,6 +711,7 @@ def test_dspark_configs_are_qwen3_family(self): self.assertEqual( set(dspark_configs), { + "deepseek-v4-flash-dspark-moe.json", "deepseek-v4-flash-dspark.json", "glm-5.2-dspark.json", "inkling-dspark.json", From 3b91aa25087bc807fd9adbee5987abe59b25503b Mon Sep 17 00:00:00 2001 From: Kan Wu Date: Thu, 3 Sep 2026 07:04:44 +0000 Subject: [PATCH 06/11] fix: make MoE drafter exports servable and reloadable - MoEConfig.serving_fields(): the resolved recipe in the DeepSeek HF config vocabulary (scoring_func, topk_method=noaux_tc, routed_scaling_factor, n_group/topk_group, norm_topk_prob, swiglu_limit when enabled, ...). export --to hf writes them into config.json so a serving engine needs no knowledge of SpecForge presets. - AutoDraftModel.from_pretrained: HF assigns tensors by key and cannot regroup per-expert files into the stacked module layout (and refuses an explicit state_dict next to a path), so for MoE configs read the safetensors, convert through from_checkpoint_state_dict, and load into a freshly built module. Warm starts from HF export dirs go through this path. - Tests: serving fields; an end-to-end export -> config.json -> from_pretrained round trip on a tiny MoE DFlash draft. Co-Authored-By: Claude Fable 5.1 --- specforge/export/to_hf.py | 8 ++- specforge/modeling/auto.py | 46 ++++++++++++++++ specforge/modeling/draft/moe/config.py | 24 +++++++++ tests/test_modeling/test_moe_deepseek_v4.py | 58 +++++++++++++++++++++ 4 files changed, 135 insertions(+), 1 deletion(-) diff --git a/specforge/export/to_hf.py b/specforge/export/to_hf.py index d224e5c0c..8ed702321 100644 --- a/specforge/export/to_hf.py +++ b/specforge/export/to_hf.py @@ -31,7 +31,7 @@ materialize_draft, resolve_training_state, ) -from specforge.modeling.draft.moe import to_checkpoint_state_dict +from specforge.modeling.draft.moe import resolve_moe_config, to_checkpoint_state_dict def _load_embedding_tensor(source: str, key: str) -> torch.Tensor: @@ -89,6 +89,12 @@ def export_to_hf( state, draft_config_path, vocab_mapping_path=vocab_mapping_path ) full_state = dict(to_checkpoint_state_dict(model.state_dict())) + moe_cfg = resolve_moe_config(model.config) + if moe_cfg is not None: + # Serving engines read the routing recipe from config.json; the draft + # JSON only names a preset, so materialize the resolved fields. + for key, value in moe_cfg.serving_fields().items(): + setattr(model.config, key, value) owns_embedding = hasattr(model, "embed_tokens") if owns_embedding and "embed_tokens.weight" not in state["draft_state_dict"]: if not embedding_source: diff --git a/specforge/modeling/auto.py b/specforge/modeling/auto.py index de19dc5c7..686b2b7e8 100644 --- a/specforge/modeling/auto.py +++ b/specforge/modeling/auto.py @@ -62,6 +62,22 @@ def filtered_warning(msg): config = AutoConfig.from_pretrained(pretrained_model_name_or_path) model_cls = cls._model_cls_from_config(config) kwargs = {**kwargs, "config": config} + state_dict = _native_state_dict_for(config, pretrained_model_name_or_path) + if state_dict is not None: + # HF from_pretrained assigns tensors by key and refuses an + # explicit state_dict alongside a path, so build the module and + # load the converted state ourselves. + torch_dtype = kwargs.pop("torch_dtype", kwargs.pop("dtype", None)) + model = model_cls._from_config(config, torch_dtype=torch_dtype) + result = model.load_state_dict(state_dict, strict=False) + missing = [k for k in result.missing_keys if "embed_tokens" not in k] + if missing or result.unexpected_keys: + raise ValueError( + f"{pretrained_model_name_or_path!r} does not match " + f"{model_cls.__name__}: missing {missing[:5]}, " + f"unexpected {list(result.unexpected_keys)[:5]}" + ) + return model.eval() model = model_cls.from_pretrained( pretrained_model_name_or_path, *model_args, **kwargs ) @@ -71,6 +87,36 @@ def filtered_warning(msg): return model +def _native_state_dict_for(config, pretrained_model_name_or_path): + """Checkpoint files use the official parameter naming; modules may use a + native layout (MoE experts). HF ``from_pretrained`` assigns tensors by key + and cannot regroup them, so read the files and convert at this boundary. + Returns ``None`` when no conversion is needed (dense drafts).""" + from specforge.modeling.draft.moe import ( + from_checkpoint_state_dict, + is_moe_config, + ) + + if not is_moe_config(config): + return None + import glob + + from safetensors.torch import load_file + + path = str(pretrained_model_name_or_path) + if not os.path.isdir(path): + from huggingface_hub import snapshot_download + + path = snapshot_download(path, allow_patterns=["*.safetensors", "*.json"]) + files = sorted(glob.glob(os.path.join(path, "*.safetensors"))) + if not files: + raise FileNotFoundError(f"no safetensors weights under {path!r}") + state = {} + for file in files: + state.update(load_file(file)) + return from_checkpoint_state_dict(state) + + class AutoDraftModelConfig: @classmethod def from_file(cls, config_path: str): diff --git a/specforge/modeling/draft/moe/config.py b/specforge/modeling/draft/moe/config.py index d2a55a802..66fb2c761 100644 --- a/specforge/modeling/draft/moe/config.py +++ b/specforge/modeling/draft/moe/config.py @@ -125,6 +125,30 @@ def group_limited(self) -> bool: def as_dict(self) -> Dict[str, Any]: return {f.name: getattr(self, f.name) for f in fields(self)} + def serving_fields(self) -> Dict[str, Any]: + """The resolved recipe in the DeepSeek HF config vocabulary. + + Exports write these to ``config.json`` so a serving engine reads the + complete routing recipe without knowing SpecForge presets. Only the + keys a DeepSeek-style MoE reads; ``swiglu_limit`` is omitted when the + clamp is off (a serving engine treats 0 as a clamp at 0). + """ + out: Dict[str, Any] = { + "n_routed_experts": self.n_routed_experts, + "num_experts_per_tok": self.num_experts_per_tok, + "moe_intermediate_size": self.moe_intermediate_size, + "n_shared_experts": self.n_shared_experts, + "scoring_func": self.scoring_func, + "norm_topk_prob": self.norm_topk_prob, + "routed_scaling_factor": self.routed_scaling_factor, + "n_group": self.n_group, + "topk_group": self.topk_group, + "topk_method": "noaux_tc" if self.balance == "noaux_tc" else "greedy", + } + if self.swiglu_limit > 0: + out["swiglu_limit"] = self.swiglu_limit + return out + #: preset name -> architecture defaults (a subset of ARCHITECTURE_KEYS). MOE_PRESETS: Registry[Dict[str, Any]] = Registry("MoE preset") diff --git a/tests/test_modeling/test_moe_deepseek_v4.py b/tests/test_modeling/test_moe_deepseek_v4.py index 5085d5b4e..2db92fbe5 100644 --- a/tests/test_modeling/test_moe_deepseek_v4.py +++ b/tests/test_modeling/test_moe_deepseek_v4.py @@ -352,6 +352,64 @@ def test_apply_plan_copies_selected_experts_gate_rows_and_shared(self): apply_warm_start(_layer(n_routed_experts=4), plan, source) +class TestServingExport(unittest.TestCase): + def test_serving_fields_carry_the_resolved_recipe(self): + fields = resolve_moe_config(_json()).serving_fields() + self.assertEqual(fields["topk_method"], "noaux_tc") + self.assertEqual(fields["scoring_func"], "sqrtsoftplus") + self.assertEqual(fields["routed_scaling_factor"], 1.5) + self.assertEqual(fields["swiglu_limit"], 10.0) + self.assertEqual((fields["n_group"], fields["topk_group"]), (1, 1)) + # a disabled clamp is omitted rather than exported as a clamp at 0 + self.assertNotIn( + "swiglu_limit", resolve_moe_config(_json(swiglu_limit=0)).serving_fields() + ) + + def test_hf_export_reloads_and_carries_serving_config(self): + import os + import tempfile + + from specforge.export import export_to_hf + from specforge.modeling.auto import AutoDraftModel + + torch.manual_seed(1) + config = _dflash_config() + model = DFlashDraftModel(config).to(torch.bfloat16) + for layer in iter_moe_layers(model): + layer.balance.bias.uniform_(-1.0, 1.0) + workdir = tempfile.mkdtemp(prefix="moe_export_") + config_path = os.path.join(workdir, "draft.json") + config.save_pretrained(workdir) + os.replace(os.path.join(workdir, "config.json"), config_path) + ckpt_dir = os.path.join(workdir, "run-step1") + os.makedirs(ckpt_dir) + torch.save( + { + "draft_state_dict": to_checkpoint_state_dict(model.state_dict()), + "strategy": "dflash", + "global_step": 1, + }, + os.path.join(ckpt_dir, "training_state.pt"), + ) + out = export_to_hf(ckpt_dir, config_path, os.path.join(workdir, "hf")) + exported = json.loads((Path(out) / "config.json").read_text()) + self.assertEqual(exported["topk_method"], "noaux_tc") + self.assertEqual(exported["scoring_func"], "sqrtsoftplus") + self.assertEqual(exported["n_routed_experts"], 8) + from safetensors import safe_open + + with safe_open(os.path.join(out, "model.safetensors"), "pt") as f: + keys = set(f.keys()) + self.assertIn("layers.0.mlp.experts.0.w1.weight", keys) + self.assertIn("layers.0.mlp.gate.bias", keys) + # HF from_pretrained assigns tensors by key; SpecForge's loader must + # convert the official naming back into the stacked module layout. + reloaded = AutoDraftModel.from_pretrained(out, torch_dtype=torch.bfloat16) + fresh = reloaded.state_dict() + for key, value in model.state_dict().items(): + self.assertTrue(torch.equal(value.float(), fresh[key].float()), key) + + class TestDFlashIntegration(unittest.TestCase): def _forward(self, model): return model( From 1c7ae9e08e0d90e40f9d93a6eb0f92d324947e36 Mon Sep 17 00:00:00 2001 From: Kan Wu Date: Thu, 3 Sep 2026 07:27:38 +0000 Subject: [PATCH 07/11] feat: optional auxiliary balance loss for the noaux_tc controller dflash_config.moe_aux_loss_coeff > 0 adds DeepSeek-V3's complementary sequence-wise balance loss (coeff * sum_e f_e P_e over the micro-batch) built from the router's differentiable scores; DFlash/DSpark strategies add it to the training loss and it appears as moe/aux_loss. Why: the bias controller only reorders selection. A from-scratch DFlash-family drafter feeds the router near-identical mask-token embeddings early on, so every token routed to the same experts (load_max_ratio pinned at E/k, ~85% of experts unused through step 100) while the gate logits grew faster than the bias could follow. The differentiable term bounds the logit gaps so input-driven routing can emerge. Co-Authored-By: Claude Fable 5.1 --- specforge/modeling/draft/moe/noaux_tc.py | 32 +++++++++++++++++++-- specforge/modeling/draft/moe/topk_router.py | 4 ++- specforge/training/strategies/base.py | 15 ++++++++-- tests/test_modeling/test_moe_deepseek_v4.py | 29 +++++++++++++++++++ 4 files changed, 75 insertions(+), 5 deletions(-) diff --git a/specforge/modeling/draft/moe/noaux_tc.py b/specforge/modeling/draft/moe/noaux_tc.py index 14182ccf3..063aec649 100644 --- a/specforge/modeling/draft/moe/noaux_tc.py +++ b/specforge/modeling/draft/moe/noaux_tc.py @@ -5,6 +5,16 @@ still use the raw scores. A sign controller moves the bias against the all-reduced expert load, so it is updated by the trainer loop, not gradients. +Optionally (``dflash_config.moe_aux_loss_coeff`` > 0) the controller also +emits DeepSeek-V3's complementary sequence-wise balance loss, +``coeff * sum_e f_e * P_e`` with ``f_e`` the (E/(k*T))-scaled routed-token +fraction and ``P_e`` the mean normalized affinity, computed over the +micro-batch. The bias alone only *reorders* selection; from-scratch drafters +whose early router inputs are near-identical (mask-token embeddings) collapse +onto a few experts while the gate logits grow faster than the bias can +follow. The differentiable term bounds the logit gaps so input-driven routing +can emerge. + Checkpoint naming: the bias is stored as ``.gate.bias`` (the DeepSeek native key SGLang maps onto ``e_score_correction_bias``); the module keeps it at ``gate.balance.bias``, converted at the state-dict boundary. @@ -19,6 +29,7 @@ from .balance import BalanceController, MetricValue, register_balance_controller from .config import MoEConfig +from .router import RoutingResult from .state_dict import register_state_dict_converter @@ -27,8 +38,10 @@ class NoAuxTCController(BalanceController): def __init__(self, cfg: MoEConfig, n_experts: int) -> None: super().__init__(cfg, n_experts) self.update_rate = float(cfg.bias_update_rate) + self.aux_loss_coeff = float(cfg.aux_loss_coeff) self.register_buffer("bias", torch.zeros(n_experts, dtype=torch.float32)) self._pending_counts: Optional[torch.Tensor] = None + self._aux_loss: Optional[torch.Tensor] = None self.last_load: Optional[torch.Tensor] = None def _apply(self, fn, recurse=True): @@ -42,10 +55,23 @@ def _apply(self, fn, recurse=True): def adjust_selection_scores(self, scores: torch.Tensor) -> torch.Tensor: return scores + self.bias - def observe(self, counts: torch.Tensor) -> None: + def observe(self, routing: RoutingResult) -> None: # Overwrite, never accumulate: a checkpoint recompute re-runs the # forward and must leave identical state behind. - self._pending_counts = counts + self._pending_counts = routing.counts + self._aux_loss = None + scores = routing.scores + if self.aux_loss_coeff <= 0 or scores is None or not scores.requires_grad: + return + tokens, n_experts = scores.shape + # f_e: routed fraction scaled so that a uniform load gives 1. + f = routing.counts.float() * (n_experts / (routing.topk * tokens)) + # P_e: mean normalized affinity (the differentiable side). + p = (scores / scores.sum(dim=-1, keepdim=True).clamp_min(1e-20)).mean(0) + self._aux_loss = self.aux_loss_coeff * (f * p).sum() + + def aux_loss(self) -> Optional[torch.Tensor]: + return self._aux_loss def apply_pending_update(self) -> None: import torch.distributed as dist @@ -64,6 +90,8 @@ def apply_pending_update(self) -> None: def metrics(self) -> Dict[str, MetricValue]: out: Dict[str, MetricValue] = {"bias_abs_max": self.bias.abs().max()} + if self._aux_loss is not None: + out["aux_loss"] = self._aux_loss.detach() if self.last_load is not None: mean = self.last_load.mean().clamp_min(1e-9) out["global_load_max_ratio"] = self.last_load.max() / mean diff --git a/specforge/modeling/draft/moe/topk_router.py b/specforge/modeling/draft/moe/topk_router.py index 0888dd006..c4376d632 100644 --- a/specforge/modeling/draft/moe/topk_router.py +++ b/specforge/modeling/draft/moe/topk_router.py @@ -82,4 +82,6 @@ def forward(self, x: torch.Tensor) -> RoutingResult: counts = torch.zeros( self.n_experts, dtype=torch.long, device=x.device ).scatter_add_(0, flat, torch.ones_like(flat)) - return RoutingResult(weights=weights, indices=indices, counts=counts) + return RoutingResult( + weights=weights, indices=indices, counts=counts, scores=scores + ) diff --git a/specforge/training/strategies/base.py b/specforge/training/strategies/base.py index 175545b54..0f1ddf526 100644 --- a/specforge/training/strategies/base.py +++ b/specforge/training/strategies/base.py @@ -66,6 +66,17 @@ def _moe_metrics(model_wrapper: nn.Module) -> Dict[str, Any]: return collect_moe_metrics(draft_model) +def _with_moe_aux_loss(model_wrapper: nn.Module, loss: torch.Tensor) -> torch.Tensor: + """Add the MoE balance policies' auxiliary loss (if any) for this forward.""" + from specforge.modeling.draft.moe import collect_moe_aux_loss + + draft_model = getattr(model_wrapper, "draft_model", None) + if draft_model is None: + return loss + aux = collect_moe_aux_loss(draft_model) + return loss if aux is None else loss + aux.to(loss.dtype) + + def linear_lambda_base( global_step: int, total_steps: int, @@ -558,7 +569,7 @@ def forward_loss( metrics["selector_loss_alpha"] = model_metrics["selector_loss_alpha"] metrics.update(_moe_metrics(self.dflash_model)) return StepOutput( - loss=loss, + loss=_with_moe_aux_loss(self.dflash_model, loss), metrics=metrics, ratio_metrics=model_metrics.get("ratio_metrics", {}), loss_terms=model_metrics.get("loss_terms"), @@ -629,7 +640,7 @@ def forward_loss( metrics[name] = model_metrics[name] metrics.update(_moe_metrics(self.dspark_model)) return StepOutput( - loss=loss, + loss=_with_moe_aux_loss(self.dspark_model, loss), metrics=metrics, ratio_metrics=ratio_metrics, ) diff --git a/tests/test_modeling/test_moe_deepseek_v4.py b/tests/test_modeling/test_moe_deepseek_v4.py index 2db92fbe5..cde7b0649 100644 --- a/tests/test_modeling/test_moe_deepseek_v4.py +++ b/tests/test_modeling/test_moe_deepseek_v4.py @@ -202,6 +202,35 @@ def test_bias_stays_fp32_and_moves_selection_only(self): self.assertTrue((routing.indices == 3).any(dim=-1).all()) self.assertTrue(torch.allclose(routing.weights.sum(-1), torch.full((4,), 1.5))) + def test_aux_balance_loss_is_differentiable_and_uniform_at_balance(self): + layer = _layer(dflash_config={"moe_aux_loss_coeff": 0.5}).train() + x = torch.randn(64, 32) + layer(x) + aux = layer.aux_loss() + self.assertIsNotNone(aux) + self.assertTrue(aux.requires_grad) + aux.backward() + self.assertIsNotNone(layer.gate.weight.grad) + self.assertGreater(float(layer.gate.weight.grad.abs().sum()), 0.0) + # perfectly uniform routing and affinities give exactly coeff * 1 + routing = layer.gate(x) + n_experts = layer.cfg.n_routed_experts + uniform_scores = torch.full((64, n_experts), 0.25, requires_grad=True) + counts = torch.full((n_experts,), 64 * routing.topk // n_experts) + layer.balance.observe( + type(routing)(routing.weights, routing.indices, counts, uniform_scores) + ) + self.assertAlmostEqual(float(layer.balance.aux_loss()), 0.5, places=5) + self.assertIn("aux_loss", layer.balance.metrics()) + # disabled by default, and never built without a gradient signal + layer = _layer().train() + layer(torch.randn(8, 32)) + self.assertIsNone(layer.aux_loss()) + with torch.no_grad(): + aux_layer = _layer(dflash_config={"moe_aux_loss_coeff": 0.5}).train() + aux_layer(torch.randn(8, 32)) + self.assertIsNone(aux_layer.aux_loss()) + def test_metrics_include_bias_and_global_load(self): layer = _layer().train() layer(torch.randn(6, 32)) From 337960ed3d1c55dc67b42bb5450107bf3717fa34 Mon Sep 17 00:00:00 2001 From: Kan Wu Date: Fri, 4 Sep 2026 06:44:46 +0000 Subject: [PATCH 08/11] feat: warm-start MoE drafters from a DeepSeek-V4 target's experts - moe/deepseek_v4_target.py: dequantize one target layer's ffn.* tensors (packed fp4 e2m1 experts with per-32 ue8m0 scales, fp8 e4m3 shared expert with 128x128-block scales, bf16 gate, fp32 noaux bias) into the official naming apply_warm_start() consumes; hash-routed layers are rejected. - scripts/warm_start_moe_drafter.py: build the draft from its config, seed each MoE layer from a chosen target layer (identity expert mapping when the shapes match, strided subset otherwise), record provenance and the serving fields, and write an HF dir usable as model.draft_checkpoint_path. - Tests for the fp4/fp8 conventions and the layer dequant/naming. Co-Authored-By: Claude Fable 5.1 --- scripts/warm_start_moe_drafter.py | 112 +++++++++++++++ .../modeling/draft/moe/deepseek_v4_target.py | 130 ++++++++++++++++++ tests/test_modeling/test_moe_deepseek_v4.py | 55 ++++++++ 3 files changed, 297 insertions(+) create mode 100644 scripts/warm_start_moe_drafter.py create mode 100644 specforge/modeling/draft/moe/deepseek_v4_target.py diff --git a/scripts/warm_start_moe_drafter.py b/scripts/warm_start_moe_drafter.py new file mode 100644 index 000000000..449f9525b --- /dev/null +++ b/scripts/warm_start_moe_drafter.py @@ -0,0 +1,112 @@ +#!/usr/bin/env python3 +"""Build a warm-start source for an MoE DFlash-family drafter from a DeepSeek-V4 target. + +Constructs the draft model from its config (random init for attention, heads, +norms, projections), seeds every MoE layer's routed experts, gate weight, +noaux bias and shared expert from one target layer (dequantized fp4/fp8 -> +bf16), and writes an HF-format directory usable as +``model.draft_checkpoint_path``. Requires the draft's expert shape to match the +target's (256 x 2048 for DeepSeek-V4-Flash); a smaller draft takes a strided +subset of experts. + +Example: + python scripts/warm_start_moe_drafter.py \ + --draft-config examples/configs/kan-ablations/deepseek-v4-flash-dspark-moe256-auxbal.json \ + --target-snapshot /cluster-storage/models/models--deepseek-ai--DeepSeek-V4-Flash-0731/snapshots/ \ + --target-layers 3,11,21,31,41 --output-dir warm-starts/moe256-from-0731 +""" + +from __future__ import annotations + +import argparse +import json +import os +import time + +import torch + + +def main() -> int: + ap = argparse.ArgumentParser( + description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter + ) + ap.add_argument("--draft-config", required=True) + ap.add_argument( + "--target-snapshot", required=True, help="local HF snapshot dir of the target" + ) + ap.add_argument( + "--target-layers", + required=True, + help="comma list, one target layer per draft layer", + ) + ap.add_argument("--output-dir", required=True) + ap.add_argument("--seed", type=int, default=42) + ap.add_argument( + "--select", + default="strided", + help="expert selection when draft has fewer experts", + ) + args = ap.parse_args() + + from specforge.modeling.auto import AutoDraftModel, AutoDraftModelConfig + from specforge.modeling.draft.moe import ( + apply_warm_start, + iter_moe_layers, + plan_warm_start, + resolve_moe_config, + to_checkpoint_state_dict, + ) + from specforge.modeling.draft.moe.deepseek_v4_target import load_target_moe_layer + + target_layers = [int(x) for x in args.target_layers.split(",")] + config = AutoDraftModelConfig.from_file(args.draft_config) + moe_cfg = resolve_moe_config(config) + if moe_cfg is None: + raise SystemExit("draft config is dense (n_routed_experts == 0)") + torch.manual_seed(args.seed) + t0 = time.time() + model = AutoDraftModel.from_config(config, torch_dtype=torch.bfloat16) + layers = list(iter_moe_layers(model)) + if len(layers) != len(target_layers): + raise SystemExit( + f"{len(layers)} MoE layers but {len(target_layers)} target layers given" + ) + print( + f"built draft ({sum(p.numel() for p in model.parameters())/1e9:.2f}B params) in {time.time()-t0:.0f}s" + ) + + target_cfg = json.load(open(os.path.join(args.target_snapshot, "config.json"))) + n_target = int(target_cfg["n_routed_experts"]) + for i, (layer, tl) in enumerate(zip(layers, target_layers)): + t1 = time.time() + source = load_target_moe_layer(args.target_snapshot, tl) + plan = plan_warm_start( + layer.cfg, n_target_experts=n_target, strategy=args.select + ) + loaded = apply_warm_start(layer, plan, source) + print( + f"draft layer {i} <- target layer {tl}: {len(loaded)} tensors, " + f"experts {plan.target_expert_ids[:3]}...{plan.target_expert_ids[-1]} ({time.time()-t1:.0f}s)" + ) + + for key, value in moe_cfg.serving_fields().items(): + setattr(model.config, key, value) + model.config.moe_warm_start = { + "target": os.path.basename( + os.path.dirname(os.path.dirname(args.target_snapshot.rstrip("/"))) + ), + "snapshot": os.path.basename(args.target_snapshot.rstrip("/")), + "target_layers": target_layers, + "select": args.select, + "seed": args.seed, + } + os.makedirs(args.output_dir, exist_ok=True) + model.save_pretrained( + args.output_dir, state_dict=to_checkpoint_state_dict(model.state_dict()) + ) + print(f"wrote {args.output_dir} in {time.time()-t0:.0f}s total") + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/specforge/modeling/draft/moe/deepseek_v4_target.py b/specforge/modeling/draft/moe/deepseek_v4_target.py new file mode 100644 index 000000000..17a76ed36 --- /dev/null +++ b/specforge/modeling/draft/moe/deepseek_v4_target.py @@ -0,0 +1,130 @@ +# coding=utf-8 +"""Read one DeepSeek-V4 target MoE layer as a warm-start source. + +The DeepSeek-V4-Flash checkpoints store routed experts as packed FP4 e2m1 +(two values per int8, low nibble first) with per-32 ue8m0 scales, the shared +expert and other linears as FP8 E4M3 with 128x128-block ue8m0 scales, the +gate in bf16 and the ``noaux_tc`` bias in fp32. This module dequantizes one +layer's ``ffn.*`` tensors to bf16 in the official naming +:func:`specforge.modeling.draft.moe.init.apply_warm_start` consumes +(``experts.{i}.w{1,2,3}.weight``, ``gate.weight``, ``gate.bias``, +``shared_experts.w{1,2,3}.weight``). Conventions follow the reference +``inference/convert.py``. + +Layers below ``num_hash_layers`` route by token hash (``gate.tid2eid``) and +carry no learned gate; they are rejected as warm-start sources. +""" + +from __future__ import annotations + +import json +import os +from typing import Dict, Iterable, Mapping, Tuple + +import torch + +FP4_TABLE = torch.tensor( + [ + 0.0, + 0.5, + 1.0, + 1.5, + 2.0, + 3.0, + 4.0, + 6.0, + 0.0, + -0.5, + -1.0, + -1.5, + -2.0, + -3.0, + -4.0, + -6.0, + ], + dtype=torch.float32, +) +FP8_BLOCK = 128 +FP4_GROUP = 32 + + +def dequant_fp8_block(weight: torch.Tensor, scale: torch.Tensor) -> torch.Tensor: + """FP8 E4M3 ``[out, in]`` with e8m0 scale ``[out/128, in/128]`` -> bf16.""" + out_dim, in_dim = weight.shape + if out_dim % FP8_BLOCK or in_dim % FP8_BLOCK: + raise ValueError(f"fp8 weight {tuple(weight.shape)} is not 128-block aligned") + w = weight.float().view( + out_dim // FP8_BLOCK, FP8_BLOCK, in_dim // FP8_BLOCK, FP8_BLOCK + ) + w = w * scale.float()[:, None, :, None] + return w.view(out_dim, in_dim).to(torch.bfloat16) + + +def dequant_fp4_packed(weight: torch.Tensor, scale: torch.Tensor) -> torch.Tensor: + """Packed FP4 int8 ``[out, in/2]`` with e8m0 scale ``[out, in/32]`` -> bf16. + + Low nibble = even element, high nibble = odd element.""" + if weight.dtype != torch.int8: + raise TypeError(f"packed fp4 weight must be int8, got {weight.dtype}") + out_dim, half_in = weight.shape + in_dim = half_in * 2 + x = weight.view(torch.uint8) + decoded = torch.stack( + [FP4_TABLE[(x & 0x0F).long()], FP4_TABLE[((x >> 4) & 0x0F).long()]], dim=-1 + ).view(out_dim, in_dim // FP4_GROUP, FP4_GROUP) + decoded = decoded * scale.float()[:, :, None] + return decoded.view(out_dim, in_dim).to(torch.bfloat16) + + +def dequantize_ffn_tensors( + raw: Mapping[str, torch.Tensor], prefix: str +) -> Dict[str, torch.Tensor]: + """``{prefix}...`` raw tensors of one MoE layer -> official-relative bf16 dict.""" + out: Dict[str, torch.Tensor] = {} + for name, tensor in raw.items(): + if not name.startswith(prefix) or name.endswith(".scale"): + continue + rel = name[len(prefix) :] + if tensor.dtype == torch.float8_e4m3fn: + value = dequant_fp8_block(tensor, raw[name[: -len(".weight")] + ".scale"]) + elif tensor.dtype == torch.int8: + value = dequant_fp4_packed(tensor, raw[name[: -len(".weight")] + ".scale"]) + elif rel == "gate.bias": + value = tensor.float() + elif rel == "gate.tid2eid": + raise ValueError(f"{prefix} is a hash-routed layer (no learned gate)") + else: + value = tensor.to(torch.bfloat16) + out[rel] = value + if "gate.bias" not in out or "gate.weight" not in out: + raise ValueError( + f"{prefix} has no learned gate; pick a layer >= num_hash_layers" + ) + return out + + +def _iter_layer_tensors( + snapshot_dir: str, prefix: str +) -> Iterable[Tuple[str, torch.Tensor]]: + from safetensors.torch import safe_open + + index = json.load(open(os.path.join(snapshot_dir, "model.safetensors.index.json"))) + weight_map = index["weight_map"] + shards = sorted({v for k, v in weight_map.items() if k.startswith(prefix)}) + if not shards: + raise KeyError(f"no tensors with prefix {prefix!r} in {snapshot_dir}") + for shard in shards: + with safe_open( + os.path.join(snapshot_dir, shard), framework="pt", device="cpu" + ) as h: + for name in h.keys(): + if name.startswith(prefix): + yield name, h.get_tensor(name) + + +def load_target_moe_layer(snapshot_dir: str, layer_id: int) -> Dict[str, torch.Tensor]: + """Dequantized ``layers.{layer_id}.ffn.*`` of a DeepSeek-V4 checkpoint dir.""" + prefix = f"layers.{layer_id}.ffn." + return dequantize_ffn_tensors( + dict(_iter_layer_tensors(snapshot_dir, prefix)), prefix + ) diff --git a/tests/test_modeling/test_moe_deepseek_v4.py b/tests/test_modeling/test_moe_deepseek_v4.py index cde7b0649..004fa2c16 100644 --- a/tests/test_modeling/test_moe_deepseek_v4.py +++ b/tests/test_modeling/test_moe_deepseek_v4.py @@ -439,6 +439,61 @@ def test_hf_export_reloads_and_carries_serving_config(self): self.assertTrue(torch.equal(value.float(), fresh[key].float()), key) +class TestDeepseekV4TargetDequant(unittest.TestCase): + def test_fp4_and_fp8_dequant_conventions(self): + from specforge.modeling.draft.moe.deepseek_v4_target import ( + FP4_TABLE, + dequant_fp4_packed, + dequant_fp8_block, + dequantize_ffn_tensors, + ) + + # one row of 32 fp4 values: nibbles 0..15 twice; low nibble = even index + codes = torch.arange(16, dtype=torch.uint8) + packed = ( + (codes | (codes << 4)).repeat(2).view(1, 32).to(torch.int8) + ) # 64 values + scale = torch.tensor([[2.0, 0.5]], dtype=torch.float32) # two groups of 32 + out = dequant_fp4_packed(packed, scale) + self.assertEqual(tuple(out.shape), (1, 64)) + expected = FP4_TABLE[codes.long()].repeat_interleave(2).repeat(2) + expected[:32] *= 2.0 + expected[32:] *= 0.5 + self.assertTrue(torch.equal(out.float()[0], expected)) + w8 = torch.full((128, 256), 1.0).to(torch.float8_e4m3fn) + s8 = torch.tensor([[1.0, 4.0]]) + d8 = dequant_fp8_block(w8, s8).float() + self.assertTrue((d8[:, :128] == 1.0).all() and (d8[:, 128:] == 4.0).all()) + raw = { + "layers.3.ffn.gate.weight": torch.randn(4, 8), + "layers.3.ffn.gate.bias": torch.randn(4), + "layers.3.ffn.experts.0.w1.weight": packed.repeat(2, 1), + "layers.3.ffn.experts.0.w1.scale": scale.repeat(2, 1), + "layers.3.ffn.shared_experts.w1.weight": w8, + "layers.3.ffn.shared_experts.w1.scale": s8, + } + rel = dequantize_ffn_tensors(raw, "layers.3.ffn.") + self.assertEqual( + set(rel), + { + "gate.weight", + "gate.bias", + "experts.0.w1.weight", + "shared_experts.w1.weight", + }, + ) + self.assertEqual(rel["gate.bias"].dtype, torch.float32) + self.assertEqual(rel["experts.0.w1.weight"].dtype, torch.bfloat16) + with self.assertRaisesRegex(ValueError, "hash-routed"): + dequantize_ffn_tensors( + { + "layers.1.ffn.gate.tid2eid": torch.zeros(4), + "layers.1.ffn.gate.weight": torch.zeros(4, 8), + }, + "layers.1.ffn.", + ) + + class TestDFlashIntegration(unittest.TestCase): def _forward(self, model): return model( From 174b59911917bb38bb8b3f91d6934cab7f46686b Mon Sep 17 00:00:00 2001 From: Kan Wu Date: Fri, 4 Sep 2026 07:21:20 +0000 Subject: [PATCH 09/11] fix: memory-lean, rank-staggered warm start for large MoE drafters - AutoDraftModel.from_pretrained (MoE path): honor output_loading_info, which the warm-start loader passes; read safetensors memory-mapped and drop the regrouped state once loaded. - warm start: for drafts with a native module layout, read and regroup the files directly instead of materializing a second full model per rank, and let ranks take turns loading. Eight ranks each building a 65 GB CPU model and then loading another 65 GB state + model exceeded the node's RAM and got the trainer OOM-killed on the 256-expert warm-started run. - Test: warm_start_draft_model round-trips a tiny MoE export. Co-Authored-By: Claude Fable 5.1 --- specforge/modeling/auto.py | 29 +++++-- specforge/training/model_loading.py | 86 +++++++++++++++------ tests/test_modeling/test_moe_deepseek_v4.py | 12 +++ 3 files changed, 97 insertions(+), 30 deletions(-) diff --git a/specforge/modeling/auto.py b/specforge/modeling/auto.py index 686b2b7e8..8b52a0f45 100644 --- a/specforge/modeling/auto.py +++ b/specforge/modeling/auto.py @@ -62,14 +62,16 @@ def filtered_warning(msg): config = AutoConfig.from_pretrained(pretrained_model_name_or_path) model_cls = cls._model_cls_from_config(config) kwargs = {**kwargs, "config": config} - state_dict = _native_state_dict_for(config, pretrained_model_name_or_path) + state_dict = load_native_state_dict(config, pretrained_model_name_or_path) if state_dict is not None: # HF from_pretrained assigns tensors by key and refuses an # explicit state_dict alongside a path, so build the module and # load the converted state ourselves. torch_dtype = kwargs.pop("torch_dtype", kwargs.pop("dtype", None)) + output_loading_info = bool(kwargs.pop("output_loading_info", False)) model = model_cls._from_config(config, torch_dtype=torch_dtype) result = model.load_state_dict(state_dict, strict=False) + del state_dict missing = [k for k in result.missing_keys if "embed_tokens" not in k] if missing or result.unexpected_keys: raise ValueError( @@ -77,7 +79,15 @@ def filtered_warning(msg): f"{model_cls.__name__}: missing {missing[:5]}, " f"unexpected {list(result.unexpected_keys)[:5]}" ) - return model.eval() + model.eval() + if output_loading_info: + return model, { + "missing_keys": list(result.missing_keys), + "unexpected_keys": list(result.unexpected_keys), + "mismatched_keys": [], + "error_msgs": [], + } + return model model = model_cls.from_pretrained( pretrained_model_name_or_path, *model_args, **kwargs ) @@ -87,11 +97,13 @@ def filtered_warning(msg): return model -def _native_state_dict_for(config, pretrained_model_name_or_path): +def load_native_state_dict(config, pretrained_model_name_or_path): """Checkpoint files use the official parameter naming; modules may use a native layout (MoE experts). HF ``from_pretrained`` assigns tensors by key - and cannot regroup them, so read the files and convert at this boundary. - Returns ``None`` when no conversion is needed (dense drafts).""" + and cannot regroup them, so read the files (memory-mapped) and convert at + this boundary. Returns ``None`` when no conversion is needed (dense + drafts). Callers that only need the tensors (warm start) use this directly + instead of materializing a second model.""" from specforge.modeling.draft.moe import ( from_checkpoint_state_dict, is_moe_config, @@ -101,7 +113,7 @@ def _native_state_dict_for(config, pretrained_model_name_or_path): return None import glob - from safetensors.torch import load_file + from safetensors import safe_open path = str(pretrained_model_name_or_path) if not os.path.isdir(path): @@ -113,7 +125,10 @@ def _native_state_dict_for(config, pretrained_model_name_or_path): raise FileNotFoundError(f"no safetensors weights under {path!r}") state = {} for file in files: - state.update(load_file(file)) + # mmap-backed views: only the regrouped tensors are materialized. + with safe_open(file, framework="pt", device="cpu") as handle: + for key in handle.keys(): + state[key] = handle.get_tensor(key) return from_checkpoint_state_dict(state) diff --git a/specforge/training/model_loading.py b/specforge/training/model_loading.py index 6ad8ee6d9..38286ac83 100644 --- a/specforge/training/model_loading.py +++ b/specforge/training/model_loading.py @@ -379,7 +379,15 @@ def _load_pretrained_draft_state( cache_dir: Optional[str], trust_remote_code: bool, ) -> Dict[str, Any]: - from specforge.modeling.auto import AutoDraftModel + from specforge.modeling.auto import AutoDraftModel, load_native_state_dict + + # Drafts with a native module layout (MoE experts): read and regroup the + # files directly. Building a full second model per rank doubled the CPU + # footprint (~65 GB each for a 256-expert drafter) and OOM-killed 8-rank + # trainers during warm start. + native = load_native_state_dict(draft_config, source) + if native is not None: + return {key: value.detach().cpu() for key, value in native.items()} loaded, loading_info = AutoDraftModel.from_pretrained( source, @@ -398,6 +406,25 @@ def _load_pretrained_draft_state( return state +def _rank_staggered(fn): + """Run ``fn`` on one distributed rank at a time (all ranks call this).""" + try: + import torch.distributed as dist + + active = dist.is_available() and dist.is_initialized() + except ImportError: # pragma: no cover + active = False + if not active or dist.get_world_size() == 1: + return fn() + rank, world = dist.get_rank(), dist.get_world_size() + result = None + for turn in range(world): + if turn == rank: + result = fn() + dist.barrier() + return result + + def warm_start_draft_model( model: Any, source: str, @@ -411,30 +438,43 @@ def warm_start_draft_model( """Load only draft weights, never optimizer/counters/RNG training state.""" runtime_state = _runtime_state_file(source) - if runtime_state is not None: - checkpoint_format: Literal["specforge", "pretrained"] = "specforge" - state = _load_specforge_draft_state(runtime_state, expected_strategy=strategy) - else: - checkpoint_format = "pretrained" - state = _load_pretrained_draft_state( - source, - draft_config=draft_config, - cache_dir=cache_dir, - trust_remote_code=trust_remote_code, - ) + checkpoint_format: Literal["specforge", "pretrained"] = ( + "specforge" if runtime_state is not None else "pretrained" + ) - if not state: - raise ValueError(f"warm-start checkpoint {source!r} contains no draft weights") - try: - # Files use the official naming; modules may use a native MoE layout. - from specforge.modeling.draft.moe import from_checkpoint_state_dict + def _load() -> Tuple[Dict[str, Any], Any]: + if runtime_state is not None: + state = _load_specforge_draft_state( + runtime_state, expected_strategy=strategy + ) + else: + state = _load_pretrained_draft_state( + source, + draft_config=draft_config, + cache_dir=cache_dir, + trust_remote_code=trust_remote_code, + ) + if not state: + raise ValueError( + f"warm-start checkpoint {source!r} contains no draft weights" + ) + try: + # Files use the official naming; modules may use a native MoE layout. + from specforge.modeling.draft.moe import from_checkpoint_state_dict - result = model.load_state_dict(from_checkpoint_state_dict(state), strict=False) - except RuntimeError as exc: - raise ValueError( - f"warm-start checkpoint {source!r} has incompatible draft tensor " - f"shapes: {exc}" - ) from exc + result = model.load_state_dict( + from_checkpoint_state_dict(state), strict=False + ) + except RuntimeError as exc: + raise ValueError( + f"warm-start checkpoint {source!r} has incompatible draft tensor " + f"shapes: {exc}" + ) from exc + return state, result + + # Every rank loads the full draft state; for large drafts the transient + # host memory of N simultaneous loads can exceed the node, so take turns. + state, result = _rank_staggered(_load) loaded_keys = len(state) - len(result.unexpected_keys) if result.unexpected_keys or loaded_keys == 0: raise ValueError( diff --git a/tests/test_modeling/test_moe_deepseek_v4.py b/tests/test_modeling/test_moe_deepseek_v4.py index 004fa2c16..f701a3467 100644 --- a/tests/test_modeling/test_moe_deepseek_v4.py +++ b/tests/test_modeling/test_moe_deepseek_v4.py @@ -437,6 +437,18 @@ def test_hf_export_reloads_and_carries_serving_config(self): fresh = reloaded.state_dict() for key, value in model.state_dict().items(): self.assertTrue(torch.equal(value.float(), fresh[key].float()), key) + # the trainer's weights-only warm start reads the same directory + from specforge.training.model_loading import warm_start_draft_model + + target = DFlashDraftModel(_dflash_config()).to(torch.bfloat16) + report = warm_start_draft_model( + target, out, draft_config=config, strategy="dflash" + ) + self.assertEqual(report.checkpoint_format, "pretrained") + for key, value in model.state_dict().items(): + self.assertTrue( + torch.equal(value.float(), target.state_dict()[key].float()), key + ) class TestDeepseekV4TargetDequant(unittest.TestCase): From fd0745e3630b3daf0cefcb7dd8d22afedc6a22cf Mon Sep 17 00:00:00 2001 From: Kan Wu Date: Fri, 4 Sep 2026 08:22:10 +0000 Subject: [PATCH 10/11] feat: freeze warm-started MoE experts and keep them replicated under FSDP dflash_config.moe_freeze_experts keeps the routed experts fixed (router, shared expert and the rest of the draft still train). RoutedExperts opts into FSDP replication when fully frozen, and the backend's frozen-module scan now honors that opt-in next to lm_head/embed_tokens. Why: with DeepSeek-V4-Flash's 256x2048 experts (32.7B params) sharded, every micro-batch re-gathered ~65 GB of weights and reduce-scattered their grads (~3 TB/step, 46 s/step on 8 B200s), and at the drafter's LR (6e-4) AdamW would have overwritten the warm-started experts within ~100 steps anyway. Frozen + replicated: no expert communication, dense-like step time, and the pretrained experts are what gets served. Co-Authored-By: Claude Fable 5.1 --- specforge/modeling/draft/moe/config.py | 6 ++++ specforge/modeling/draft/moe/experts.py | 4 +++ specforge/modeling/draft/moe/layer.py | 2 ++ specforge/training/backend.py | 18 +++++++++--- tests/test_modeling/test_moe_deepseek_v4.py | 31 +++++++++++++++++++++ 5 files changed, 57 insertions(+), 4 deletions(-) diff --git a/specforge/modeling/draft/moe/config.py b/specforge/modeling/draft/moe/config.py index 66fb2c761..ee6c58d87 100644 --- a/specforge/modeling/draft/moe/config.py +++ b/specforge/modeling/draft/moe/config.py @@ -49,6 +49,7 @@ "moe_bias_update_rate": "bias_update_rate", "moe_aux_loss_coeff": "aux_loss_coeff", "moe_dispatch": "dispatch", + "moe_freeze_experts": "freeze_experts", } @@ -80,6 +81,11 @@ class MoEConfig: bias_update_rate: float = 0.0 aux_loss_coeff: float = 0.0 dispatch: str = "sorted_loop" + #: Keep the routed experts fixed (e.g. warm-started from the target) and + #: train only the router, shared expert and the rest of the draft. Frozen + #: experts are replicated by the FSDP backend instead of sharded, which + #: removes the per-micro-batch weight all-gathers that dominate large MoEs. + freeze_experts: bool = False def __post_init__(self) -> None: if self.n_routed_experts <= 0: diff --git a/specforge/modeling/draft/moe/experts.py b/specforge/modeling/draft/moe/experts.py index c5d0f3f94..4d29cacee 100644 --- a/specforge/modeling/draft/moe/experts.py +++ b/specforge/modeling/draft/moe/experts.py @@ -24,6 +24,10 @@ class RoutedExperts(nn.Module, abc.ABC): + #: The training backend keeps a fully frozen instance replicated (outside + #: FSDP sharding): no weight all-gathers or gradient reduce-scatters. + fsdp_replicate_when_frozen = True + def __init__(self, cfg: MoEConfig, hidden_size: int) -> None: super().__init__() self.cfg = cfg diff --git a/specforge/modeling/draft/moe/layer.py b/specforge/modeling/draft/moe/layer.py index 612ca8c1b..9857d98dc 100644 --- a/specforge/modeling/draft/moe/layer.py +++ b/specforge/modeling/draft/moe/layer.py @@ -30,6 +30,8 @@ def __init__(self, cfg: MoEConfig, hidden_size: int) -> None: balance = build_balance_controller(cfg, cfg.n_routed_experts) self.gate = build_router(cfg, hidden_size, balance) self.experts = build_routed_experts(cfg, hidden_size) + if cfg.freeze_experts: + self.experts.requires_grad_(False) self.shared_experts: Optional[nn.Module] = ( build_shared_expert(cfg, hidden_size) if cfg.n_shared_experts else None ) diff --git a/specforge/training/backend.py b/specforge/training/backend.py index e55f2fad3..c488883aa 100644 --- a/specforge/training/backend.py +++ b/specforge/training/backend.py @@ -209,11 +209,21 @@ def _frozen_target_modules(model: nn.Module) -> tuple[nn.Module, ...]: all-gather them before every optimizer window without saving optimizer memory, which is the wrong trade-off for the current trainer recipes. """ + candidates = [ + module + for name in ("lm_head", "embed_tokens") + if isinstance(module := getattr(model, name, None), nn.Module) + ] + # Submodules that opt in (e.g. frozen MoE experts warm-started from the + # target) are replicated as well: sharding tens of GB of frozen weights + # would re-gather them on every micro-batch for no optimizer savings. + candidates += [ + module + for module in model.modules() + if getattr(module, "fsdp_replicate_when_frozen", False) + ] modules = [] - for name in ("lm_head", "embed_tokens"): - module = getattr(model, name, None) - if not isinstance(module, nn.Module): - continue + for module in candidates: parameters = tuple(module.parameters()) if parameters and not any( parameter.requires_grad for parameter in parameters diff --git a/tests/test_modeling/test_moe_deepseek_v4.py b/tests/test_modeling/test_moe_deepseek_v4.py index f701a3467..e26f28c4c 100644 --- a/tests/test_modeling/test_moe_deepseek_v4.py +++ b/tests/test_modeling/test_moe_deepseek_v4.py @@ -451,6 +451,37 @@ def test_hf_export_reloads_and_carries_serving_config(self): ) +class TestFrozenExperts(unittest.TestCase): + def test_freeze_experts_trains_router_and_shared_only(self): + layer = _layer(dflash_config={"moe_freeze_experts": True}).train() + self.assertFalse(any(p.requires_grad for p in layer.experts.parameters())) + self.assertTrue(layer.gate.weight.requires_grad) + self.assertTrue(all(p.requires_grad for p in layer.shared_experts.parameters())) + y = layer(torch.randn(6, 32, requires_grad=True)) + y.float().sum().backward() + self.assertIsNotNone(layer.gate.weight.grad) + self.assertIsNone(layer.experts.w1.grad) + + def test_backend_replicates_frozen_experts(self): + from specforge.training.backend import FSDPTrainingBackend + + config = _dflash_config() + config.dflash_config = {**config.dflash_config, "moe_freeze_experts": True} + model = DFlashDraftModel(config) + ignored = FSDPTrainingBackend._frozen_target_modules(model) + self.assertEqual([type(m).__name__ for m in ignored], ["GroupedExperts"] * 2) + trainable = sum(p.numel() for p in model.parameters() if p.requires_grad) + frozen = sum(p.numel() for m in ignored for p in m.parameters()) + self.assertGreater(frozen, trainable) + # trained experts stay sharded + self.assertEqual( + FSDPTrainingBackend._frozen_target_modules( + DFlashDraftModel(_dflash_config()) + ), + (), + ) + + class TestDeepseekV4TargetDequant(unittest.TestCase): def test_fp4_and_fp8_dequant_conventions(self): from specforge.modeling.draft.moe.deepseek_v4_target import ( From 12c3f9d27de01e5bbc7f43c023f9489e99950c0b Mon Sep 17 00:00:00 2001 From: Yusheng Su Date: Mon, 5 Oct 2026 17:26:43 +0900 Subject: [PATCH 11/11] feat(moe): expert parallelism for DFlash-family MoE drafters on the FSDP2 backend Add training.expert_parallel_size (FSDP2 only). Each MoE layer's routed experts are sliced across that many consecutive ranks as DTensor Shard(0) over an `ep` mesh axis; FSDP2 shards the slice again over the `efsdp` ranks through a per-parameter shard_placement_fn while dense parameters stay on the full data-parallel mesh. The EP axis is carved out of data parallelism, so ranks keep distinct micro-batches: inside the layer the group all-gathers its tokens, every rank routes them with the replicated router and computes only the experts it owns, and a reduce-scatter returns each rank its own tokens summed over all owners (fixed-size collectives, one host sync per layer for this rank's slot bounds). Because every parameter element lives on exactly one rank, the optimizer's local-shard grad norm, per-shard Adam state and the DCP full-state-dict checkpoint path are unchanged; the full [E, ...] tensors keep the official per-expert naming. FSDP2 averages expert gradients over efsdp only while the owner already summed its group's tokens, so the backend rescales them by 1/ep at the optimizer boundary. A rank whose experts receive no token keeps the gathered input, combine weights and expert parameters on the graph through zero-valued terms so every rank issues the same collectives and FSDP2 sees the same gradient set. init_distributed builds the (efsdp, ep) mesh; ParallelConfig carries it; the FSDP1 backend and NO_SHARD reject EP. Adds the DeepSeek-V4-Flash DSpark MoE EP=4 recipe and docs. Tests (CPU gloo): 2-rank layer parity incl. a starved rank, 4-rank FSDP2TrainingBackend with ep=2 and ep=4 against a data-parallel reference, checkpoint naming and reload; schema validation. --- .../deepseek-v4-flash-dspark-disaggregated.md | 11 + .../advanced_features/customization.md | 28 ++ docs/sections/basic_usage/training.md | 9 + examples/configs/README.md | 6 + ...v4-flash-dspark-moe-ep4-disaggregated.yaml | 101 ++++++ specforge/cli.py | 1 + specforge/config/schema.py | 27 ++ specforge/distributed.py | 43 ++- specforge/modeling/draft/moe/DESIGN.md | 15 + specforge/modeling/draft/moe/__init__.py | 5 + .../modeling/draft/moe/expert_parallel.py | 172 +++++++++ .../modeling/draft/moe/grouped_experts.py | 129 +++++-- specforge/modeling/draft/moe/hooks.py | 13 + specforge/modeling/draft/moe/layer.py | 35 +- specforge/training/backend.py | 17 + specforge/training/fsdp2.py | 94 ++++- tests/test_config/test_schema.py | 42 +++ .../test_modeling/test_moe_expert_parallel.py | 336 ++++++++++++++++++ 18 files changed, 1052 insertions(+), 32 deletions(-) create mode 100644 examples/configs/online/disaggregated/external/deepseek-v4-flash-dspark-moe-ep4-disaggregated.yaml create mode 100644 specforge/modeling/draft/moe/expert_parallel.py create mode 100644 tests/test_modeling/test_moe_expert_parallel.py diff --git a/docs/recipes/deepseek-v4-flash-dspark-disaggregated.md b/docs/recipes/deepseek-v4-flash-dspark-disaggregated.md index 8876e4485..3915132c6 100644 --- a/docs/recipes/deepseek-v4-flash-dspark-disaggregated.md +++ b/docs/recipes/deepseek-v4-flash-dspark-disaggregated.md @@ -127,6 +127,17 @@ per-expert naming (`layers.N.mlp.experts.{i}.w{1,2,3}.weight`, `layers.N.mlp.gate.bias`, `layers.N.mlp.shared_experts.w{1,2,3}.weight`), so exports load into SGLang's DeepSeek-V4 MoE unchanged. +To train the experts instead of freezing warm-started ones, use the +expert-parallel variant +`examples/configs/online/disaggregated/external/deepseek-v4-flash-dspark-moe-ep4-disaggregated.yaml` +(`training.backend: fsdp2`, `training.expert_parallel_size: 4`). With the +experts sharded as plain FSDP parameters, every micro-batch re-gathers all 64 +experts of every layer; under EP each of the four trainer ranks owns 16 experts +per layer, the group exchanges its anchor tokens (fixed-size all-gather and +reduce-scatter per MoE layer) and the expert weights never move. Checkpoints, +`moe/*` metrics, warm start and exports are unchanged. See the expert +parallelism notes in `docs/sections/advanced_features/customization.md`. + ## Fresh attempts Delete the run's `outputs/` directory and, whenever a capture server was diff --git a/docs/sections/advanced_features/customization.md b/docs/sections/advanced_features/customization.md index a2ae5cbc1..e051a821f 100644 --- a/docs/sections/advanced_features/customization.md +++ b/docs/sections/advanced_features/customization.md @@ -164,6 +164,34 @@ preset registration plus whichever components it needs (score function, balance controller, experts backend, shared expert); each registers by name from its own module. +### Expert parallelism + +Experts that are trained (not `moe_freeze_experts`) can be sliced across ranks +on the FSDP2 backend: + +```yaml +training: + backend: fsdp2 + expert_parallel_size: 4 +``` + +Each MoE layer's routed experts are split across `expert_parallel_size` +consecutive ranks (the EP group). Inside the layer the group all-gathers its +tokens, every rank routes them with the replicated router and computes only the +experts it owns, and a reduce-scatter returns each rank its own tokens summed +over all owners. The expert slices are `DTensor`s over the `ep` mesh axis and +FSDP2 shards them again over the ranks that hold the same slice (`efsdp`); +everything else stays on the full data-parallel mesh. The EP axis is carved out +of data parallelism, so ranks keep distinct data and the dense part of the draft +is not computed twice. Checkpoints still gather to the full `[E, ...]` tensors +and keep the official naming; the balance controller, warm start and exports +are unchanged. Requirements: a MoE draft JSON, `backend: fsdp2`, a sharded +`fsdp_sharding`, `tp_size: 1`, no sequence parallelism, and `n_routed_experts` +divisible by `expert_parallel_size`. Expert parallelism adds one host sync per +MoE layer (this rank's slot bounds in the sorted routing) and two fixed-size +collectives; it pays off when the gathered expert weights, not the tokens, +dominate the step. + ## Draft architectures Draft classes register through `@register_draft`. The key defaults to the diff --git a/docs/sections/basic_usage/training.md b/docs/sections/basic_usage/training.md index 8221f9266..e5e25a9da 100644 --- a/docs/sections/basic_usage/training.md +++ b/docs/sections/basic_usage/training.md @@ -511,6 +511,15 @@ foundation for future DTensor-based parallelism. Throughput and peak memory still depend on the model and sharding policy; selecting FSDP2 alone does not guarantee a speedup or enable tensor parallelism. +MoE drafts can add expert parallelism on FSDP2 with +`training.expert_parallel_size`: each MoE layer's routed experts are sliced +across that many consecutive ranks, the group exchanges its tokens around the +expert computation, and FSDP2 shards the slice over the ranks that own the same +experts while dense parameters stay on the full mesh. It requires a sharded +`fsdp_sharding`, `tp_size: 1`, no sequence parallelism, and a world size and +`n_routed_experts` divisible by it; see the MoE section of the customization +guide. + The launcher creates every process group from the typed run config: - Online target TP/EP belongs to each external SGLang capture server, not the diff --git a/examples/configs/README.md b/examples/configs/README.md index 465e5598f..ee2d6ae30 100644 --- a/examples/configs/README.md +++ b/examples/configs/README.md @@ -93,6 +93,11 @@ drafter-architecture ablation: the same recipe with `configs/deepseek-v4-flash-dspark-moe.json`, whose `moe_preset: deepseek_v4` swaps the dense MLP for the target's routing (64 routed + 1 shared experts, top-6, width 2048); see the runbook's MoE section. +`deepseek-v4-flash-dspark-moe-ep4-disaggregated.yaml` is that arm with the +experts trained under expert parallelism on the FSDP2 backend +(`training.backend: fsdp2`, `training.expert_parallel_size: 4`): each of the +four trainer ranks owns 16 of the 64 experts instead of re-gathering all of +them every micro-batch. `qwen3.8-27b-dflash2-disaggregated.yaml` (external services, two nodes) and its managed-local siblings `qwen3.8-27b-dflash2-4server-dp4-disaggregated.yaml` @@ -294,6 +299,7 @@ Common fields: | `training.tp_size` | `1` | Online disaggregated consumers must keep it at 1; configure target TP on capture servers. Offline non-USP ranks consume disjoint data. | | `training.sp_ulysses_size` | `1` | Ulysses sequence-parallel factor for offline EAGLE3 USP. | | `training.sp_ring_size` | `1` | Ring sequence-parallel factor for offline EAGLE3 USP. | +| `training.expert_parallel_size` | `1` | Expert parallelism for MoE drafts on the FSDP2 backend: each MoE layer's routed experts are sliced across this many consecutive ranks (the group all-gathers its tokens and reduce-scatters the expert outputs) and FSDP2 shards the slice over the remaining ranks. Requires `backend: fsdp2`, a sharded `fsdp_sharding`, `tp_size: 1`, no sequence parallelism, a world size divisible by it, and `n_routed_experts` divisible by it. | | `training.dist_timeout` | `10` | Positive distributed-operation timeout in minutes. | | `training.save_interval` | `0` | Save every N optimizer steps; 0 disables periodic saves. A final checkpoint is still written. | | `training.eval_interval` | `0` | Evaluate every N optimizer steps; 0 disables evaluation. | diff --git a/examples/configs/online/disaggregated/external/deepseek-v4-flash-dspark-moe-ep4-disaggregated.yaml b/examples/configs/online/disaggregated/external/deepseek-v4-flash-dspark-moe-ep4-disaggregated.yaml new file mode 100644 index 000000000..ca1576b1c --- /dev/null +++ b/examples/configs/online/disaggregated/external/deepseek-v4-flash-dspark-moe-ep4-disaggregated.yaml @@ -0,0 +1,101 @@ +# Expert-parallel MoE-FFN arm of the DSpark drafter ablation for +# DeepSeek-V4-Flash: deepseek-v4-flash-dspark-moe-disaggregated.yaml with the +# experts TRAINED under expert parallelism on the FSDP2 backend. The four +# trainer ranks form one EP group: each owns 16 of the 64 routed experts per +# layer, the group all-gathers its anchor tokens and reduce-scatters the +# expert outputs, and the expert weights are never re-gathered. Everything else +# (capture servers, data, objective, checkpoints) is the MoE arm unchanged. +# Requires n_routed_experts divisible by expert_parallel_size and a trainer +# world size divisible by it (nproc_per_node: 4 here). +model: + target_model_path: deepseek-ai/DeepSeek-V4-Flash-0731 + draft_model_config: configs/deepseek-v4-flash-dspark-moe.json + target_backend: sglang + trust_remote_code: true + # DeepSeek-V4 checkpoints use the native inference weight layout. + embedding_key: embed.weight + lm_head_key: head.weight + mask_token_id: 128799 + torch_dtype: bfloat16 + sglang_mem_fraction_static: 0.85 + # data.max_length plus headroom: capture rejects inputs at exactly context length. + sglang_context_length: 8704 + sglang_max_running_requests: 8 + # Routed experts are fp4; the default MoE path cannot run them on B200. + sglang_moe_runner_backend: flashinfer_mxfp4 + +data: + train_data_path: ./cache/dataset/sharegpt_train.jsonl + max_length: 8192 + chat_template: deepseek-v4 + cache_dir: cache + build_dataset_num_proc: 64 + # Async prefetch: overlaps per-sample Mooncake feature fetches with compute. + dataloader_num_workers: 8 + +training: + strategy: dspark + backend: fsdp2 + expert_parallel_size: 4 + num_epochs: 2 + # 4 ranks x 32 microbatches -> global batch 128. + batch_size: 1 + accumulation_steps: 32 + learning_rate: 0.0006 + lr_scheduler: constant + warmup_ratio: 0 + max_grad_norm: 1 + attention_backend: flex_attention + num_anchors: 512 + loss_decay_gamma: 4.0 + objective_chunk_blocks: 128 + dspark_ce_loss_alpha: 0.1 + dspark_l1_loss_alpha: 0.9 + dspark_confidence_head_alpha: 1.0 + save_interval: 128 + log_interval: 10 + dist_timeout: 30 + seed: 42 + prompt_seed: 1 + +tracking: + report_to: wandb + wandb_project: specforge + wandb_name: deepseek-v4-flash-dspark-moe-ep4-disaggregated + wandb_dir: outputs/deepseek-v4-flash-dspark-moe-ep4-disaggregated/wandb + +runtime: + producer_lease: 8 + producer_concurrency: 8 + # Keep two 128-sample optimizer quanta in flight (~0.4 GiB features/sample); + # the low watermark is one full quantum so complete windows always dispatch. + in_flight_high_watermark: 256 + in_flight_low_watermark: 128 + resident_high_watermark_bytes: 137438953472 + resident_low_watermark_bytes: 103079215104 + feature_store_max_resident_bytes: 171798691840 + +run_id: deepseek-v4-flash-dspark-moe-ep4-disaggregated +output_dir: outputs/deepseek-v4-flash-dspark-moe-ep4-disaggregated + +deployment: + mode: disaggregated + trainer: + nnodes: 1 + nproc_per_node: 4 + disaggregated: + control_dir: outputs/deepseek-v4-flash-dspark-moe-ep4-disaggregated/control + consumer_state_dir: outputs/deepseek-v4-flash-dspark-moe-ep4-disaggregated/consumer-state + backend: mooncake + store_id: deepseek-v4-flash-dspark-moe-ep4-disaggregated + # Two TP2 servers out-produce one TP4 server (TP prefill scaling is sublinear). + server_urls: + - http://127.0.0.1:30000 + - http://127.0.0.1:30001 + mooncake_metadata_server: http://127.0.0.1:35880/metadata + mooncake_master_server_addr: 127.0.0.1:35551 + mooncake_local_hostname: 127.0.0.1 + mooncake_protocol: tcp + client_buffer_size: 1073741824 + idle_timeout_s: 7200 + peer_wait_timeout_s: 7200 diff --git a/specforge/cli.py b/specforge/cli.py index 89677480e..352d45920 100644 --- a/specforge/cli.py +++ b/specforge/cli.py @@ -137,6 +137,7 @@ def _train(resolved) -> int: tp_size=cfg.training.tp_size, sp_ulysses_size=cfg.training.sp_ulysses_size, sp_ring_size=cfg.training.sp_ring_size, + expert_parallel_size=cfg.training.expert_parallel_size, ) failed = True try: diff --git a/specforge/config/schema.py b/specforge/config/schema.py index f8728bf32..ed42dcb5e 100644 --- a/specforge/config/schema.py +++ b/specforge/config/schema.py @@ -917,6 +917,11 @@ class TrainingConfig(StrictConfigModel): tp_size: int = Field(default=1, gt=0) sp_ulysses_size: int = Field(default=1, gt=0) sp_ring_size: int = Field(default=1, gt=0) + #: Expert parallelism for MoE drafts on the FSDP2 backend: each MoE layer's + #: routed experts are sliced across this many consecutive ranks and FSDP2 + #: shards the slice over the remaining ranks. Requires a MoE draft JSON, + #: ``backend: fsdp2`` and a sharded ``fsdp_sharding``. + expert_parallel_size: int = Field(default=1, gt=0) dist_timeout: int = Field(default=10, gt=0) #: Acceptance-aware token objective. DFlash-family hard targets make #: ``alpha`` equivalent to CE; ``lambda`` mixes CE and TV. @@ -1016,6 +1021,22 @@ def _validate_training_shape(self): "training.sp_ulysses_size/sp_ring_size require " "training.attention_backend=usp" ) + if self.expert_parallel_size > 1: + if self.backend != "fsdp2": + raise ValueError( + "training.expert_parallel_size > 1 requires training.backend=fsdp2" + ) + if self.fsdp_sharding == "NO_SHARD": + raise ValueError( + "training.expert_parallel_size > 1 requires a sharded " + "training.fsdp_sharding (SHARD_GRAD_OP or FULL_SHARD), " + "not NO_SHARD" + ) + if self.tp_size != 1 or sp_size != 1: + raise ValueError( + "training.expert_parallel_size > 1 currently requires " + "training.tp_size=1 and no sequence parallelism" + ) return self @@ -1351,6 +1372,12 @@ def validate_world_size(self, world_size: int) -> None: f"parallel size {sp_size} " "(sp_ulysses_size * sp_ring_size)" ) + ep_size = self.training.expert_parallel_size + if world_size % ep_size: + raise ValueError( + f"world_size={world_size} must be divisible by " + f"training.expert_parallel_size={ep_size}" + ) @classmethod def from_file(cls, path: str) -> "Config": diff --git a/specforge/distributed.py b/specforge/distributed.py index 44acb80f6..42ed0fad0 100644 --- a/specforge/distributed.py +++ b/specforge/distributed.py @@ -14,6 +14,7 @@ _DP_GROUP = None _DRAFT_DP_GROUP = None _DRAFT_SP_GROUP = None +_DRAFT_EP_MESH = None _SP_ULYSSES_GROUP = None _SP_RING_GROUP = None @@ -131,14 +132,32 @@ def get_sp_ring_group(): return _SP_RING_GROUP +def get_draft_ep_mesh(): + """2-D ``(efsdp, ep)`` mesh of the draft's expert parallelism, or ``None``.""" + global _DRAFT_EP_MESH + return _DRAFT_EP_MESH + + +def _draft_ep_groups(): + if _DRAFT_EP_MESH is None: + return () + return tuple(_DRAFT_EP_MESH.get_group(name) for name in ("ep", "efsdp")) + + def init_distributed( - timeout: int = 10, tp_size: int = 1, sp_ulysses_size: int = 1, sp_ring_size: int = 1 + timeout: int = 10, + tp_size: int = 1, + sp_ulysses_size: int = 1, + sp_ring_size: int = 1, + expert_parallel_size: int = 1, ): """Initialize distributed training. Args: timeout(int): Timeout for collective communication in minutes tp_size(int): The degree of tensor parallelism + expert_parallel_size(int): Ranks that share one MoE layer's experts + (each owns a disjoint slice); FSDP shards over the remaining ranks. """ device_type = get_device_type() backend = _distributed_backend(device_type) @@ -182,6 +201,21 @@ def init_distributed( if set_seq_parallel_pg is not None: set_seq_parallel_pg(sp_ulysses_size, sp_ring_size, dist.get_rank(), world_size) + draft_ep_mesh = None + if expert_parallel_size > 1: + assert world_size % expert_parallel_size == 0, ( + f"World size ({world_size}) cannot be evenly divided by " + f"expert_parallel_size ({expert_parallel_size})" + ) + # ``ep`` is the fast axis: an expert-parallel group is ``ep`` consecutive + # ranks (one node for ep <= GPUs per node), ``efsdp`` the ranks that + # own the same expert slice and shard it with FSDP2. + draft_ep_mesh = dist.device_mesh.init_device_mesh( + device_type, + (world_size // expert_parallel_size, expert_parallel_size), + mesh_dim_names=("efsdp", "ep"), + ) + print_with_rank(f"device mesh: {device_mesh}") tp_group = device_mesh.get_group("tp") dp_group = device_mesh.get_group("dp") @@ -199,7 +233,7 @@ def init_distributed( # we need to create a 1D submesh tp_device_mesh = dist.DeviceMesh.from_group(tp_group, device_type=device_type) - global _TP_GROUP, _DP_GROUP, _DEVICE_MESH, _TP_DEVICE_MESH, _DP_DEVICE_MESH, _SP_RING_GROUP, _SP_ULYSSES_GROUP, _DRAFT_DP_GROUP, _DRAFT_SP_GROUP + global _TP_GROUP, _DP_GROUP, _DEVICE_MESH, _TP_DEVICE_MESH, _DP_DEVICE_MESH, _SP_RING_GROUP, _SP_ULYSSES_GROUP, _DRAFT_DP_GROUP, _DRAFT_SP_GROUP, _DRAFT_EP_MESH _DEVICE_MESH = device_mesh _TP_GROUP = tp_group _TP_DEVICE_MESH = tp_device_mesh @@ -208,6 +242,7 @@ def init_distributed( _DP_GROUP = dp_group _DRAFT_DP_GROUP = draft_dp_group _DRAFT_SP_GROUP = draft_sp_group + _DRAFT_EP_MESH = draft_ep_mesh _DP_DEVICE_MESH = dist.DeviceMesh.from_group(dp_group, device_type=device_type) @@ -220,7 +255,7 @@ def destroy_distributed(*, abort: bool = False): """ global _DEVICE_MESH, _TP_DEVICE_MESH, _TP_GROUP global _DP_DEVICE_MESH, _DP_GROUP, _DRAFT_DP_GROUP, _DRAFT_SP_GROUP - global _SP_ULYSSES_GROUP, _SP_RING_GROUP + global _SP_ULYSSES_GROUP, _SP_RING_GROUP, _DRAFT_EP_MESH # Teardown must never crash the process. Several handles can alias the same # underlying group (e.g. DP and draft-DP when there is no sequence # parallelism), and degenerate single-rank SP groups (created when @@ -236,6 +271,7 @@ def destroy_distributed(*, abort: bool = False): _SP_RING_GROUP, _DRAFT_DP_GROUP, _DRAFT_SP_GROUP, + *_draft_ep_groups(), default_group, # may alias the DP group; the seen-set dedups it ): if group is None or id(group) in seen: @@ -262,6 +298,7 @@ def destroy_distributed(*, abort: bool = False): _DP_GROUP = None _DRAFT_DP_GROUP = None _DRAFT_SP_GROUP = None + _DRAFT_EP_MESH = None _SP_ULYSSES_GROUP = None _SP_RING_GROUP = None diff --git a/specforge/modeling/draft/moe/DESIGN.md b/specforge/modeling/draft/moe/DESIGN.md index 15bd32538..611076a4e 100644 --- a/specforge/modeling/draft/moe/DESIGN.md +++ b/specforge/modeling/draft/moe/DESIGN.md @@ -52,6 +52,8 @@ hooks.py apply_pending_balance_updates / collect_moe_aux_loss / collect_moe_metrics over any module tree. state_dict.py to/from_checkpoint_state_dict: module layout <-> official names. init.py WarmStartPlan: which target experts seed which draft experts. +expert_parallel.py EP layout, DTensor expert slicing, differentiable token + all-gather / output reduce-scatter seams used by MoELayer. ``` Implementations register into these registries at import time (imported at @@ -97,6 +99,19 @@ differently and raise `CheckpointError`. controller metrics; the DFlash/DSpark strategies add them to `StepOutput.metrics`, and the trainer DP-averages and logs them like any other scalar. +**Expert parallelism is a layout, not a different layer.** After +`MoELayer.apply_expert_parallel(ep_mesh)` the experts backend holds its +`[E/ep, ...]` slice as a `DTensor` (`Shard(0)` over `ep`) and the layer wraps the +forward in `gather_tokens` / `scatter_outputs`: the replicated router runs on +the group's gathered tokens, each rank computes its own experts for all of +them, and the fp32 partial outputs are reduce-scattered. The FSDP2 backend +applies this before wrapping and shards the slice further over `efsdp` through +a per-parameter placement; `get_model_state_dict(full_state_dict=True)` gathers +the full tensors, so the checkpoint boundary above is untouched. A rank whose +experts receive no token keeps the gathered input, the combine weights and its +expert parameters on the graph through zero-valued terms, so every rank issues +the same collectives and FSDP2 sees the same gradient set. + **Aux losses are collected, not yet consumed.** `collect_moe_aux_loss` sums scaled layer losses; wiring it into an objective is done with the first preset whose balancing policy emits one (aux-loss-free policies do not). diff --git a/specforge/modeling/draft/moe/__init__.py b/specforge/modeling/draft/moe/__init__.py index 3cd8ad691..de83887e5 100644 --- a/specforge/modeling/draft/moe/__init__.py +++ b/specforge/modeling/draft/moe/__init__.py @@ -24,6 +24,7 @@ - :mod:`.hooks` model-level plumbing: balance updates, aux loss, metrics - :mod:`.state_dict` module layout <-> official checkpoint naming boundary - :mod:`.init` warm-start plans from a target model's experts +- :mod:`.expert_parallel` expert slicing + token gather/scatter seams for EP Implementations register into the registries from their own modules (:mod:`.topk_router`, :mod:`.noaux_tc`, :mod:`.grouped_experts`, @@ -53,6 +54,7 @@ register_moe_preset, resolve_moe_config, ) +from .expert_parallel import ExpertParallelLayout from .experts import ( EXPERTS_BACKENDS, RoutedExperts, @@ -60,6 +62,7 @@ register_experts_backend, ) from .hooks import ( + apply_expert_parallel, apply_pending_balance_updates, collect_moe_aux_loss, collect_moe_metrics, @@ -95,6 +98,8 @@ ) __all__ = [ + "ExpertParallelLayout", + "apply_expert_parallel", "BALANCE_CONTROLLERS", "BalanceController", "EXPERTS_BACKENDS", diff --git a/specforge/modeling/draft/moe/expert_parallel.py b/specforge/modeling/draft/moe/expert_parallel.py new file mode 100644 index 000000000..1a847144b --- /dev/null +++ b/specforge/modeling/draft/moe/expert_parallel.py @@ -0,0 +1,172 @@ +# coding=utf-8 +"""Expert parallelism (EP) for the routed experts: layout and autograd seams. + +The EP axis is carved out of data parallelism, so every rank keeps its own +micro-batch and the dense part of the draft is never computed twice. Inside an +MoE layer the ``ep`` ranks of a group + +1. all-gather their tokens (:func:`gather_tokens`), +2. run the replicated router on the gathered tokens, +3. compute only the experts they own, for all gathered tokens, and +4. reduce-scatter the partial outputs (:func:`scatter_outputs`) so each rank + gets its own tokens back, summed over every expert owner. + +Each rank owns a disjoint slice of the stacked expert weights: the ``[E, ...]`` +parameters become ``DTensor`` ``Shard(0)`` over the 1-D ``ep`` mesh +(:func:`shard_expert_parameter`). FSDP2 then shards that slice again over the +``efsdp`` mesh through a per-parameter ``shard_placement_fn`` (see +``specforge/training/fsdp2.py``), so FSDP and EP compose and +``get_model_state_dict(full_state_dict=True)`` still yields the full ``[E, ...]`` +tensors the checkpoint converters expect. + +Both collectives have fixed sizes, so EP adds no device-to-host sync of its +own; the only host sync is reading this rank's slot bounds out of the routing +counts, which the ``sorted_loop`` dispatch already pays. + +Gradient semantics: a rank's expert gradient covers the tokens of its whole EP +group. FSDP2 averages it over the ``efsdp`` ranks only, while dense gradients +are averaged over the full data-parallel world, so the FSDP2 backend rescales +expert gradients by ``1 / ep`` at the optimizer boundary. +""" + +from __future__ import annotations + +from dataclasses import dataclass + +import torch +import torch.distributed as dist +from torch import nn +from torch.distributed.device_mesh import DeviceMesh +from torch.distributed.tensor import DTensor, Shard, distribute_tensor + + +@dataclass(frozen=True) +class ExpertParallelLayout: + """Which experts this rank owns, and the group it shares tokens with.""" + + mesh: DeviceMesh + rank: int + size: int + n_experts: int + + @property + def n_local_experts(self) -> int: + return self.n_experts // self.size + + @property + def expert_start(self) -> int: + return self.rank * self.n_local_experts + + @property + def expert_end(self) -> int: + return self.expert_start + self.n_local_experts + + @property + def group(self): + return self.mesh.get_group() + + +def expert_parallel_layout(ep_mesh: DeviceMesh, n_experts: int) -> ExpertParallelLayout: + """Resolve the 1-D ``ep`` mesh into this rank's expert slice.""" + if ep_mesh.ndim != 1: + raise ValueError(f"expert parallelism needs a 1-D mesh, got {ep_mesh.ndim}-D") + size = ep_mesh.size() + if size < 1 or n_experts % size: + raise ValueError( + f"n_routed_experts={n_experts} is not divisible by " + f"expert_parallel_size={size}" + ) + return ExpertParallelLayout( + mesh=ep_mesh, rank=ep_mesh.get_local_rank(), size=size, n_experts=n_experts + ) + + +def shard_expert_parameter(param: nn.Parameter, ep_mesh: DeviceMesh) -> nn.Parameter: + """Replace a full stacked ``[E, ...]`` parameter by this rank's EP shard. + + The result is a ``DTensor`` with ``Shard(0)`` over ``ep_mesh``; the slice is + scattered from rank 0 of the mesh so every EP group starts from the same + weights even if the ranks' initializations drifted. Frozen parameters stay + frozen. + """ + if isinstance(param, DTensor): + raise ValueError( + "parameter is already a DTensor; expert parallelism was applied twice" + ) + sharded = distribute_tensor(param.detach(), ep_mesh, [Shard(0)]) + return nn.Parameter(sharded, requires_grad=param.requires_grad) + + +def local_expert_weight(param: torch.Tensor) -> torch.Tensor: + """This rank's ``[E/ep, ...]`` slice of a (possibly FSDP-unsharded) expert weight.""" + return param.to_local() if isinstance(param, DTensor) else param + + +class _GatherTokens(torch.autograd.Function): + """All-gather ``[T, D]`` -> ``[ep * T, D]``; backward reduce-scatters (sum).""" + + @staticmethod + def forward(ctx, x: torch.Tensor, group): + ctx.group = group + world = dist.get_world_size(group) + x = x.contiguous() + out = x.new_empty((world * x.shape[0],) + tuple(x.shape[1:])) + dist.all_gather_into_tensor(out, x, group=group) + return out + + @staticmethod + def backward(ctx, grad_output: torch.Tensor): + world = dist.get_world_size(ctx.group) + grad_output = grad_output.contiguous() + grad = grad_output.new_empty( + (grad_output.shape[0] // world,) + tuple(grad_output.shape[1:]) + ) + dist.reduce_scatter_tensor( + grad, grad_output, op=dist.ReduceOp.SUM, group=ctx.group + ) + return grad, None + + +class _ScatterOutputs(torch.autograd.Function): + """Reduce-scatter (sum) ``[ep * T, D]`` -> ``[T, D]``; backward all-gathers.""" + + @staticmethod + def forward(ctx, partial: torch.Tensor, group): + ctx.group = group + world = dist.get_world_size(group) + partial = partial.contiguous() + out = partial.new_empty((partial.shape[0] // world,) + tuple(partial.shape[1:])) + dist.reduce_scatter_tensor(out, partial, op=dist.ReduceOp.SUM, group=group) + return out + + @staticmethod + def backward(ctx, grad_output: torch.Tensor): + world = dist.get_world_size(ctx.group) + grad_output = grad_output.contiguous() + grad = grad_output.new_empty( + (world * grad_output.shape[0],) + tuple(grad_output.shape[1:]) + ) + dist.all_gather_into_tensor(grad, grad_output, group=ctx.group) + return grad, None + + +def gather_tokens(x: torch.Tensor, layout: ExpertParallelLayout) -> torch.Tensor: + """Every rank's tokens, in EP-rank order; differentiable.""" + return _GatherTokens.apply(x, layout.group) + + +def scatter_outputs( + partial: torch.Tensor, layout: ExpertParallelLayout +) -> torch.Tensor: + """Sum the per-owner partial outputs and return this rank's tokens; differentiable.""" + return _ScatterOutputs.apply(partial, layout.group) + + +__all__ = [ + "ExpertParallelLayout", + "expert_parallel_layout", + "gather_tokens", + "local_expert_weight", + "scatter_outputs", + "shard_expert_parameter", +] diff --git a/specforge/modeling/draft/moe/grouped_experts.py b/specforge/modeling/draft/moe/grouped_experts.py index 7a511a60e..1c39c3d00 100644 --- a/specforge/modeling/draft/moe/grouped_experts.py +++ b/specforge/modeling/draft/moe/grouped_experts.py @@ -17,17 +17,34 @@ - ``"grouped_mm"``: the same segments through ``torch._grouped_mm`` with on-device offsets (no host sync). Used on CUDA when available; falls back to the loop elsewhere. Same math up to bf16 rounding. + +Expert parallelism (:mod:`.expert_parallel`): after +:meth:`GroupedExperts.apply_expert_parallel` the stacked parameters are +``DTensor`` ``Shard(0)`` slices over the ``ep`` mesh and the forward receives the +EP group's gathered tokens. Sorting by expert makes this rank's experts one +contiguous run of slots, so the local work is a slice of the sorted order; the +returned fp32 output is this rank's PARTIAL sum, which :class:`MoELayer` +reduce-scatters across the group. """ from __future__ import annotations import re +from typing import Optional import torch import torch.nn.functional as F from torch import nn +from torch.distributed.device_mesh import DeviceMesh +from torch.distributed.tensor import DTensor from .config import MoEConfig +from .expert_parallel import ( + ExpertParallelLayout, + expert_parallel_layout, + local_expert_weight, + shard_expert_parameter, +) from .experts import RoutedExperts, register_experts_backend from .router import RoutingResult from .state_dict import register_state_dict_converter @@ -61,50 +78,115 @@ def __init__(self, cfg: MoEConfig, hidden_size: int) -> None: self.w1 = nn.Parameter(torch.empty(e, i, d)) self.w2 = nn.Parameter(torch.empty(e, d, i)) self.w3 = nn.Parameter(torch.empty(e, i, d)) + #: Set by :meth:`apply_expert_parallel`; ``None`` keeps every expert local. + self.ep: Optional[ExpertParallelLayout] = None + + @property + def n_local_experts(self) -> int: + return self.ep.n_local_experts if self.ep is not None else self.n_experts + + def apply_expert_parallel(self, ep_mesh: DeviceMesh) -> ExpertParallelLayout: + """Keep only this rank's ``[E/ep]`` slice of the stacked experts. + + Call after initialization and warm start (the full tensors are sliced + in place) and before FSDP wrapping, which shards the slice further. + """ + if self.ep is not None: + raise RuntimeError( + "expert parallelism was already applied to these experts" + ) + layout = expert_parallel_layout(ep_mesh, self.n_experts) + for name in self._WEIGHT_NAMES: + setattr(self, name, shard_expert_parameter(getattr(self, name), ep_mesh)) + self.ep = layout + return layout def reset_parameters(self, std: float) -> None: if self.w1.device.type == "meta": return for name in self._WEIGHT_NAMES: - nn.init.normal_(getattr(self, name), mean=0.0, std=std) + nn.init.normal_(local_expert_weight(getattr(self, name)), mean=0.0, std=std) def forward(self, x: torch.Tensor, routing: RoutingResult) -> torch.Tensor: flat_expert = routing.indices.flatten() # [T*k] order = flat_expert.argsort(stable=True) - token_of = order // routing.topk # routed token index per sorted slot - x_sorted = x.index_select(0, token_of) - w_sorted = routing.weights.reshape(-1, 1).index_select(0, order).float() counts = routing.counts - - if self.grouped_mm and x.is_cuda: - offs = counts.cumsum(0).to(torch.int32) - gate = torch._grouped_mm(x_sorted, self.w1.transpose(-1, -2), offs=offs) - up = torch._grouped_mm(x_sorted, self.w3.transpose(-1, -2), offs=offs) + ep = self.ep + w1, w2, w3 = (local_expert_weight(getattr(self, n)) for n in self._WEIGHT_NAMES) + + counts_list = None + n_local_tokens = None + if ep is None: + order_local = order + local_counts = counts + else: + # Sorting by expert makes this rank's experts one contiguous run of + # slots; reading its bounds is the one host sync EP costs per MoE + # layer (the sorted_loop path pays it anyway). + counts_list = counts.tolist() + start_slot = sum(counts_list[: ep.expert_start]) + n_local_tokens = sum(counts_list[ep.expert_start : ep.expert_end]) + order_local = order[start_slot : start_slot + n_local_tokens] + local_counts = counts[ep.expert_start : ep.expert_end] + + token_of = order_local // routing.topk # routed token index per sorted slot + x_sorted = x.index_select(0, token_of) + w_sorted = routing.weights.reshape(-1, 1).index_select(0, order_local).float() + + y_routed = None + if n_local_tokens == 0: + # No token routed to this rank's experts this micro-batch: routine + # on an imbalanced route under EP. The zero terms below keep the + # graph (and so the collective order) identical on every rank. + pass + elif self.grouped_mm and x.is_cuda: + offs = local_counts.cumsum(0).to(torch.int32) + gate = torch._grouped_mm(x_sorted, w1.transpose(-1, -2), offs=offs) + up = torch._grouped_mm(x_sorted, w3.transpose(-1, -2), offs=offs) h = w_sorted * swiglu_clamped(gate, up, self.swiglu_limit) - y_routed = torch._grouped_mm( - h.to(x.dtype), self.w2.transpose(-1, -2), offs=offs - ) + y_routed = torch._grouped_mm(h.to(x.dtype), w2.transpose(-1, -2), offs=offs) else: - counts_list = counts.tolist() # one host sync per MoE forward + if counts_list is None: + counts_list = counts.tolist() # one host sync per MoE forward + local_list = ( + counts_list + if ep is None + else counts_list[ep.expert_start : ep.expert_end] + ) parts = [] offset = 0 - for i, n in enumerate(counts_list): + for i, n in enumerate(local_list): if n == 0: continue seg = x_sorted[offset : offset + n] h = w_sorted[offset : offset + n] * swiglu_clamped( - F.linear(seg, self.w1[i]), - F.linear(seg, self.w3[i]), + F.linear(seg, w1[i]), + F.linear(seg, w3[i]), self.swiglu_limit, ) - parts.append(F.linear(h.to(seg.dtype), self.w2[i])) + parts.append(F.linear(h.to(seg.dtype), w2[i])) offset += n - if not parts: - return torch.zeros_like(x) - y_routed = torch.cat(parts, dim=0) + if parts: + y_routed = torch.cat(parts, dim=0) y = torch.zeros(x.shape, dtype=torch.float32, device=x.device) - y = y.index_add(0, token_of, y_routed.float()) + if ep is not None: + # Keep the gathered input, the combine weights and every expert + # parameter on the graph even when this rank computed nothing: + # autograd must reach the EP collectives and FSDP2 must see the same + # set of gradients on every rank. Adds nothing to the value. + y = ( + y + + x.reshape(-1)[0].float() * 0.0 + + routing.weights.float().sum() * 0.0 + + sum(w.reshape(-1)[0].float() * 0.0 for w in (w1, w2, w3)) + ) + if y_routed is not None: + y = y.index_add(0, token_of, y_routed.float()) + if ep is not None: + # This rank's partial sum in fp32; MoELayer reduce-scatters it over + # the EP group and casts afterwards. + return y return y.to(x.dtype) @@ -122,6 +204,11 @@ def unstack_grouped_expert_state_dict(state: dict) -> dict: if m is None or not isinstance(value, torch.Tensor) or value.dim() != 3: out[key] = value continue + if isinstance(value, DTensor): + raise TypeError( + f"{key} is still an expert-parallel DTensor shard; gather the full " + "state (get_model_state_dict(full_state_dict=True)) before converting" + ) for i in range(value.shape[0]): out[f"{m['base']}.{i}.{m['w']}.weight"] = value[i] return out diff --git a/specforge/modeling/draft/moe/hooks.py b/specforge/modeling/draft/moe/hooks.py index 2867d370a..b22d03c31 100644 --- a/specforge/modeling/draft/moe/hooks.py +++ b/specforge/modeling/draft/moe/hooks.py @@ -10,6 +10,9 @@ - :func:`collect_moe_aux_loss` to add to the objective when a balance policy emits one; - :func:`collect_moe_metrics` for per-step diagnostics (``moe/...``). + +The training backend uses :func:`apply_expert_parallel` to slice every layer's +experts over the expert-parallel mesh before it wraps the model. """ from __future__ import annotations @@ -18,6 +21,7 @@ import torch from torch import nn +from torch.distributed.device_mesh import DeviceMesh from .balance import MetricValue from .layer import MoELayer @@ -34,6 +38,15 @@ def apply_pending_balance_updates(module: nn.Module) -> None: layer.apply_pending_balance_update() +def apply_expert_parallel(module: nn.Module, ep_mesh: DeviceMesh) -> int: + """Shard every MoE layer's experts over ``ep_mesh``; returns the layer count.""" + count = 0 + for layer in iter_moe_layers(module): + layer.apply_expert_parallel(ep_mesh) + count += 1 + return count + + def collect_moe_aux_loss(module: nn.Module) -> Optional[torch.Tensor]: """Sum of the layers' (already scaled) auxiliary losses, or ``None``.""" total: Optional[torch.Tensor] = None diff --git a/specforge/modeling/draft/moe/layer.py b/specforge/modeling/draft/moe/layer.py index 9857d98dc..4b622db31 100644 --- a/specforge/modeling/draft/moe/layer.py +++ b/specforge/modeling/draft/moe/layer.py @@ -4,6 +4,11 @@ Attribute names follow the official DeepSeek-style checkpoint layout (``gate``, ``experts``, ``shared_experts``) so that per-implementation converters only need to handle their own internals. + +Under expert parallelism (:meth:`MoELayer.apply_expert_parallel`) the layer +all-gathers its tokens over the EP group, routes and computes the locally owned +experts on all of them, and reduce-scatters the partial outputs back. The +shared expert and everything outside the layer stay data-parallel. """ from __future__ import annotations @@ -12,9 +17,11 @@ import torch from torch import nn +from torch.distributed.device_mesh import DeviceMesh from .balance import MetricValue, build_balance_controller from .config import MoEConfig, resolve_moe_config +from .expert_parallel import ExpertParallelLayout, gather_tokens, scatter_outputs from .experts import build_routed_experts from .router import RoutingResult, build_router from .shared import build_shared_expert @@ -35,6 +42,8 @@ def __init__(self, cfg: MoEConfig, hidden_size: int) -> None: self.shared_experts: Optional[nn.Module] = ( build_shared_expert(cfg, hidden_size) if cfg.n_shared_experts else None ) + #: Expert-parallel layout once :meth:`apply_expert_parallel` ran. + self.ep: Optional[ExpertParallelLayout] = None # Detached per-expert counts of the last training forward, for metrics. self.last_counts: Optional[torch.Tensor] = None @@ -42,14 +51,36 @@ def __init__(self, cfg: MoEConfig, hidden_size: int) -> None: def balance(self): return self.gate.balance + def apply_expert_parallel(self, ep_mesh: DeviceMesh) -> ExpertParallelLayout: + """Slice the routed experts over ``ep_mesh`` (a 1-D mesh of the EP group). + + The router, balance controller and shared expert stay replicated; the + experts backend decides how its weights are sliced. + """ + apply = getattr(self.experts, "apply_expert_parallel", None) + if apply is None: + raise NotImplementedError( + f"{type(self.experts).__name__} does not support expert parallelism" + ) + self.ep = apply(ep_mesh) + return self.ep + def forward(self, x: torch.Tensor) -> torch.Tensor: shape = x.shape x = x.reshape(-1, self.hidden_size) - routing: RoutingResult = self.gate(x) + ep = self.ep + # Under EP every rank routes and computes its experts for the whole + # group's tokens; the replicated router sees identical inputs on every + # rank of the group, so no routing has to be exchanged. + routed_in = gather_tokens(x, ep) if ep is not None else x + routing: RoutingResult = self.gate(routed_in) if self.training: self.last_counts = routing.counts.detach() self.balance.observe(routing) - y = self.experts(x, routing) + y = self.experts(routed_in, routing) + if ep is not None: + # fp32 partial sums from every expert owner -> this rank's tokens. + y = scatter_outputs(y, ep).to(x.dtype) if self.shared_experts is not None: y = y + self.shared_experts(x) return y.view(shape) diff --git a/specforge/training/backend.py b/specforge/training/backend.py index 651a4411f..ff2712bb4 100644 --- a/specforge/training/backend.py +++ b/specforge/training/backend.py @@ -62,6 +62,8 @@ class ParallelConfig: tp_size: int = 1 sp_ulysses_size: int = 1 sp_ring_size: int = 1 + #: MoE experts sliced across this many consecutive ranks (FSDP2 only). + expert_parallel_size: int = 1 sharding_strategy: str = "SHARD_GRAD_OP" param_dtype: torch.dtype = torch.bfloat16 fsdp_process_group: Any = None @@ -71,6 +73,8 @@ class ParallelConfig: sp_ulysses_group: Any = None sp_ring_group: Any = None draft_sp_group: Any = None + #: 2-D ``(efsdp, ep)`` mesh from init_distributed when expert_parallel_size > 1. + draft_ep_mesh: Any = None device_mesh: Any = None tp_device_mesh: Any = None extra: dict = field(default_factory=dict) @@ -114,6 +118,7 @@ def from_distributed( ("sp_ulysses_group", "get_sp_ulysses_group"), ("sp_ring_group", "get_sp_ring_group"), ("draft_sp_group", "get_draft_sp_group"), + ("draft_ep_mesh", "get_draft_ep_mesh"), ("device_mesh", "get_device_mesh"), ("tp_device_mesh", "get_tp_device_mesh"), ): @@ -133,11 +138,14 @@ def from_distributed( "ParallelConfig.from_distributed: distributed handles unavailable: %s", exc, ) + ep_mesh = handles.get("draft_ep_mesh") + expert_parallel_size = int(ep_mesh["ep"].size()) if ep_mesh is not None else 1 return cls( world_size=dist.get_world_size(), tp_size=tp_size, sp_ulysses_size=sp_ulysses_size, sp_ring_size=sp_ring_size, + expert_parallel_size=expert_parallel_size, sharding_strategy=sharding_strategy, param_dtype=param_dtype, fsdp_process_group=dist.group.WORLD, @@ -265,6 +273,11 @@ def prepare_model( for module in model.modules() if type(module).__name__ in block_names } + if pc.expert_parallel_size > 1 and pc.sharding_strategy == "NO_SHARD": + raise ValueError( + "expert parallelism requires a sharded fsdp_sharding " + "(SHARD_GRAD_OP or FULL_SHARD), not NO_SHARD" + ) if pc.sharding_strategy == "NO_SHARD": # PyTorch deprecated FSDP's NO_SHARD mode in favor of DDP. # DDP gives this small draft model replicated-param execution @@ -488,6 +501,10 @@ class FSDPTrainingBackend(DistributedTrainingBackend): name = "fsdp" def _shard_model(self, model, block_classes, ignored_frozen_modules): + if self.parallel_config.expert_parallel_size > 1: + raise ValueError( + "training.expert_parallel_size > 1 requires training.backend=fsdp2" + ) import functools from torch.distributed.fsdp import BackwardPrefetch diff --git a/specforge/training/fsdp2.py b/specforge/training/fsdp2.py index 8616134e6..420ad8401 100644 --- a/specforge/training/fsdp2.py +++ b/specforge/training/fsdp2.py @@ -1,15 +1,30 @@ -"""Composable FSDP2 backend with the same training contract as FSDP1.""" +"""Composable FSDP2 backend with the same training contract as FSDP1. + +Expert parallelism (``training.expert_parallel_size > 1``) lives here as well: +each MoE layer's experts are sliced over the ``ep`` mesh before wrapping, and a +per-parameter ``shard_placement_fn`` lets FSDP2 shard the slices over the +``efsdp`` ranks while every other parameter stays on the full data-parallel +mesh. Gradient accumulation, grad-norm reduction over local shards and the +full-state-dict checkpoint path are unchanged. +""" import torch from torch.distributed.device_mesh import DeviceMesh from torch.distributed.fsdp import MixedPrecisionPolicy, fully_shard +from torch.distributed.tensor import Shard from specforge.training.backend import DistributedTrainingBackend +from specforge.training.params import local_tensor class FSDP2TrainingBackend(DistributedTrainingBackend): name = "fsdp2" + def __init__(self, parallel_config, *, optimizer_factory=None) -> None: + super().__init__(parallel_config, optimizer_factory=optimizer_factory) + self._expert_params: list = [] + self._expert_grad_scale = 1.0 + def _shard_model(self, model, block_classes, ignored_frozen_modules): pc = self.parallel_config if pc.sharding_strategy not in ("FULL_SHARD", "SHARD_GRAD_OP"): @@ -20,8 +35,14 @@ def _shard_model(self, model, block_classes, ignored_frozen_modules): mesh = DeviceMesh.from_group( pc.fsdp_process_group or torch.distributed.group.WORLD, device_type=device.type, + mesh_dim_names=("fsdp",), ) - ignored_params = { + kwargs = dict(mesh=mesh) + if pc.expert_parallel_size > 1: + # Slice the experts first: the frozen/ignored parameter sets below + # must refer to the sliced parameters. + kwargs["shard_placement_fn"] = self._apply_expert_parallel(model) + kwargs["ignored_params"] = { p for module in ignored_frozen_modules for p in module.parameters() } # FSDP1 keeps floating buffers (e.g. RoPE) in FP32. FSDP2 does not @@ -32,10 +53,6 @@ def _shard_model(self, model, block_classes, ignored_frozen_modules): for name, buffer in module.named_buffers(recurse=False): if buffer.is_floating_point(): setattr(module, name, buffer.float()) - kwargs = dict( - mesh=mesh, - ignored_params=ignored_params, - ) # Bottom-up application preserves the FSDP1 block/root boundaries. # Leave non-block draft parameters on the composite root: some losses # call custom draft head methods after draft_model.forward has returned. @@ -57,8 +74,66 @@ def _shard_model(self, model, block_classes, ignored_frozen_modules): mp_policy=MixedPrecisionPolicy(param_dtype=pc.param_dtype), **kwargs, ) + if pc.expert_parallel_size > 1: + # fully_shard replaced every parameter with its FSDP-sharded DTensor; + # gradient rescaling must address those, not the pre-wrap slices. + from specforge.modeling.draft.moe import iter_moe_layers + + self._expert_params = [ + parameter + for layer in iter_moe_layers(model) + for parameter in layer.experts.parameters() + if parameter.requires_grad + ] return model + def _apply_expert_parallel(self, model): + """Slice MoE experts over the ``ep`` mesh; return FSDP2's placement fn. + + Routed experts end up ``Shard(0)`` over ``ep`` (their own slice) and + ``Shard(0)`` over ``efsdp`` (FSDP across the ranks that own the same + slice); dense parameters keep the default placement on the full mesh. + """ + from torch.distributed.fsdp._fully_shard._fsdp_common import ( + ShardPlacementResult, + ) + from torch.distributed.fsdp._fully_shard._fsdp_init import _get_mesh_info + + from specforge.modeling.draft.moe import apply_expert_parallel, iter_moe_layers + + pc = self.parallel_config + ep_mesh = pc.draft_ep_mesh + if ep_mesh is None: + raise RuntimeError( + "training.expert_parallel_size > 1 but init_distributed built no " + "expert-parallel mesh" + ) + if apply_expert_parallel(model, ep_mesh["ep"]) == 0: + raise ValueError( + "training.expert_parallel_size > 1 requires a MoE draft " + "(n_routed_experts > 0 in the draft JSON)" + ) + expert_params = { + parameter + for layer in iter_moe_layers(model) + for parameter in layer.experts.parameters() + } + expert_mesh_info = _get_mesh_info(ep_mesh["efsdp"]) + + def shard_placement_fn(param): + if param in expert_params: + return ShardPlacementResult( + placement=Shard(0), mesh_info=expert_mesh_info + ) + return None + + # FSDP2 averages every parameter's gradient over its own reduce-scatter + # group: ``efsdp`` for the experts, the full mesh for dense parameters. + # An expert gradient already sums its whole EP group's tokens, so bring + # it to the dense per-token scale before clipping and the optimizer. + self._expert_grad_scale = 1.0 / pc.expert_parallel_size + return shard_placement_fn + def backward(self, loss: torch.Tensor, *, is_boundary: bool = True) -> None: if self._wrapper_kind != "fsdp2": return super().backward(loss, is_boundary=is_boundary) @@ -74,6 +149,12 @@ def backward(self, loss: torch.Tensor, *, is_boundary: bool = True) -> None: finally: self.module.set_requires_gradient_sync(True) self.module.set_reshard_after_backward(True) + if is_boundary and self._expert_params: + # Reduced expert gradients exist only after the boundary backward. + with torch.no_grad(): + for parameter in self._expert_params: + if parameter.grad is not None: + local_tensor(parameter.grad).mul_(self._expert_grad_scale) def _sharded_model_state_dict(self) -> dict: from torch.distributed.checkpoint.state_dict import ( @@ -83,6 +164,7 @@ def _sharded_model_state_dict(self) -> dict: # All ranks participate; full_state_dict + cpu_offload returns the # gathered ordinary tensors on rank zero and an empty dict elsewhere. + # Expert-parallel DTensor slices gather to the full [E, ...] tensors. return get_model_state_dict( self.module, options=StateDictOptions(full_state_dict=True, cpu_offload=True), diff --git a/tests/test_config/test_schema.py b/tests/test_config/test_schema.py index 55b7ee280..c66e8c9f9 100644 --- a/tests/test_config/test_schema.py +++ b/tests/test_config/test_schema.py @@ -856,3 +856,45 @@ def test_load_config_applies_overrides(self): if __name__ == "__main__": unittest.main(verbosity=2) + + +class ExpertParallelSchemaTest(unittest.TestCase): + """``training.expert_parallel_size`` is an FSDP2-only, sharded-only knob.""" + + def _payload(self, **training): + payload = copy.deepcopy(MINIMAL) + payload["training"] = dict(training) + return payload + + def test_defaults_to_one(self): + config = Config.model_validate(self._payload()) + self.assertEqual(config.training.expert_parallel_size, 1) + config.validate_world_size(3) + + def test_requires_the_fsdp2_backend(self): + with self.assertRaisesRegex(ValidationError, "expert_parallel_size"): + Config.model_validate(self._payload(expert_parallel_size=2)) + config = Config.model_validate( + self._payload(backend="fsdp2", expert_parallel_size=2) + ) + self.assertEqual(config.training.expert_parallel_size, 2) + + def test_rejects_no_shard_and_model_parallel_combinations(self): + with self.assertRaisesRegex(ValidationError, "NO_SHARD"): + Config.model_validate( + self._payload( + backend="fsdp2", expert_parallel_size=2, fsdp_sharding="NO_SHARD" + ) + ) + with self.assertRaisesRegex(ValidationError, "tp_size=1"): + Config.model_validate( + self._payload(backend="fsdp2", expert_parallel_size=2, tp_size=2) + ) + + def test_world_size_must_be_divisible(self): + config = Config.model_validate( + self._payload(backend="fsdp2", expert_parallel_size=2) + ) + config.validate_world_size(4) + with self.assertRaisesRegex(ValueError, "expert_parallel_size"): + config.validate_world_size(3) diff --git a/tests/test_modeling/test_moe_expert_parallel.py b/tests/test_modeling/test_moe_expert_parallel.py new file mode 100644 index 000000000..7cd59c33e --- /dev/null +++ b/tests/test_modeling/test_moe_expert_parallel.py @@ -0,0 +1,336 @@ +# coding=utf-8 +"""Expert parallelism over the grouped routed experts, and on the FSDP2 backend. + +The layer contract: an EP-sharded layer is numerically the unsharded layer run +on the EP group's tokens. Each rank gets its own tokens' outputs and input +gradients; expert gradients equal the reference's slice (the owner sees the +whole group's tokens); the replicated router and shared expert hold partial +gradients that sum over the group to the reference's. A rank whose experts +receive no token must still complete every collective. + +The backend contract: under ``FSDP2TrainingBackend`` with +``expert_parallel_size > 1`` every gradient matches a data-parallel reference +(dense ones through FSDP averaging, expert ones after the ``1 / ep`` rescale), +the full checkpoint keeps the official per-expert naming, and it reloads. + +CPU gloo with 2 and 4 processes; no CUDA needed. +""" + +import os +import tempfile +import unittest + +import torch +import torch.distributed as dist +import torch.multiprocessing as mp +from torch import nn +from torch.distributed.device_mesh import init_device_mesh +from torch.distributed.tensor import DTensor +from torch.testing import assert_close + +from specforge.modeling.draft.moe import ( + MoELayer, + iter_moe_layers, + resolve_moe_config, + to_checkpoint_state_dict, +) + +N_EXPERTS = 8 +TOPK = 2 +HIDDEN = 32 +TOKENS = 6 + + +def _config(): + return resolve_moe_config( + dict( + moe_preset="deepseek_v4", + n_routed_experts=N_EXPERTS, + num_experts_per_tok=TOPK, + moe_intermediate_size=16, + dflash_config={"moe_bias_update_rate": 1e-3}, + ) + ) + + +def _init_layer(layer: MoELayer) -> MoELayer: + layer.reset_parameters(std=0.05) + for parameter in layer.shared_experts.parameters(): + nn.init.normal_(parameter, std=0.05) + return layer + + +def _layer(seed: int = 0) -> MoELayer: + torch.manual_seed(seed) + return _init_layer(MoELayer(_config(), HIDDEN)) + + +def _init_process_group(rank: int, world: int, init_file: str) -> None: + if os.uname().sysname == "Darwin": + os.environ.setdefault("GLOO_SOCKET_IFNAME", "lo0") + dist.init_process_group( + "gloo", init_method=f"file://{init_file}", rank=rank, world_size=world + ) + + +def _gather_all(tensor: torch.Tensor, world: int) -> torch.Tensor: + parts = [torch.empty_like(tensor) for _ in range(world)] + dist.all_gather(parts, tensor) + return torch.cat(parts) + + +def _full(tensor: torch.Tensor) -> torch.Tensor: + return tensor.full_tensor() if isinstance(tensor, DTensor) else tensor + + +# -------------------------------------------------------------------------- +# Layer level: a 1-D ``ep`` mesh over two ranks, no FSDP. +# -------------------------------------------------------------------------- +def _layer_worker(rank, init_file, starve_a_rank): + world = 2 + _init_process_group(rank, world, init_file) + try: + reference = _layer() + if starve_a_rank: + # Force every token onto experts 0..3 so the rank owning 4..7 sees a + # micro-batch with no selected expert at all. Through the + # selection-only balancing bias, so the combine scores are untouched. + with torch.no_grad(): + reference.gate.balance.bias[N_EXPERTS // 2 :].fill_(-1e9) + sharded = _layer() + sharded.load_state_dict(reference.state_dict()) + mesh = init_device_mesh("cpu", (world,), mesh_dim_names=("ep",)) + layout = sharded.apply_expert_parallel(mesh) + local = N_EXPERTS // world + assert layout.size == world and layout.expert_start == rank * local + assert isinstance(sharded.experts.w1, DTensor) + assert sharded.experts.w1.to_local().shape[0] == local + expert_slice = slice(rank * local, (rank + 1) * local) + for name in ("w1", "w2", "w3"): + assert_close( + getattr(sharded.experts, name).to_local(), + getattr(reference.experts, name)[expert_slice], + rtol=0, + atol=0, + ) + + torch.manual_seed(100 + rank) + x_local = torch.randn(TOKENS, HIDDEN) + cotangent_local = torch.randn(TOKENS, HIDDEN) + x_all = _gather_all(x_local, world).requires_grad_(True) + cotangent_all = _gather_all(cotangent_local, world) + token_slice = slice(rank * TOKENS, (rank + 1) * TOKENS) + + # Reference: the unsharded layer over the whole group's tokens, with the + # sum of every rank's loss. + reference_output = reference(x_all) + (reference_output * cotangent_all).sum().backward() + + sharded_input = x_local.clone().requires_grad_(True) + sharded_output = sharded(sharded_input) + (sharded_output * cotangent_local).sum().backward() + + assert_close( + sharded_output, reference_output[token_slice], rtol=1e-5, atol=1e-6 + ) + assert_close(sharded_input.grad, x_all.grad[token_slice], rtol=1e-5, atol=1e-6) + # The owner's expert gradient covers the whole group's tokens. + for name in ("w1", "w2", "w3"): + sharded_grad = getattr(sharded.experts, name).grad + assert sharded_grad is not None, f"{name} received no gradient" + assert_close( + sharded_grad.to_local(), + getattr(reference.experts, name).grad[expert_slice], + rtol=1e-5, + atol=1e-6, + ) + # Replicated parameters hold partial gradients that sum over the group. + replicated = [(sharded.gate.weight, reference.gate.weight)] + replicated += list( + zip( + sharded.shared_experts.parameters(), + reference.shared_experts.parameters(), + ) + ) + for sharded_param, reference_param in replicated: + total = sharded_param.grad.clone() + dist.all_reduce(total) + assert_close(total, reference_param.grad, rtol=1e-5, atol=1e-6) + # Routing statistics describe the gathered tokens. + assert int(sharded.last_counts.sum()) == world * TOKENS * TOPK + # The full state keeps the stacked layout and converts to official naming. + full = {key: _full(value) for key, value in sharded.state_dict().items()} + assert_close(full["experts.w1"], reference.experts.w1, rtol=0, atol=0) + official = to_checkpoint_state_dict(full) + assert "experts.0.w1.weight" in official + assert f"experts.{N_EXPERTS - 1}.w3.weight" in official + assert "gate.bias" in official + finally: + dist.destroy_process_group() + + +# -------------------------------------------------------------------------- +# Backend level: FSDP2 over four ranks with a 2-D ``(efsdp, ep)`` mesh. +# -------------------------------------------------------------------------- +class _Block(nn.Module): + def __init__(self, cfg): + super().__init__() + self.proj = nn.Linear(HIDDEN, HIDDEN, bias=False) + self.mlp = MoELayer(cfg, HIDDEN) + + def forward(self, x): + return x + self.mlp(self.proj(x)) + + +class _TinyMoEDraft(nn.Module): + _no_split_modules = ["_Block"] + + def __init__(self, cfg, n_layers: int = 2): + super().__init__() + self.layers = nn.ModuleList([_Block(cfg) for _ in range(n_layers)]) + self.norm = nn.LayerNorm(HIDDEN) + + def forward(self, x): + for layer in self.layers: + x = layer(x) + return self.norm(x) + + +def _tiny(seed: int = 0) -> _TinyMoEDraft: + torch.manual_seed(seed) + model = _TinyMoEDraft(_config()) + for layer in iter_moe_layers(model): + _init_layer(layer) + return model + + +def _backend_worker(rank, init_file, ep): + world = 4 + _init_process_group(rank, world, init_file) + try: + from specforge.training.backend import ParallelConfig + from specforge.training.fsdp2 import FSDP2TrainingBackend + + reference = _tiny() + model = _tiny() + model.load_state_dict(reference.state_dict()) + ep_mesh = init_device_mesh( + "cpu", (world // ep, ep), mesh_dim_names=("efsdp", "ep") + ) + parallel = ParallelConfig( + world_size=world, + expert_parallel_size=ep, + sharding_strategy="FULL_SHARD", + param_dtype=torch.float32, + fsdp_process_group=dist.group.WORLD, + draft_ep_mesh=ep_mesh, + ) + backend = FSDP2TrainingBackend(parallel) + wrapped = backend.prepare_model(model, optimizer_target=model) + assert backend.auto_wrap_block_classes == {_Block} + for layer in iter_moe_layers(wrapped): + assert layer.ep is not None and layer.ep.size == ep + assert isinstance(layer.experts.w1, DTensor) + assert "ep" in layer.experts.w1.device_mesh.mesh_dim_names + + torch.manual_seed(200 + rank) + x_local = torch.randn(TOKENS, HIDDEN) + cotangent_local = torch.randn(TOKENS, HIDDEN) + x_all = _gather_all(x_local, world) + cotangent_all = _gather_all(cotangent_local, world) + token_slice = slice(rank * TOKENS, (rank + 1) * TOKENS) + + # Data-parallel reference: FSDP averages the per-rank losses. + reference_output = reference(x_all) + ((reference_output * cotangent_all).sum() / world).backward() + + output = wrapped(x_local) + backend.backward((output * cotangent_local).sum(), is_boundary=True) + assert_close(output, reference_output[token_slice], rtol=1e-5, atol=1e-6) + + reference_params = dict(reference.named_parameters()) + for name, parameter in wrapped.named_parameters(): + assert parameter.grad is not None, name + assert_close( + _full(parameter.grad), + reference_params[name].grad, + rtol=1e-4, + atol=1e-6, + msg=lambda m, n=name: f"{n}: {m}", + ) + + # Full checkpoint in the official naming, gathered on rank zero. + state = backend.state_dict() + model_state = state["model"] + if rank == 0: + assert "layers.0.mlp.experts.0.w1.weight" in model_state + assert f"layers.1.mlp.experts.{N_EXPERTS - 1}.w2.weight" in model_state + assert "layers.0.mlp.experts.w1" not in model_state + assert "layers.0.mlp.gate.bias" in model_state + for name, value in model_state.items(): + assert not isinstance(value, DTensor), name + assert_close( + model_state["layers.1.mlp.experts.5.w2.weight"], + reference.layers[1].mlp.experts.w2[5], + rtol=0, + atol=0, + ) + else: + assert model_state == {} + # Every rank reads the same file in production; broadcast the payload + # here, perturb the live weights and reload. + payload = [state if rank == 0 else None] + dist.broadcast_object_list(payload, src=0) + with torch.no_grad(): + for parameter in wrapped.parameters(): + local = ( + parameter.to_local() + if isinstance(parameter, DTensor) + else parameter + ) + local.zero_() + backend.load_state_dict(payload[0]) + for name, parameter in wrapped.named_parameters(): + assert_close(_full(parameter), reference_params[name], rtol=0, atol=0) + finally: + dist.destroy_process_group() + + +@unittest.skipUnless(dist.is_available(), "torch.distributed is unavailable") +class ExpertParallelMoETest(unittest.TestCase): + def _spawn(self, worker, nprocs, *args): + if dist.is_initialized(): + self.skipTest("requires ownership of the singleton process group") + with tempfile.TemporaryDirectory() as directory: + mp.spawn( + worker, + args=(os.path.join(directory, "gloo-init"), *args), + nprocs=nprocs, + join=True, + ) + + def test_two_rank_sharding_matches_the_unsharded_layer(self): + self._spawn(_layer_worker, 2, False) + + def test_a_rank_with_no_selected_expert_still_reaches_the_collectives(self): + self._spawn(_layer_worker, 2, True) + + def test_fsdp2_backend_with_ep2_matches_data_parallel_reference(self): + self._spawn(_backend_worker, 4, 2) + + def test_fsdp2_backend_with_ep4_matches_data_parallel_reference(self): + self._spawn(_backend_worker, 4, 4) + + def test_without_expert_parallelism_the_layer_is_unchanged(self): + layer = _layer() + self.assertIsNone(layer.ep) + self.assertEqual(layer.experts.n_local_experts, N_EXPERTS) + self.assertFalse(isinstance(layer.experts.w1, DTensor)) + output = layer(torch.randn(TOKENS, HIDDEN)) + self.assertEqual(tuple(output.shape), (TOKENS, HIDDEN)) + official = to_checkpoint_state_dict(layer.state_dict()) + self.assertIn("experts.0.w1.weight", official) + + +if __name__ == "__main__": + unittest.main(verbosity=2)