From 3acf1aebdbfebbc211c6ce3b710490006bf47973 Mon Sep 17 00:00:00 2001 From: "Ethan (Yusheng) Su" Date: Fri, 2 Oct 2026 09:11:01 +0900 Subject: [PATCH 1/9] 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 2/9] 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 177dc643351e13362c7ba08f0a6ef9db67617aee Mon Sep 17 00:00:00 2001 From: "Ethan (Yusheng) Su" Date: Sun, 4 Oct 2026 14:02:25 +0900 Subject: [PATCH 3/9] Add training.compile_blocks for the FSDP2 backend BackendOptions carries opt-in backend behaviors from the typed config to the backend. With compile_blocks the FSDP2 backend compiles every draft block (the _no_split_modules classes, or the EAGLE midlayer) in place with nn.Module.compile before fully_shard, so the block keeps its class and its parameter FQNs; Dynamo skips the FSDP2 hooks (skip_fsdp_hooks), which run eagerly around the compiled block body. FSDP1 rejects the option because its blocks become FullyShardedDataParallel wrappers. Stacks directly on #915; the BackendOptions scaffolding is shared verbatim with the other FSDP2 option PRs so they can land in any order. --- specforge/config/schema.py | 4 + specforge/launch.py | 8 ++ specforge/training/DESIGN.md | 3 +- specforge/training/assembly.py | 2 + specforge/training/backend.py | 56 ++++++++- specforge/training/fsdp2.py | 21 ++++ specforge/training/trainer.py | 6 +- .../test_runtime/test_fsdp2_compile_blocks.py | 109 ++++++++++++++++++ 8 files changed, 205 insertions(+), 4 deletions(-) create mode 100644 tests/test_runtime/test_fsdp2_compile_blocks.py diff --git a/specforge/config/schema.py b/specforge/config/schema.py index f8728bf32..6ec592a7e 100644 --- a/specforge/config/schema.py +++ b/specforge/config/schema.py @@ -895,6 +895,8 @@ def _validate_mode(self): class TrainingConfig(StrictConfigModel): strategy: str = "eagle3" backend: Literal["fsdp", "fsdp2"] = "fsdp" + #: ``torch.compile`` every draft block (or the EAGLE midlayer) in place before FSDP2 sharding. Requires ``backend: fsdp2``. + compile_blocks: bool = False 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) @@ -1001,6 +1003,8 @@ def _validate_training_shape(self): "training.down_sample_ratio_min must be in " "(0, training.down_sample_ratio]" ) + if self.compile_blocks and self.backend != "fsdp2": + raise ValueError("training.compile_blocks requires training.backend=fsdp2") sp_size = self.sp_ulysses_size * self.sp_ring_size if self.attention_backend == "usp": if self.batch_size != 1: diff --git a/specforge/launch.py b/specforge/launch.py index 77f3aca3b..ac2b6eff4 100644 --- a/specforge/launch.py +++ b/specforge/launch.py @@ -48,6 +48,7 @@ def _assemble_trainer( optimizer_factory, training_backend: str = "fsdp", fsdp_sharding: Optional[str] = None, + backend_options=None, run_id: str, output_dir: str, batch_size: int, @@ -125,6 +126,7 @@ def _assemble_trainer( optimizer_factory=optimizer_factory, training_backend=training_backend, fsdp_sharding=fsdp_sharding, + backend_options=backend_options, run_id=run_id, output_dir=output_dir, batch_size=batch_size, @@ -552,6 +554,7 @@ def build_offline_runtime( optimizer_factory, training_backend: str = "fsdp", fsdp_sharding: Optional[str] = None, + backend_options=None, run_id: str, output_dir: str, ttt_length: int = 7, @@ -640,6 +643,7 @@ def refs_for_epoch(epoch): optimizer_factory=optimizer_factory, training_backend=training_backend, fsdp_sharding=fsdp_sharding, + backend_options=backend_options, run_id=run_id, output_dir=output_dir, batch_size=batch_size, @@ -682,6 +686,7 @@ def build_disagg_offline_runtime( optimizer_factory, training_backend: str = "fsdp", fsdp_sharding: Optional[str] = None, + backend_options=None, run_id: str, output_dir: str, ttt_length: int = 7, @@ -766,6 +771,7 @@ def refs_for_epoch(epoch): optimizer_factory=optimizer_factory, training_backend=training_backend, fsdp_sharding=fsdp_sharding, + backend_options=backend_options, run_id=run_id, output_dir=output_dir, batch_size=batch_size, @@ -1539,6 +1545,7 @@ def build_disagg_online_consumer( optimizer_factory, training_backend: str = "fsdp", fsdp_sharding: Optional[str] = None, + backend_options=None, run_id: str, output_dir: str, target_head=None, @@ -1892,6 +1899,7 @@ def stop_distributor_and_drain() -> None: optimizer_factory=optimizer_factory, training_backend=training_backend, fsdp_sharding=fsdp_sharding, + backend_options=backend_options, run_id=run_id, output_dir=output_dir, batch_size=batch_size, diff --git a/specforge/training/DESIGN.md b/specforge/training/DESIGN.md index d47a0ce05..74e3cf25c 100644 --- a/specforge/training/DESIGN.md +++ b/specforge/training/DESIGN.md @@ -28,7 +28,8 @@ 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. +defaulting to the original `fsdp` backend. `BackendOptions` carries opt-in +behaviors; `training.compile_blocks` compiles each draft block in place before FSDP2 sharding so the composable hooks stay outside Dynamo. 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()`. diff --git a/specforge/training/assembly.py b/specforge/training/assembly.py index 4a17fef87..0489d3b1e 100644 --- a/specforge/training/assembly.py +++ b/specforge/training/assembly.py @@ -35,6 +35,7 @@ from specforge.algorithms.contracts import FeatureMode from specforge.algorithms.registry import AlgorithmRegistration from specforge.config import Config +from specforge.training.backend import BackendOptions from specforge.training.provenance import ( model_resume_provenance as _model_resume_provenance, ) @@ -555,6 +556,7 @@ def _common_launch_kwargs( optimizer_factory=_optimizer_factory(cfg), training_backend=t.backend, fsdp_sharding=t.fsdp_sharding, + backend_options=BackendOptions(compile_blocks=t.compile_blocks), 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 9279f664a..b31bab9f0 100644 --- a/specforge/training/backend.py +++ b/specforge/training/backend.py @@ -145,6 +145,20 @@ def from_distributed( ) +@dataclass(frozen=True) +class BackendOptions: + """Opt-in backend behaviors selected by ``training.*`` config fields. + + Options act on the draft blocks (the ``_no_split_modules`` classes, or the + EAGLE ``midlayer`` when a draft advertises none) before sharding. FSDP1 + wraps blocks in ``FullyShardedDataParallel`` modules and supports none of + them; the FSDP2 backend applies them in ``_prepare_blocks``. + """ + + #: ``torch.compile`` every draft block (or the EAGLE midlayer) in place before FSDP2 sharding. + compile_blocks: bool = False + + class TrainingBackend(abc.ABC): name: str #: Whether ``step(loss_denominator=...)`` validates the global loss @@ -187,9 +201,11 @@ def __init__( parallel_config: ParallelConfig, *, optimizer_factory=None, + options: Optional[BackendOptions] = None, ) -> None: self.parallel_config = parallel_config self._optimizer_factory = optimizer_factory + self.options = options or BackendOptions() self.module: Optional[nn.Module] = None self.optimizer = None self._wrapped = False @@ -255,6 +271,7 @@ def prepare_model( for module in model.modules() if type(module).__name__ in block_names } + self._prepare_blocks(model, block_classes, optimizer_target) 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 @@ -294,6 +311,37 @@ def prepare_model( @abc.abstractmethod def _shard_model(self, model, block_classes, ignored_frozen_modules): ... + @staticmethod + def _block_targets( + model: nn.Module, block_classes, optimizer_target: Optional[nn.Module] + ) -> List[nn.Module]: + """Draft blocks that ``BackendOptions`` transforms act on. + + DFlash-family drafts advertise their decoder block class; EAGLE drafts + expose a single ``midlayer``. Working at this granularity keeps the + per-block FSDP boundaries and the custom head kernels untouched. + """ + targets = [m for m in model.modules() if type(m) in block_classes] + if not targets: + midlayer = getattr(optimizer_target, "midlayer", None) + if isinstance(midlayer, nn.Module): + targets = [midlayer] + return targets + + def _prepare_blocks( + self, model: nn.Module, block_classes, optimizer_target: Optional[nn.Module] + ) -> None: + """Hook for pre-sharding block transforms; FSDP1 supports none.""" + for option in BackendOptions.__dataclass_fields__: + if getattr(self.options, option): + raise ValueError( + f"BackendOptions.{option} is not supported by the " + f"{self.name!r} backend; use training.backend=fsdp2" + ) + + def _after_optimizer_step(self) -> None: + """Hook run after every optimizer step.""" + def set_optimizer(self, optimizer) -> None: self.optimizer = optimizer self._configure_optimizer_grad_norm() @@ -351,8 +399,11 @@ def step( 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) + grad_norm = self.optimizer.step() + else: + grad_norm = self.optimizer.step(loss_denominator=loss_denominator) + self._after_optimizer_step() + return grad_norm def state_dict(self) -> dict: """Model, optimizer, RNG and sharding metadata for resume. @@ -531,6 +582,7 @@ def create_training_backend(name: str, parallel_config: ParallelConfig, **kwargs __all__ = [ + "BackendOptions", "ParallelConfig", "TrainingBackend", "DistributedTrainingBackend", diff --git a/specforge/training/fsdp2.py b/specforge/training/fsdp2.py index 8616134e6..b5a67a301 100644 --- a/specforge/training/fsdp2.py +++ b/specforge/training/fsdp2.py @@ -9,6 +9,27 @@ class FSDP2TrainingBackend(DistributedTrainingBackend): name = "fsdp2" + compiled_blocks: int = 0 + + def _prepare_blocks(self, model, block_classes, optimizer_target) -> None: + if not self.options.compile_blocks: + return + targets = self._block_targets(model, block_classes, optimizer_target) + if not targets: + raise ValueError( + "BackendOptions.compile_blocks found no draft blocks to compile: " + "the draft advertises no _no_split_modules and has no midlayer" + ) + # ``nn.Module.compile`` compiles in place, so the block keeps its class + # (the ``fully_shard`` boundary below still matches ``block_classes``) + # and its parameter names (no ``_orig_mod.`` checkpoint prefix). The + # FSDP2 hooks registered afterwards run inside the compiled call, but + # Dynamo skips them (``torch._dynamo.config.skip_fsdp_hooks``), so they + # execute eagerly around the compiled block body. Dynamo starts static + # and marks shapes dynamic only after a recompilation. + for module in targets: + module.compile() + self.compiled_blocks = len(targets) def _shard_model(self, model, block_classes, ignored_frozen_modules): pc = self.parallel_config diff --git a/specforge/training/trainer.py b/specforge/training/trainer.py index 68ed08eaa..0dcfb7902 100644 --- a/specforge/training/trainer.py +++ b/specforge/training/trainer.py @@ -89,6 +89,7 @@ def __init__( optimizer_factory, training_backend: str = "fsdp", fsdp_sharding: Optional[str] = None, + backend_options=None, run_id: str, output_dir: str, batch_size: int, @@ -434,7 +435,10 @@ def set_epoch(epoch): sp_ring_size=sp_ring_size, ) backend = create_training_backend( - training_backend, parallel, optimizer_factory=optimizer_factory + training_backend, + parallel, + optimizer_factory=optimizer_factory, + options=backend_options, ) # 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 diff --git a/tests/test_runtime/test_fsdp2_compile_blocks.py b/tests/test_runtime/test_fsdp2_compile_blocks.py new file mode 100644 index 000000000..6b13da14b --- /dev/null +++ b/tests/test_runtime/test_fsdp2_compile_blocks.py @@ -0,0 +1,109 @@ +"""``training.compile_blocks``: in-place block compilation ahead of FSDP2 sharding.""" + +import json +import os +import tempfile +import unittest + +import torch +import torch.nn as nn + +from specforge.training.backend import ( + BackendOptions, + ParallelConfig, + create_training_backend, +) + + +class TestCompileBlocksSelection(unittest.TestCase): + def test_fsdp1_rejects_compile_blocks(self): + from tests.test_runtime.test_fsdp2_backend import TinyComposite + + pc = ParallelConfig(world_size=1) + backend = create_training_backend( + "fsdp", pc, options=BackendOptions(compile_blocks=True) + ) + with self.assertRaisesRegex(ValueError, "compile_blocks"): + backend.prepare_model(TinyComposite(), optimizer_target=None) + + def test_config_requires_fsdp2(self): + from specforge.config.schema import TrainingConfig + + with self.assertRaisesRegex(ValueError, "compile_blocks"): + TrainingConfig(compile_blocks=True) + cfg = TrainingConfig(backend="fsdp2", compile_blocks=True) + self.assertTrue(cfg.compile_blocks) + + def test_block_targets_fall_back_to_midlayer(self): + from specforge.training.backend import DistributedTrainingBackend + + class Draft(nn.Module): + def __init__(self): + super().__init__() + self.midlayer = nn.Linear(3, 3) + + class Composite(nn.Module): + def __init__(self): + super().__init__() + self.draft_model = Draft() + + model = Composite() + targets = DistributedTrainingBackend._block_targets( + model, set(), model.draft_model + ) + self.assertEqual(targets, [model.draft_model.midlayer]) + + +def _worker(rank, world_size, port, results_dir): + from tests.test_runtime import _fixtures as fx + from tests.test_runtime.test_fsdp2_backend import TinyComposite, _optimizer + + fx.init_rank_distributed(rank, world_size, port=str(port)) + torch.manual_seed(0) + reference = TinyComposite().cuda() + compiled = TinyComposite().cuda() + compiled.load_state_dict(reference.state_dict()) + x = torch.randn(4, 7, device="cuda") + + losses = {} + for name, model, opts in ( + ("eager", reference, BackendOptions()), + ("compiled", compiled, BackendOptions(compile_blocks=True)), + ): + # fp32 test models: match the backend's compute dtype like PR 915's gates. + pc = ParallelConfig.from_distributed(param_dtype=torch.float32) + backend = create_training_backend( + "fsdp2", pc, optimizer_factory=_optimizer, options=opts + ) + wrapped = backend.prepare_model(model, optimizer_target=model.draft_model) + blocks = [m for m in wrapped.modules() if type(m).__name__.endswith("TinyBlock")] + compiled_flags = [m._compiled_call_impl is not None for m in blocks] + step_losses = [] + for _ in range(3): + loss = wrapped(x).float().pow(2).mean() + backend.backward(loss, is_boundary=True) + backend.step() + step_losses.append(float(loss.detach())) + losses[name] = {"losses": step_losses, "compiled": compiled_flags} + with open(os.path.join(results_dir, f"rank{rank}.json"), "w") as f: + json.dump(losses, f) + + +@unittest.skipUnless(torch.cuda.device_count() >= 2, "requires two CUDA devices") +class TestCompileBlocksDistributed(unittest.TestCase): + def test_compiled_blocks_match_eager(self): + import torch.multiprocessing as mp + + results_dir = tempfile.mkdtemp(prefix="compile_blocks_") + mp.spawn(_worker, args=(2, 29613, results_dir), nprocs=2, join=True) + for rank in range(2): + with open(os.path.join(results_dir, f"rank{rank}.json")) as f: + out = json.load(f) + self.assertTrue(all(out["compiled"]["compiled"])) + self.assertFalse(any(out["eager"]["compiled"])) + for a, b in zip(out["eager"]["losses"], out["compiled"]["losses"]): + self.assertAlmostEqual(a, b, places=4) + + +if __name__ == "__main__": + unittest.main() From b60fe73aa6aa2a32c1a84c49c81ee93aa59ff85c Mon Sep 17 00:00:00 2001 From: "Ethan (Yusheng) Su" Date: Sun, 4 Oct 2026 16:56:13 +0900 Subject: [PATCH 4/9] Forward training.compile_blocks through the disaggregated launch path The disaggregated runtime (specforge/training/disaggregated.py) builds the offline and online trainers through its own call sites, which did not carry backend_options, so training.compile_blocks was silently ignored whenever training ran disaggregated, i.e. in every online run. The option is now resolved once, in assembly._backend_options, and used by the single-process path and both disaggregated paths. Both call sites are now covered by tests that drive _build_offline and _build_online with the launch entry points stubbed. --- specforge/training/assembly.py | 7 +- specforge/training/disaggregated.py | 4 + .../test_runtime/test_fsdp2_compile_blocks.py | 97 +++++++++++++++++++ 3 files changed, 107 insertions(+), 1 deletion(-) diff --git a/specforge/training/assembly.py b/specforge/training/assembly.py index 0489d3b1e..2d772871d 100644 --- a/specforge/training/assembly.py +++ b/specforge/training/assembly.py @@ -536,6 +536,11 @@ def _profiling_options(cfg: Config): ) +def _backend_options(cfg: Config) -> BackendOptions: + """``training.*`` -> typed backend options, shared by every launch path.""" + return BackendOptions(compile_blocks=cfg.training.compile_blocks) + + def _common_launch_kwargs( cfg: Config, bundle: ModelBundle, @@ -556,7 +561,7 @@ def _common_launch_kwargs( optimizer_factory=_optimizer_factory(cfg), training_backend=t.backend, fsdp_sharding=t.fsdp_sharding, - backend_options=BackendOptions(compile_blocks=t.compile_blocks), + backend_options=_backend_options(cfg), run_id=cfg.run_id, output_dir=cfg.output_dir, batch_size=t.batch_size, diff --git a/specforge/training/disaggregated.py b/specforge/training/disaggregated.py index 7752fc268..ecfa00c73 100644 --- a/specforge/training/disaggregated.py +++ b/specforge/training/disaggregated.py @@ -438,6 +438,7 @@ def _build_offline( ): from specforge.training.assembly import ( TrainingRun, + _backend_options, _dataloader_num_workers, _profiling_options, ) @@ -540,6 +541,7 @@ def produce() -> int: optimizer_factory=optimizer_factory(cfg), training_backend=cfg.training.backend, fsdp_sharding=cfg.training.fsdp_sharding, + backend_options=_backend_options(cfg), run_id=cfg.run_id, output_dir=cfg.output_dir, ttt_length=cfg.training.ttt_length, @@ -604,6 +606,7 @@ def _build_online( ): from specforge.training.assembly import ( TrainingRun, + _backend_options, _dataloader_num_workers, _load_input_tools, _profiling_options, @@ -821,6 +824,7 @@ def produce() -> int: optimizer_factory=optimizer_factory(cfg), training_backend=cfg.training.backend, fsdp_sharding=cfg.training.fsdp_sharding, + backend_options=_backend_options(cfg), run_id=cfg.run_id, output_dir=cfg.output_dir, batch_size=cfg.training.batch_size, diff --git a/tests/test_runtime/test_fsdp2_compile_blocks.py b/tests/test_runtime/test_fsdp2_compile_blocks.py index 6b13da14b..1a280197a 100644 --- a/tests/test_runtime/test_fsdp2_compile_blocks.py +++ b/tests/test_runtime/test_fsdp2_compile_blocks.py @@ -3,7 +3,9 @@ import json import os import tempfile +import types import unittest +from unittest import mock import torch import torch.nn as nn @@ -105,5 +107,100 @@ def test_compiled_blocks_match_eager(self): self.assertAlmostEqual(a, b, places=4) +class TestDisaggregatedLaunchForwardsCompileBlocks(unittest.TestCase): + """The disaggregated runtime builds its trainers through its own call sites + (offline consumer and online consumer); ``training.compile_blocks`` must + reach them exactly like the single-process path.""" + + @staticmethod + def _training(): + return { + "strategy": "dflash", + "role": "consumer", + "max_steps": 1, + "backend": "fsdp2", + "compile_blocks": True, + } + + def test_online_consumer_receives_the_option(self): + from specforge.algorithms.builtin import builtin_algorithm_registry + from specforge.config import Config + from specforge.training.disaggregated import _build_online + + cfg = Config.model_validate( + { + "model": {"target_model_path": "t", "draft_model_config": "d"}, + "data": {"prompts_path": "prompts.jsonl"}, + "training": self._training(), + "deployment": { + "mode": "disaggregated", + "disaggregated": { + "control_dir": "/shared/compile_blocks", + "backend": "mooncake", + "server_urls": ["http://capture:30000"], + }, + }, + } + ) + bundle = types.SimpleNamespace(model=object(), target_head=None, strategy_kwargs={}) + with ( + mock.patch.dict(os.environ, {"DISAGG_REF_CHANNEL": "/shared/refs"}), + mock.patch("specforge.runtime.data_plane.streaming_ref_channel.StreamingRefChannel"), + mock.patch("specforge.training.disaggregated._mooncake_store", return_value=mock.Mock()), + mock.patch("specforge.launch.build_disagg_online_consumer", return_value=_FakeFitTrainer()) as build, + ): + _build_online( + cfg, + algorithm=builtin_algorithm_registry().resolve("dflash"), + build_model_bundle=lambda _cfg: bundle, + prepare_prompts=mock.Mock(), + optimizer_factory=mock.Mock(), + logger=None, + ) + options = build.call_args.kwargs["backend_options"] + self.assertIsInstance(options, BackendOptions) + self.assertTrue(options.compile_blocks) + + def test_offline_consumer_receives_the_option(self): + from specforge.algorithms.builtin import builtin_algorithm_registry + from specforge.config import Config + from specforge.training.disaggregated import _build_offline + + cfg = Config.model_validate( + { + "model": {"target_model_path": "t", "draft_model_config": "d"}, + "data": {"hidden_states_path": "features"}, + "training": self._training(), + "deployment": { + "mode": "disaggregated", + "disaggregated": {"control_dir": "/shared/compile_blocks", "backend": "mooncake"}, + }, + } + ) + bundle = types.SimpleNamespace(model=object(), target_head=None, strategy_kwargs={}) + with ( + mock.patch.dict(os.environ, {"DISAGG_MANIFEST": "/shared/manifest.json"}), + mock.patch("specforge.training.disaggregated._wait_for"), + mock.patch("specforge.training.disaggregated._offline_store", return_value=mock.Mock()), + mock.patch("specforge.runtime.data_plane.disagg_ingest.read_ref_manifest", return_value=[]), + mock.patch("specforge.launch.build_disagg_offline_runtime", return_value=_FakeFitTrainer()) as build, + ): + _build_offline( + cfg, + algorithm=builtin_algorithm_registry().resolve("dflash"), + build_model_bundle=lambda _cfg: bundle, + optimizer_factory=mock.Mock(), + logger=None, + ) + options = build.call_args.kwargs["backend_options"] + self.assertIsInstance(options, BackendOptions) + self.assertTrue(options.compile_blocks) + + +class _FakeFitTrainer: + def fit(self): + return 1 + + if __name__ == "__main__": unittest.main() From d8e3cf32ca55c6d5cc4c1a3325637f18cd1adb0d Mon Sep 17 00:00:00 2001 From: "Ethan (Yusheng) Su" Date: Sun, 4 Oct 2026 17:02:58 +0900 Subject: [PATCH 5/9] Pass backend options to create_training_backend only when set Trainer always forwarded options=backend_options, which broke backend factories injected with the pre-options signature (test_domain_trainer builds one). Without configured options the call is now unchanged. --- specforge/training/trainer.py | 7 ++++++- 1 file changed, 6 insertions(+), 1 deletion(-) diff --git a/specforge/training/trainer.py b/specforge/training/trainer.py index 0dcfb7902..7bd9d9952 100644 --- a/specforge/training/trainer.py +++ b/specforge/training/trainer.py @@ -434,11 +434,16 @@ def set_epoch(epoch): sp_ulysses_size=sp_ulysses_size, sp_ring_size=sp_ring_size, ) + backend_kwargs = {} + if backend_options is not None: + # Injected backend factories (tests, extensions) keep seeing the + # pre-options call signature until options are configured. + backend_kwargs["options"] = backend_options backend = create_training_backend( training_backend, parallel, optimizer_factory=optimizer_factory, - options=backend_options, + **backend_kwargs, ) # 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 From 7815e576568f7007f1df99b34a46d4c6bd73e306 Mon Sep 17 00:00:00 2001 From: "Ethan (Yusheng) Su" Date: Sun, 4 Oct 2026 17:58:45 +0900 Subject: [PATCH 6/9] Add training.static_shapes: one batch length and anchor count per run compile_blocks needs every micro-batch to look the same. With real data the padded context length changes from one micro-batch to the next, Dynamo recompiles with that dimension symbolic, and torch 2.13's Inductor fails to lower flex_attention for the DFlash blocks (CantSplit: 8*s50 + 65536 not divisible by s50 + 8192); where it does compile, dynamic-shape kernels give back none of the fixed-shape speed-up. training.static_shapes pads every micro-batch to data.max_length (the algorithm collators take pad_to; the EAGLE3 offline collator too) and keeps num_anchors anchor slots per sample in the DFlash family, masking the slots past a row's valid anchors exactly like today's short rows. compile_blocks now requires it. The single-process path and both disaggregated call sites forward it; the online consumer also receives data.max_length to pad to. --- specforge/algorithms/common/collation.py | 13 +- .../algorithms/common/dflash_family_model.py | 18 +- .../algorithms/common/hidden_states_data.py | 17 +- specforge/algorithms/eagle3/data.py | 8 +- specforge/algorithms/model_providers.py | 1 + specforge/config/schema.py | 7 + specforge/data/utils.py | 12 +- specforge/launch.py | 48 ++++- specforge/training/DESIGN.md | 2 +- specforge/training/assembly.py | 1 + specforge/training/disaggregated.py | 3 + .../test_runtime/test_fsdp2_compile_blocks.py | 13 +- tests/test_runtime/test_static_shapes.py | 165 ++++++++++++++++++ 13 files changed, 287 insertions(+), 21 deletions(-) create mode 100644 tests/test_runtime/test_static_shapes.py diff --git a/specforge/algorithms/common/collation.py b/specforge/algorithms/common/collation.py index 44f506225..e8f8c120f 100644 --- a/specforge/algorithms/common/collation.py +++ b/specforge/algorithms/common/collation.py @@ -2,7 +2,7 @@ from __future__ import annotations -from typing import Mapping, Sequence +from typing import Mapping, Optional, Sequence def concatenate_features(features): @@ -27,11 +27,14 @@ def pad_and_concatenate_features( sequence_axes: Mapping[str, int], required_keys: Sequence[str], optional_keys: Sequence[str] = (), + pad_to: Optional[int] = None, ): """Zero-pad configured tensor axes to the longest input sequence. ``optional_keys`` are collated when every sample carries them and omitted - when none does; a batch that mixes both raises. + when none does; a batch that mixes both raises. With ``pad_to`` every + batch is padded to that fixed length instead (``training.static_shapes``), + and a longer sample raises. """ if not features: @@ -56,6 +59,12 @@ def pad_and_concatenate_features( "omitted from every sample" ) max_length = max(int(feature["input_ids"].shape[-1]) for feature in features) + if pad_to is not None: + if max_length > int(pad_to): + raise ValueError( + f"sample length {max_length} exceeds the static batch length {pad_to}" + ) + max_length = int(pad_to) import torch diff --git a/specforge/algorithms/common/dflash_family_model.py b/specforge/algorithms/common/dflash_family_model.py index 719b12232..aa040be3c 100644 --- a/specforge/algorithms/common/dflash_family_model.py +++ b/specforge/algorithms/common/dflash_family_model.py @@ -362,6 +362,7 @@ def __init__( block_size: int = 16, attention_backend: str = "flex_attention", num_anchors: int = 512, + static_anchor_count: bool = False, loss_decay_gamma: Optional[float] = None, objective_chunk_blocks: int = 128, loss_type: str = "dflash", @@ -405,6 +406,10 @@ def __init__( self.mask_token_id = mask_token_id self.attention_backend = attention_backend self.num_anchors = num_anchors + # ``training.static_shapes``: always sample ``num_anchors`` slots so the + # draft-block input keeps one shape; slots beyond the valid anchors of a + # row are masked exactly like today's short rows. + self.static_anchor_count = bool(static_anchor_count) self.loss_decay_gamma = loss_decay_gamma self.objective_chunk_blocks = int(objective_chunk_blocks) self.loss_type = loss_type @@ -461,11 +466,16 @@ def _sample_anchor_positions( # Training strategies pass the CPU-computed value and avoid this # synchronizing fallback on CUDA. max_valid_anchors = int(valid_counts.max().item()) - width = min(self.num_anchors, max(0, int(max_valid_anchors))) - if width == 0: + max_valid = max(0, int(max_valid_anchors)) + if max_valid == 0: raise ValueError( "DFlash-family training requires two consecutive supervised tokens" ) + # ``getattr``: unit tests drive this sampler with bare stand-ins. + if getattr(self, "static_anchor_count", False): + width = self.num_anchors + else: + width = min(self.num_anchors, max_valid) random_values = torch.rand(valid.shape, device=device) random_values.masked_fill_(~valid, 2.0) @@ -1830,6 +1840,7 @@ def __init__( block_size: int = 16, attention_backend: str = "flex_attention", num_anchors: int = 512, + static_anchor_count: bool = False, loss_decay_gamma: Optional[float] = None, objective_chunk_blocks: int = 128, shift_label: bool = False, @@ -1842,6 +1853,7 @@ def __init__( block_size=block_size, attention_backend=attention_backend, num_anchors=num_anchors, + static_anchor_count=static_anchor_count, loss_decay_gamma=loss_decay_gamma, objective_chunk_blocks=objective_chunk_blocks, loss_type="dflash", @@ -2117,6 +2129,7 @@ def __init__( block_size: int = 16, attention_backend: str = "flex_attention", num_anchors: int = 512, + static_anchor_count: bool = False, loss_decay_gamma: Optional[float] = None, dspark_ce_loss_alpha: float = 0.1, dspark_l1_loss_alpha: float = 0.9, @@ -2131,6 +2144,7 @@ def __init__( block_size=block_size, attention_backend=attention_backend, num_anchors=num_anchors, + static_anchor_count=static_anchor_count, loss_decay_gamma=loss_decay_gamma, objective_chunk_blocks=objective_chunk_blocks, loss_type="dflash", diff --git a/specforge/algorithms/common/hidden_states_data.py b/specforge/algorithms/common/hidden_states_data.py index 82d7fd63d..ffae28f2f 100644 --- a/specforge/algorithms/common/hidden_states_data.py +++ b/specforge/algorithms/common/hidden_states_data.py @@ -153,7 +153,7 @@ def build_dspark_offline_normalizer(max_len, **_topology): return partial(normalize_dspark_offline_sample, max_len=max_len) -def _padded_collator(required_keys, optional_keys=()): +def _padded_collator(required_keys, optional_keys=(), pad_to=None): """Build a collator that zero-pads every listed key along the sequence axis.""" sequence_axes = {key: 1 for key in (*required_keys, *optional_keys)} @@ -164,23 +164,26 @@ def collate(features): sequence_axes=sequence_axes, required_keys=required_keys, optional_keys=optional_keys, + pad_to=pad_to, ) return collate -def build_collator(): +def build_collator(pad_to=None): # The target's final hidden state rides along when the capture layout # carries it; offline v1 DFlash and Domino batches omit it. return _padded_collator( ("input_ids", "loss_mask", "hidden_states"), optional_keys=("target_last_hidden_states",), + pad_to=pad_to, ) -def build_dspark_collator(): +def build_dspark_collator(pad_to=None): return _padded_collator( - ("input_ids", "loss_mask", "hidden_states", "target_last_hidden_states") + ("input_ids", "loss_mask", "hidden_states", "target_last_hidden_states"), + pad_to=pad_to, ) @@ -244,8 +247,10 @@ def build_mtp_offline_normalizer(max_len, **_topology): return partial(normalize_mtp_offline_sample, max_len=max_len) -def build_mtp_collator(): - return _padded_collator(("input_ids", "loss_mask", "target_last_hidden_states")) +def build_mtp_collator(pad_to=None): + return _padded_collator( + ("input_ids", "loss_mask", "target_last_hidden_states"), pad_to=pad_to + ) __all__ = [ diff --git a/specforge/algorithms/eagle3/data.py b/specforge/algorithms/eagle3/data.py index 7521fcfeb..90f36152a 100644 --- a/specforge/algorithms/eagle3/data.py +++ b/specforge/algorithms/eagle3/data.py @@ -79,15 +79,17 @@ def build_offline_normalizer( ) -def build_offline_collator(): +def build_offline_collator(pad_to=None): # This retained collator owns USP-aware padding. It can move here once the # distributed helper contracts are independent of specforge.data. from specforge.data.utils import DataCollatorWithPadding - return DataCollatorWithPadding() + return DataCollatorWithPadding(pad_to=pad_to) -def build_server_collator(): +def build_server_collator(pad_to=None): + # Server-streamed EAGLE3 batches already arrive with equal shapes (the + # concatenate contract), so a static length needs no extra padding here. from specforge.algorithms.common.collation import concatenate_features return concatenate_features diff --git a/specforge/algorithms/model_providers.py b/specforge/algorithms/model_providers.py index 3e154e948..00c813b40 100644 --- a/specforge/algorithms/model_providers.py +++ b/specforge/algorithms/model_providers.py @@ -364,6 +364,7 @@ def _build_dflash_family_model( "block_size": int(draft_model.block_size), "attention_backend": cfg.training.attention_backend, "num_anchors": cfg.training.num_anchors, + "static_anchor_count": cfg.training.static_shapes, "loss_decay_gamma": cfg.training.loss_decay_gamma, "objective_chunk_blocks": cfg.training.objective_chunk_blocks, } diff --git a/specforge/config/schema.py b/specforge/config/schema.py index 6ec592a7e..4875f2403 100644 --- a/specforge/config/schema.py +++ b/specforge/config/schema.py @@ -897,6 +897,8 @@ class TrainingConfig(StrictConfigModel): backend: Literal["fsdp", "fsdp2"] = "fsdp" #: ``torch.compile`` every draft block (or the EAGLE midlayer) in place before FSDP2 sharding. Requires ``backend: fsdp2``. compile_blocks: bool = False + #: Pad every micro-batch to ``data.max_length`` and every DFlash-family anchor set to ``num_anchors`` so the draft blocks see one input shape per run. Required by ``compile_blocks``. + static_shapes: bool = False 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) @@ -1005,6 +1007,11 @@ def _validate_training_shape(self): ) if self.compile_blocks and self.backend != "fsdp2": raise ValueError("training.compile_blocks requires training.backend=fsdp2") + if self.compile_blocks and not self.static_shapes: + raise ValueError( + "training.compile_blocks requires training.static_shapes=true: " + "compiled draft blocks need one input shape per run" + ) sp_size = self.sp_ulysses_size * self.sp_ring_size if self.attention_backend == "usp": if self.batch_size != 1: diff --git a/specforge/data/utils.py b/specforge/data/utils.py index dc6979808..2a1e12d46 100644 --- a/specforge/data/utils.py +++ b/specforge/data/utils.py @@ -34,7 +34,10 @@ class DataCollatorWithPadding: Datacollator that will dynamically pad the inputs for batching. """ - def __init__(self): + def __init__(self, pad_to=None): + # ``pad_to``: fixed batch length (``training.static_shapes``) instead of + # the longest sample of each batch. + self.pad_to = None if pad_to is None else int(pad_to) if torch.distributed.is_available() and torch.distributed.is_initialized(): self.sp_degree = torch.distributed.get_world_size(get_draft_sp_group()) self.ulysses_degree = torch.distributed.get_world_size( @@ -120,6 +123,13 @@ def __call__(self, features: List[Dict[str, Any]]) -> Dict[str, Any]: - loss_mask: torch.Tensor of shape (B, N) """ max_length = max(item["input_ids"].shape[1] for item in features) + if self.pad_to is not None: + if max_length > self.pad_to: + raise ValueError( + f"sample length {max_length} exceeds the static batch length " + f"{self.pad_to}" + ) + max_length = self.pad_to # pad for sequence parrel max_length = ( diff --git a/specforge/launch.py b/specforge/launch.py index ac2b6eff4..b53e7a4d7 100644 --- a/specforge/launch.py +++ b/specforge/launch.py @@ -169,16 +169,31 @@ def _offline_io( *, ttt_length: int, use_usp_preprocess: bool, + static_shapes: bool = False, ): """Resolve the algorithm-owned normalizer and collator for one modality.""" provider = algorithm.providers.offline_for(modality) - return provider.build_collator(), provider.build_normalizer( + collate_fn = _static_or_dynamic_collator(provider.build_collator, max_len if static_shapes else None) + return collate_fn, provider.build_normalizer( max_len, ttt_length=ttt_length, use_usp_preprocess=use_usp_preprocess, ) +def _static_or_dynamic_collator(build_collator, pad_to): + """``training.static_shapes`` asks the algorithm's collator for a fixed length.""" + if pad_to is None: + return build_collator() + try: + return build_collator(pad_to=int(pad_to)) + except TypeError as exc: + raise ValueError( + "training.static_shapes is not supported by this algorithm's collator " + f"({build_collator!r} takes no pad_to)" + ) from exc + + def _shard_offline_refs( refs, *, @@ -258,6 +273,7 @@ def _make_offline_eval_data_factory( ttt_length: int, use_usp_preprocess: bool, dataloader_num_workers: int, + static_shapes: bool = False, ): """Build a fresh re-iterable eval loader over the offline feature path.""" provider = algorithm.providers.offline_for(modality) @@ -267,6 +283,7 @@ def _make_offline_eval_data_factory( max_len, ttt_length=ttt_length, use_usp_preprocess=use_usp_preprocess, + static_shapes=static_shapes, ) eval_run_id = f"{run_id}-eval" refs = provider.build_reader( @@ -295,15 +312,27 @@ def build_loader(): return build_loader +def _static_pad_length(static_shapes: bool, max_len) -> Optional[int]: + if not static_shapes: + return None + if max_len is None: + raise ValueError("static_shapes needs max_len (data.max_length) to pad to") + return int(max_len) + + def _streaming_collate( algorithm: AlgorithmRegistration, modality: str, collate_fn, + *, + pad_to=None, ): """Resolve an algorithm-owned server-streaming collator.""" if collate_fn is not None: return collate_fn - return algorithm.providers.server_streaming_for(modality).build_collator() + return _static_or_dynamic_collator( + algorithm.providers.server_streaming_for(modality).build_collator, pad_to + ) def _resolve_metadata_store( @@ -559,6 +588,7 @@ def build_offline_runtime( output_dir: str, ttt_length: int = 7, max_len: int = 2048, + static_shapes: bool = False, batch_size: int = 1, accumulation_steps: int = 1, num_epochs: int = 1, @@ -595,6 +625,7 @@ def build_offline_runtime( max_len, ttt_length=ttt_length, use_usp_preprocess=use_usp_preprocess, + static_shapes=static_shapes, ) controller = DataFlowController( run_id, @@ -630,6 +661,7 @@ def refs_for_epoch(epoch): ttt_length=ttt_length, use_usp_preprocess=use_usp_preprocess, dataloader_num_workers=dataloader_num_workers, + static_shapes=static_shapes, ) return _assemble_trainer( algorithm=algorithm, @@ -691,6 +723,7 @@ def build_disagg_offline_runtime( output_dir: str, ttt_length: int = 7, max_len: int = 2048, + static_shapes: bool = False, batch_size: int = 1, accumulation_steps: int = 1, num_epochs: int = 1, @@ -726,6 +759,7 @@ def build_disagg_offline_runtime( max_len, ttt_length=ttt_length, use_usp_preprocess=use_usp_preprocess, + static_shapes=static_shapes, ) source_refs = list(refs) @@ -758,6 +792,7 @@ def refs_for_epoch(epoch): ttt_length=ttt_length, use_usp_preprocess=use_usp_preprocess, dataloader_num_workers=dataloader_num_workers, + static_shapes=static_shapes, ) return _assemble_trainer( algorithm=algorithm, @@ -1557,6 +1592,8 @@ def build_disagg_online_consumer( eval_interval: int = 0, eval_data_factory=None, collate_fn=None, + static_shapes: bool = False, + max_len: Optional[int] = None, idle_timeout_s: Optional[float] = None, metadata_store: Optional[MetadataStore] = None, metadata_db_path: Optional[str] = None, @@ -1912,7 +1949,12 @@ def stop_distributor_and_drain() -> None: eval_data_factory=eval_data_factory, logger=logger, log_interval=log_interval, - collate_fn=_streaming_collate(algorithm, modality, collate_fn), + collate_fn=_streaming_collate( + algorithm, + modality, + collate_fn, + pad_to=_static_pad_length(static_shapes, max_len), + ), strategy_kwargs=strategy_kwargs, per_sample_transform=None, max_checkpoints=max_checkpoints, diff --git a/specforge/training/DESIGN.md b/specforge/training/DESIGN.md index 74e3cf25c..ccbcee699 100644 --- a/specforge/training/DESIGN.md +++ b/specforge/training/DESIGN.md @@ -29,7 +29,7 @@ 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. `BackendOptions` carries opt-in -behaviors; `training.compile_blocks` compiles each draft block in place before FSDP2 sharding so the composable hooks stay outside Dynamo. +behaviors; `training.static_shapes` pads every micro-batch to `data.max_length` and every DFlash-family anchor set to `num_anchors`, so compiled graphs see one input shape (torch 2.13's Inductor cannot lower `flex_attention` with a symbolic context length); `training.compile_blocks` requires it and compiles each draft block in place before FSDP2 sharding so the composable hooks stay outside Dynamo. 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()`. diff --git a/specforge/training/assembly.py b/specforge/training/assembly.py index 2d772871d..af0bb5aa9 100644 --- a/specforge/training/assembly.py +++ b/specforge/training/assembly.py @@ -562,6 +562,7 @@ def _common_launch_kwargs( training_backend=t.backend, fsdp_sharding=t.fsdp_sharding, backend_options=_backend_options(cfg), + static_shapes=t.static_shapes, run_id=cfg.run_id, output_dir=cfg.output_dir, batch_size=t.batch_size, diff --git a/specforge/training/disaggregated.py b/specforge/training/disaggregated.py index ecfa00c73..1676227fe 100644 --- a/specforge/training/disaggregated.py +++ b/specforge/training/disaggregated.py @@ -542,6 +542,7 @@ def produce() -> int: training_backend=cfg.training.backend, fsdp_sharding=cfg.training.fsdp_sharding, backend_options=_backend_options(cfg), + static_shapes=cfg.training.static_shapes, run_id=cfg.run_id, output_dir=cfg.output_dir, ttt_length=cfg.training.ttt_length, @@ -825,6 +826,8 @@ def produce() -> int: training_backend=cfg.training.backend, fsdp_sharding=cfg.training.fsdp_sharding, backend_options=_backend_options(cfg), + static_shapes=cfg.training.static_shapes, + max_len=cfg.data.max_length, run_id=cfg.run_id, output_dir=cfg.output_dir, batch_size=cfg.training.batch_size, diff --git a/tests/test_runtime/test_fsdp2_compile_blocks.py b/tests/test_runtime/test_fsdp2_compile_blocks.py index 1a280197a..61f6237d2 100644 --- a/tests/test_runtime/test_fsdp2_compile_blocks.py +++ b/tests/test_runtime/test_fsdp2_compile_blocks.py @@ -28,13 +28,16 @@ def test_fsdp1_rejects_compile_blocks(self): with self.assertRaisesRegex(ValueError, "compile_blocks"): backend.prepare_model(TinyComposite(), optimizer_target=None) - def test_config_requires_fsdp2(self): + def test_config_requires_fsdp2_and_static_shapes(self): from specforge.config.schema import TrainingConfig with self.assertRaisesRegex(ValueError, "compile_blocks"): - TrainingConfig(compile_blocks=True) - cfg = TrainingConfig(backend="fsdp2", compile_blocks=True) + TrainingConfig(compile_blocks=True, static_shapes=True) + with self.assertRaisesRegex(ValueError, "static_shapes"): + TrainingConfig(backend="fsdp2", compile_blocks=True) + cfg = TrainingConfig(backend="fsdp2", compile_blocks=True, static_shapes=True) self.assertTrue(cfg.compile_blocks) + self.assertTrue(cfg.static_shapes) def test_block_targets_fall_back_to_midlayer(self): from specforge.training.backend import DistributedTrainingBackend @@ -120,6 +123,7 @@ def _training(): "max_steps": 1, "backend": "fsdp2", "compile_blocks": True, + "static_shapes": True, } def test_online_consumer_receives_the_option(self): @@ -160,6 +164,8 @@ def test_online_consumer_receives_the_option(self): options = build.call_args.kwargs["backend_options"] self.assertIsInstance(options, BackendOptions) self.assertTrue(options.compile_blocks) + self.assertIs(build.call_args.kwargs["static_shapes"], True) + self.assertEqual(build.call_args.kwargs["max_len"], 2048) def test_offline_consumer_receives_the_option(self): from specforge.algorithms.builtin import builtin_algorithm_registry @@ -195,6 +201,7 @@ def test_offline_consumer_receives_the_option(self): options = build.call_args.kwargs["backend_options"] self.assertIsInstance(options, BackendOptions) self.assertTrue(options.compile_blocks) + self.assertIs(build.call_args.kwargs["static_shapes"], True) class _FakeFitTrainer: diff --git a/tests/test_runtime/test_static_shapes.py b/tests/test_runtime/test_static_shapes.py new file mode 100644 index 000000000..18382ff49 --- /dev/null +++ b/tests/test_runtime/test_static_shapes.py @@ -0,0 +1,165 @@ +"""``training.static_shapes``: one batch length and one anchor count per run, for compiled blocks.""" + +import types +import unittest + +import torch + +from specforge.algorithms.builtin import builtin_algorithm_registry + +ALGORITHM = builtin_algorithm_registry().resolve("dflash") + + +def _dflash_sample(length, width=4): + return { + "input_ids": torch.arange(length).unsqueeze(0), + "loss_mask": torch.ones(1, length, dtype=torch.long), + "hidden_states": torch.randn(1, length, width), + } + + +class TestPadToCollation(unittest.TestCase): + def test_pad_and_concatenate_pads_to_the_static_length(self): + from specforge.algorithms.common.collation import pad_and_concatenate_features + + batch = pad_and_concatenate_features( + [_dflash_sample(5), _dflash_sample(3)], + sequence_axes={"input_ids": 1, "loss_mask": 1, "hidden_states": 1}, + required_keys=("input_ids", "loss_mask", "hidden_states"), + pad_to=8, + ) + self.assertEqual(tuple(batch["input_ids"].shape), (2, 8)) + self.assertEqual(tuple(batch["hidden_states"].shape), (2, 8, 4)) + self.assertEqual(int(batch["loss_mask"][1, 3:].sum()), 0) + + def test_pad_to_rejects_longer_samples(self): + from specforge.algorithms.common.collation import pad_and_concatenate_features + + with self.assertRaisesRegex(ValueError, "static batch length"): + pad_and_concatenate_features( + [_dflash_sample(5)], + sequence_axes={"input_ids": 1, "loss_mask": 1, "hidden_states": 1}, + required_keys=("input_ids", "loss_mask", "hidden_states"), + pad_to=4, + ) + + def test_dflash_family_collators_accept_pad_to(self): + from specforge.algorithms.common.hidden_states_data import ( + build_collator, + build_dspark_collator, + ) + + batch = build_collator(pad_to=8)([_dflash_sample(5), _dflash_sample(3)]) + self.assertEqual(tuple(batch["input_ids"].shape), (2, 8)) + dspark = [ + dict(_dflash_sample(5), target_last_hidden_states=torch.randn(1, 5, 4)), + dict(_dflash_sample(2), target_last_hidden_states=torch.randn(1, 2, 4)), + ] + batch = build_dspark_collator(pad_to=8)(dspark) + self.assertEqual(tuple(batch["target_last_hidden_states"].shape), (2, 8, 4)) + # The default stays pad-to-longest. + batch = build_collator()([_dflash_sample(5), _dflash_sample(3)]) + self.assertEqual(tuple(batch["input_ids"].shape), (2, 5)) + + def test_eagle3_offline_collator_accepts_pad_to(self): + from specforge.data.utils import DataCollatorWithPadding + + def item(n): + ones = torch.ones(1, n, dtype=torch.long) + return {"input_ids": ones, "attention_mask": ones.clone(), "loss_mask": ones.clone()} + + batch = DataCollatorWithPadding(pad_to=8)([item(5), item(3)]) + self.assertEqual(tuple(batch["input_ids"].shape), (2, 8)) + self.assertEqual(tuple(batch["loss_mask"].shape), (2, 8)) + with self.assertRaisesRegex(ValueError, "static batch length"): + DataCollatorWithPadding(pad_to=4)([item(5)]) + + +class TestStaticAnchorCount(unittest.TestCase): + """``num_anchors`` slots are always sampled; slots past a row's valid anchors are masked.""" + + @staticmethod + def _mask(): + # Row 0 supervises positions 0-2 (anchors 0, 1 valid); row 1 positions 0-3 (anchors 0-2). + mask = torch.zeros(2, 6) + mask[0, :3] = 1 + mask[1, :4] = 1 + return mask + + def _sample(self, static, seed=0): + from specforge.algorithms.common.dflash_family_model import OnlineDFlashModel + + model = types.SimpleNamespace(num_anchors=4, static_anchor_count=static) + torch.manual_seed(seed) + return OnlineDFlashModel._sample_anchor_positions( + model, 6, self._mask(), torch.device("cpu"), max_valid_anchors=3 + ) + + def test_static_width_is_num_anchors_and_extra_slots_are_masked(self): + anchors, keep = self._sample(static=True) + self.assertEqual(tuple(anchors.shape), (2, 4)) + self.assertEqual(keep.sum(dim=1).tolist(), [2, 3]) + dyn_anchors, dyn_keep = self._sample(static=False) + self.assertEqual(tuple(dyn_anchors.shape), (2, 3)) + for row in range(2): + self.assertEqual( + sorted(anchors[row][keep[row]].tolist()), + sorted(dyn_anchors[row][dyn_keep[row]].tolist()), + ) + + def test_no_valid_anchor_still_raises(self): + from specforge.algorithms.common.dflash_family_model import OnlineDFlashModel + + model = types.SimpleNamespace(num_anchors=4, static_anchor_count=True) + with self.assertRaisesRegex(ValueError, "consecutive supervised"): + OnlineDFlashModel._sample_anchor_positions( + model, 6, torch.zeros(1, 6), torch.device("cpu"), max_valid_anchors=0 + ) + + +class TestLaunchStaticCollators(unittest.TestCase): + def test_static_or_dynamic_collator(self): + from specforge.launch import _static_or_dynamic_collator, _static_pad_length + + self.assertEqual(_static_or_dynamic_collator(lambda: "dynamic", None), "dynamic") + self.assertEqual(_static_or_dynamic_collator(lambda pad_to=None: pad_to, 8), 8) + with self.assertRaisesRegex(ValueError, "static_shapes"): + _static_or_dynamic_collator(lambda: "no pad_to", 8) + self.assertIsNone(_static_pad_length(False, None)) + self.assertEqual(_static_pad_length(True, 2048), 2048) + with self.assertRaisesRegex(ValueError, "max_len"): + _static_pad_length(True, None) + + def test_streaming_and_offline_collators_pad_to_max_len(self): + from specforge.launch import _offline_io, _streaming_collate + + collate = _streaming_collate(ALGORITHM, "text", None, pad_to=8) + batch = collate([_dflash_sample(5), _dflash_sample(3)]) + self.assertEqual(tuple(batch["hidden_states"].shape), (2, 8, 4)) + collate, _ = _offline_io( + ALGORITHM, "text", 8, ttt_length=7, use_usp_preprocess=False, static_shapes=True + ) + self.assertEqual( + tuple(collate([_dflash_sample(5), _dflash_sample(3)])["input_ids"].shape), (2, 8) + ) + collate, _ = _offline_io(ALGORITHM, "text", 8, ttt_length=7, use_usp_preprocess=False) + self.assertEqual( + tuple(collate([_dflash_sample(5), _dflash_sample(3)])["input_ids"].shape), (2, 5) + ) + + +class TestStaticShapesConfig(unittest.TestCase): + def test_compile_blocks_requires_static_shapes(self): + from specforge.config.schema import TrainingConfig + + self.assertFalse(TrainingConfig().static_shapes) + with self.assertRaisesRegex(ValueError, "static_shapes"): + TrainingConfig(backend="fsdp2", compile_blocks=True) + cfg = TrainingConfig(backend="fsdp2", compile_blocks=True, static_shapes=True) + self.assertTrue(cfg.static_shapes) + # static_shapes alone is allowed on either backend. + self.assertTrue(TrainingConfig(static_shapes=True).static_shapes) + + +if __name__ == "__main__": + unittest.main() From 63112683cd6a30ab35e68aa5149a772b20179310 Mon Sep 17 00:00:00 2001 From: "Ethan (Yusheng) Su" Date: Sun, 4 Oct 2026 19:30:23 +0900 Subject: [PATCH 7/9] Warn instead of rejecting compile_blocks without static_shapes Fixed-shape inputs (synthetic features, long-form data truncated at max_length) run compiled blocks fine without the padding that static_shapes adds, and the crash on changing shapes is a torch 2.13 + flex_attention limitation rather than a contract of the option. The validator now warns, naming the failure mode, instead of raising. --- specforge/config/schema.py | 13 +++++++++---- specforge/training/DESIGN.md | 2 +- tests/test_runtime/test_fsdp2_compile_blocks.py | 4 ++-- tests/test_runtime/test_static_shapes.py | 8 +++++--- 4 files changed, 17 insertions(+), 10 deletions(-) diff --git a/specforge/config/schema.py b/specforge/config/schema.py index 4875f2403..163a968fb 100644 --- a/specforge/config/schema.py +++ b/specforge/config/schema.py @@ -20,6 +20,7 @@ import json import os import re +import warnings from typing import Dict, List, Literal, Mapping, Optional from urllib.parse import urlparse @@ -897,7 +898,7 @@ class TrainingConfig(StrictConfigModel): backend: Literal["fsdp", "fsdp2"] = "fsdp" #: ``torch.compile`` every draft block (or the EAGLE midlayer) in place before FSDP2 sharding. Requires ``backend: fsdp2``. compile_blocks: bool = False - #: Pad every micro-batch to ``data.max_length`` and every DFlash-family anchor set to ``num_anchors`` so the draft blocks see one input shape per run. Required by ``compile_blocks``. + #: Pad every micro-batch to ``data.max_length`` and every DFlash-family anchor set to ``num_anchors`` so the draft blocks see one input shape per run. Recommended with ``compile_blocks`` (a warning without it). static_shapes: bool = False num_epochs: int = Field(default=1, gt=0) max_steps: Optional[int] = Field(default=None, gt=0) @@ -1008,9 +1009,13 @@ def _validate_training_shape(self): if self.compile_blocks and self.backend != "fsdp2": raise ValueError("training.compile_blocks requires training.backend=fsdp2") if self.compile_blocks and not self.static_shapes: - raise ValueError( - "training.compile_blocks requires training.static_shapes=true: " - "compiled draft blocks need one input shape per run" + warnings.warn( + "training.compile_blocks without training.static_shapes needs inputs " + "whose padded length and anchor count never change; with pad-to-longest " + "batches the compiled blocks recompile and, on torch 2.13 with " + "flex_attention, fail inside Inductor. Set training.static_shapes=true " + "unless your batches are already fixed-shape.", + stacklevel=2, ) sp_size = self.sp_ulysses_size * self.sp_ring_size if self.attention_backend == "usp": diff --git a/specforge/training/DESIGN.md b/specforge/training/DESIGN.md index ccbcee699..8e5e62826 100644 --- a/specforge/training/DESIGN.md +++ b/specforge/training/DESIGN.md @@ -29,7 +29,7 @@ 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. `BackendOptions` carries opt-in -behaviors; `training.static_shapes` pads every micro-batch to `data.max_length` and every DFlash-family anchor set to `num_anchors`, so compiled graphs see one input shape (torch 2.13's Inductor cannot lower `flex_attention` with a symbolic context length); `training.compile_blocks` requires it and compiles each draft block in place before FSDP2 sharding so the composable hooks stay outside Dynamo. +behaviors; `training.static_shapes` pads every micro-batch to `data.max_length` and every DFlash-family anchor set to `num_anchors`, so compiled graphs see one input shape (torch 2.13's Inductor cannot lower `flex_attention` with a symbolic context length); `training.compile_blocks` warns without it and compiles each draft block in place before FSDP2 sharding so the composable hooks stay outside Dynamo. 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()`. diff --git a/tests/test_runtime/test_fsdp2_compile_blocks.py b/tests/test_runtime/test_fsdp2_compile_blocks.py index 61f6237d2..9beb4be41 100644 --- a/tests/test_runtime/test_fsdp2_compile_blocks.py +++ b/tests/test_runtime/test_fsdp2_compile_blocks.py @@ -28,12 +28,12 @@ def test_fsdp1_rejects_compile_blocks(self): with self.assertRaisesRegex(ValueError, "compile_blocks"): backend.prepare_model(TinyComposite(), optimizer_target=None) - def test_config_requires_fsdp2_and_static_shapes(self): + def test_config_requires_fsdp2_and_recommends_static_shapes(self): from specforge.config.schema import TrainingConfig with self.assertRaisesRegex(ValueError, "compile_blocks"): TrainingConfig(compile_blocks=True, static_shapes=True) - with self.assertRaisesRegex(ValueError, "static_shapes"): + with self.assertWarnsRegex(UserWarning, "static_shapes"): TrainingConfig(backend="fsdp2", compile_blocks=True) cfg = TrainingConfig(backend="fsdp2", compile_blocks=True, static_shapes=True) self.assertTrue(cfg.compile_blocks) diff --git a/tests/test_runtime/test_static_shapes.py b/tests/test_runtime/test_static_shapes.py index 18382ff49..be3419750 100644 --- a/tests/test_runtime/test_static_shapes.py +++ b/tests/test_runtime/test_static_shapes.py @@ -149,12 +149,14 @@ def test_streaming_and_offline_collators_pad_to_max_len(self): class TestStaticShapesConfig(unittest.TestCase): - def test_compile_blocks_requires_static_shapes(self): + def test_compile_blocks_without_static_shapes_warns(self): from specforge.config.schema import TrainingConfig self.assertFalse(TrainingConfig().static_shapes) - with self.assertRaisesRegex(ValueError, "static_shapes"): - TrainingConfig(backend="fsdp2", compile_blocks=True) + with self.assertWarnsRegex(UserWarning, "static_shapes"): + cfg = TrainingConfig(backend="fsdp2", compile_blocks=True) + self.assertTrue(cfg.compile_blocks) + self.assertFalse(cfg.static_shapes) cfg = TrainingConfig(backend="fsdp2", compile_blocks=True, static_shapes=True) self.assertTrue(cfg.static_shapes) # static_shapes alone is allowed on either backend. From 7ce2eebc21524c9a3daec8d84516e78d363fba72 Mon Sep 17 00:00:00 2001 From: "Ethan (Yusheng) Su" Date: Sun, 4 Oct 2026 21:47:05 +0900 Subject: [PATCH 8/9] static_shapes: tolerate an anchor count above the candidate positions With a fixed anchor count a short batch (a 512-token bucket has 511 candidate positions, num_anchors defaults to 512) made the sampler take more columns than exist and fail with a size mismatch. The missing slots are now sentinels and masked like any other empty anchor slot. --- .../algorithms/common/dflash_family_model.py | 9 +++++++-- tests/test_runtime/test_static_shapes.py | 15 +++++++++++++++ 2 files changed, 22 insertions(+), 2 deletions(-) diff --git a/specforge/algorithms/common/dflash_family_model.py b/specforge/algorithms/common/dflash_family_model.py index aa040be3c..767cefc08 100644 --- a/specforge/algorithms/common/dflash_family_model.py +++ b/specforge/algorithms/common/dflash_family_model.py @@ -479,12 +479,17 @@ def _sample_anchor_positions( random_values = torch.rand(valid.shape, device=device) random_values.masked_fill_(~valid, 2.0) - candidates = random_values.argsort(dim=1)[:, :width] + sentinel = valid.shape[1] + # A static anchor count can exceed the candidate positions of a short + # (or short-bucketed) batch; the missing slots are sentinels, masked below. + take = min(width, num_candidates) + candidates = random_values.argsort(dim=1)[:, :take] + if take < width: + candidates = torch.nn.functional.pad(candidates, (0, width - take), value=sentinel) keep_mask = torch.arange(width, device=device).unsqueeze( 0 ) < valid_counts.clamp(max=width).unsqueeze(1) - sentinel = valid.shape[1] anchors = torch.where( keep_mask, candidates, diff --git a/tests/test_runtime/test_static_shapes.py b/tests/test_runtime/test_static_shapes.py index be3419750..e7ee14923 100644 --- a/tests/test_runtime/test_static_shapes.py +++ b/tests/test_runtime/test_static_shapes.py @@ -107,6 +107,21 @@ def test_static_width_is_num_anchors_and_extra_slots_are_masked(self): sorted(dyn_anchors[row][dyn_keep[row]].tolist()), ) + def test_static_width_beyond_the_candidate_positions(self): + from specforge.algorithms.common.dflash_family_model import OnlineDFlashModel + + # A 512-token bucket has 511 candidate positions but num_anchors may be 512. + model = types.SimpleNamespace(num_anchors=8, static_anchor_count=True) + mask = torch.ones(2, 6) + torch.manual_seed(0) + anchors, keep = OnlineDFlashModel._sample_anchor_positions( + model, 6, mask, torch.device("cpu"), max_valid_anchors=5 + ) + self.assertEqual(tuple(anchors.shape), (2, 8)) + self.assertEqual(keep.sum(dim=1).tolist(), [5, 5]) + self.assertTrue(bool((anchors[~keep] == 0).all())) + self.assertEqual(sorted(anchors[0][keep[0]].tolist()), [0, 1, 2, 3, 4]) + def test_no_valid_anchor_still_raises(self): from specforge.algorithms.common.dflash_family_model import OnlineDFlashModel From 5cb86277b3323759d891c12b08988f592ae3166f Mon Sep 17 00:00:00 2001 From: "Ethan (Yusheng) Su" Date: Sun, 4 Oct 2026 19:47:46 +0900 Subject: [PATCH 9/9] Add training.static_shape_buckets: a few static lengths for compiled blocks static_shapes pads every micro-batch to data.max_length. On corpora whose conversations are much shorter than max_length that wastes positions (the Qwen3.8-27B regeneration corpus averages 4-4.5k tokens under max_length 8192). static_shape_buckets lets the collators pad to the smallest of a few lengths instead (multiples of 128 so flex_attention blocks and float8 token counts stay aligned; data.max_length is always the last bucket); the anchor count stays fixed at num_anchors. Compiled blocks must then hold one static graph per bucket: the backend compiles them with dynamic=False and raises Dynamo's recompile limit to cover two traces per bucket (first block without grad on its input, the rest with), so no bucket change ever turns a dimension symbolic. Bucket padding without compile_blocks is a pure data-path setting and works on either backend. --- specforge/algorithms/common/collation.py | 33 ++-- specforge/config/schema.py | 20 +++ specforge/data/utils.py | 11 +- specforge/launch.py | 26 +++- specforge/training/DESIGN.md | 2 +- specforge/training/assembly.py | 16 +- specforge/training/backend.py | 4 + specforge/training/disaggregated.py | 4 + specforge/training/fsdp2.py | 31 +++- tests/test_runtime/test_shape_buckets.py | 182 +++++++++++++++++++++++ 10 files changed, 305 insertions(+), 24 deletions(-) create mode 100644 tests/test_runtime/test_shape_buckets.py diff --git a/specforge/algorithms/common/collation.py b/specforge/algorithms/common/collation.py index e8f8c120f..1d7f95f94 100644 --- a/specforge/algorithms/common/collation.py +++ b/specforge/algorithms/common/collation.py @@ -2,7 +2,9 @@ from __future__ import annotations -from typing import Mapping, Optional, Sequence +from typing import Mapping, Optional, Sequence, Union + +PadTo = Union[int, Sequence[int]] def concatenate_features(features): @@ -21,13 +23,31 @@ def concatenate_features(features): } +def resolve_static_length(longest: int, pad_to) -> int: + """Length a batch is padded to under ``training.static_shapes``. + + ``pad_to`` is one fixed length, or an ascending sequence of bucket lengths + (``training.static_shape_buckets``): the smallest bucket that fits the + longest sample wins. A sample longer than the (largest) length raises. + """ + if pad_to is None: + return int(longest) + buckets = [int(pad_to)] if isinstance(pad_to, int) else sorted(int(b) for b in pad_to) + for bucket in buckets: + if longest <= bucket: + return bucket + raise ValueError( + f"sample length {longest} exceeds the static batch length {buckets[-1]}" + ) + + def pad_and_concatenate_features( features, *, sequence_axes: Mapping[str, int], required_keys: Sequence[str], optional_keys: Sequence[str] = (), - pad_to: Optional[int] = None, + pad_to: Optional[PadTo] = None, ): """Zero-pad configured tensor axes to the longest input sequence. @@ -59,12 +79,7 @@ def pad_and_concatenate_features( "omitted from every sample" ) max_length = max(int(feature["input_ids"].shape[-1]) for feature in features) - if pad_to is not None: - if max_length > int(pad_to): - raise ValueError( - f"sample length {max_length} exceeds the static batch length {pad_to}" - ) - max_length = int(pad_to) + max_length = resolve_static_length(max_length, pad_to) import torch @@ -92,4 +107,4 @@ def pad_and_concatenate_features( return batch -__all__ = ["concatenate_features", "pad_and_concatenate_features"] +__all__ = ["concatenate_features", "pad_and_concatenate_features", "resolve_static_length"] diff --git a/specforge/config/schema.py b/specforge/config/schema.py index 163a968fb..8765c3a57 100644 --- a/specforge/config/schema.py +++ b/specforge/config/schema.py @@ -900,6 +900,8 @@ class TrainingConfig(StrictConfigModel): compile_blocks: bool = False #: Pad every micro-batch to ``data.max_length`` and every DFlash-family anchor set to ``num_anchors`` so the draft blocks see one input shape per run. Recommended with ``compile_blocks`` (a warning without it). static_shapes: bool = False + #: With ``static_shapes``: pad each micro-batch to the smallest of these lengths that fits instead of always to ``data.max_length`` (``max_length`` is always the last bucket). Multiples of 128, ascending; compiled blocks get one static graph per bucket. + static_shape_buckets: Optional[List[int]] = None 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) @@ -1008,6 +1010,19 @@ def _validate_training_shape(self): ) if self.compile_blocks and self.backend != "fsdp2": raise ValueError("training.compile_blocks requires training.backend=fsdp2") + if self.static_shape_buckets is not None: + if not self.static_shapes: + raise ValueError("training.static_shape_buckets requires training.static_shapes=true") + buckets = list(self.static_shape_buckets) + if not buckets: + raise ValueError("training.static_shape_buckets must not be empty") + if any(b <= 0 or b % 128 for b in buckets): + raise ValueError( + "training.static_shape_buckets must be positive multiples of 128 " + "(flex_attention block granularity; also keeps float8 token counts %% 16 == 0)" + ) + if any(b2 <= b1 for b1, b2 in zip(buckets, buckets[1:])): + raise ValueError("training.static_shape_buckets must be strictly increasing") if self.compile_blocks and not self.static_shapes: warnings.warn( "training.compile_blocks without training.static_shapes needs inputs " @@ -1130,6 +1145,11 @@ def _default_role_for_deployment(cls, values): def _validate_run_structure(self): """Validate topology and cross-field shape without resolving algorithms.""" mode = self.mode + if self.training.static_shape_buckets and max(self.training.static_shape_buckets) > self.data.max_length: + raise ValueError( + "training.static_shape_buckets must not exceed data.max_length " + f"({self.data.max_length}); the largest bucket is always data.max_length" + ) deployment = self.deployment.mode role = self.training.role diff --git a/specforge/data/utils.py b/specforge/data/utils.py index 2a1e12d46..6638fcdcd 100644 --- a/specforge/data/utils.py +++ b/specforge/data/utils.py @@ -37,7 +37,7 @@ class DataCollatorWithPadding: def __init__(self, pad_to=None): # ``pad_to``: fixed batch length (``training.static_shapes``) instead of # the longest sample of each batch. - self.pad_to = None if pad_to is None else int(pad_to) + self.pad_to = None if pad_to is None else (int(pad_to) if isinstance(pad_to, int) else tuple(int(b) for b in pad_to)) if torch.distributed.is_available() and torch.distributed.is_initialized(): self.sp_degree = torch.distributed.get_world_size(get_draft_sp_group()) self.ulysses_degree = torch.distributed.get_world_size( @@ -124,12 +124,9 @@ def __call__(self, features: List[Dict[str, Any]]) -> Dict[str, Any]: """ max_length = max(item["input_ids"].shape[1] for item in features) if self.pad_to is not None: - if max_length > self.pad_to: - raise ValueError( - f"sample length {max_length} exceeds the static batch length " - f"{self.pad_to}" - ) - max_length = self.pad_to + from specforge.algorithms.common.collation import resolve_static_length + + max_length = resolve_static_length(max_length, self.pad_to) # pad for sequence parrel max_length = ( diff --git a/specforge/launch.py b/specforge/launch.py index b53e7a4d7..19d130c16 100644 --- a/specforge/launch.py +++ b/specforge/launch.py @@ -170,10 +170,13 @@ def _offline_io( ttt_length: int, use_usp_preprocess: bool, static_shapes: bool = False, + static_shape_buckets=None, ): """Resolve the algorithm-owned normalizer and collator for one modality.""" provider = algorithm.providers.offline_for(modality) - collate_fn = _static_or_dynamic_collator(provider.build_collator, max_len if static_shapes else None) + collate_fn = _static_or_dynamic_collator( + provider.build_collator, _static_pad_length(static_shapes, max_len, static_shape_buckets) + ) return collate_fn, provider.build_normalizer( max_len, ttt_length=ttt_length, @@ -186,7 +189,7 @@ def _static_or_dynamic_collator(build_collator, pad_to): if pad_to is None: return build_collator() try: - return build_collator(pad_to=int(pad_to)) + return build_collator(pad_to=pad_to) except TypeError as exc: raise ValueError( "training.static_shapes is not supported by this algorithm's collator " @@ -274,6 +277,7 @@ def _make_offline_eval_data_factory( use_usp_preprocess: bool, dataloader_num_workers: int, static_shapes: bool = False, + static_shape_buckets=None, ): """Build a fresh re-iterable eval loader over the offline feature path.""" provider = algorithm.providers.offline_for(modality) @@ -284,6 +288,7 @@ def _make_offline_eval_data_factory( ttt_length=ttt_length, use_usp_preprocess=use_usp_preprocess, static_shapes=static_shapes, + static_shape_buckets=static_shape_buckets, ) eval_run_id = f"{run_id}-eval" refs = provider.build_reader( @@ -312,11 +317,17 @@ def build_loader(): return build_loader -def _static_pad_length(static_shapes: bool, max_len) -> Optional[int]: +def _static_pad_length(static_shapes: bool, max_len, buckets=None): + """``None`` (pad to the longest sample), one fixed length, or ascending bucket lengths.""" if not static_shapes: return None if max_len is None: raise ValueError("static_shapes needs max_len (data.max_length) to pad to") + if buckets: + lengths = sorted({int(b) for b in buckets} | {int(max_len)}) + if lengths[-1] != int(max_len): + raise ValueError("static_shape_buckets must not exceed max_len (data.max_length)") + return tuple(lengths) return int(max_len) @@ -589,6 +600,7 @@ def build_offline_runtime( ttt_length: int = 7, max_len: int = 2048, static_shapes: bool = False, + static_shape_buckets=None, batch_size: int = 1, accumulation_steps: int = 1, num_epochs: int = 1, @@ -626,6 +638,7 @@ def build_offline_runtime( ttt_length=ttt_length, use_usp_preprocess=use_usp_preprocess, static_shapes=static_shapes, + static_shape_buckets=static_shape_buckets, ) controller = DataFlowController( run_id, @@ -662,6 +675,7 @@ def refs_for_epoch(epoch): use_usp_preprocess=use_usp_preprocess, dataloader_num_workers=dataloader_num_workers, static_shapes=static_shapes, + static_shape_buckets=static_shape_buckets, ) return _assemble_trainer( algorithm=algorithm, @@ -724,6 +738,7 @@ def build_disagg_offline_runtime( ttt_length: int = 7, max_len: int = 2048, static_shapes: bool = False, + static_shape_buckets=None, batch_size: int = 1, accumulation_steps: int = 1, num_epochs: int = 1, @@ -760,6 +775,7 @@ def build_disagg_offline_runtime( ttt_length=ttt_length, use_usp_preprocess=use_usp_preprocess, static_shapes=static_shapes, + static_shape_buckets=static_shape_buckets, ) source_refs = list(refs) @@ -793,6 +809,7 @@ def refs_for_epoch(epoch): use_usp_preprocess=use_usp_preprocess, dataloader_num_workers=dataloader_num_workers, static_shapes=static_shapes, + static_shape_buckets=static_shape_buckets, ) return _assemble_trainer( algorithm=algorithm, @@ -1593,6 +1610,7 @@ def build_disagg_online_consumer( eval_data_factory=None, collate_fn=None, static_shapes: bool = False, + static_shape_buckets=None, max_len: Optional[int] = None, idle_timeout_s: Optional[float] = None, metadata_store: Optional[MetadataStore] = None, @@ -1953,7 +1971,7 @@ def stop_distributor_and_drain() -> None: algorithm, modality, collate_fn, - pad_to=_static_pad_length(static_shapes, max_len), + pad_to=_static_pad_length(static_shapes, max_len, static_shape_buckets), ), strategy_kwargs=strategy_kwargs, per_sample_transform=None, diff --git a/specforge/training/DESIGN.md b/specforge/training/DESIGN.md index 8e5e62826..9ba8259de 100644 --- a/specforge/training/DESIGN.md +++ b/specforge/training/DESIGN.md @@ -29,7 +29,7 @@ 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. `BackendOptions` carries opt-in -behaviors; `training.static_shapes` pads every micro-batch to `data.max_length` and every DFlash-family anchor set to `num_anchors`, so compiled graphs see one input shape (torch 2.13's Inductor cannot lower `flex_attention` with a symbolic context length); `training.compile_blocks` warns without it and compiles each draft block in place before FSDP2 sharding so the composable hooks stay outside Dynamo. +behaviors; `training.static_shapes` pads every micro-batch to `data.max_length` and every DFlash-family anchor set to `num_anchors`, so compiled graphs see one input shape (torch 2.13's Inductor cannot lower `flex_attention` with a symbolic context length); `training.static_shape_buckets` refines that into a few lengths (multiples of 128, `max_length` last) with one static graph per bucket; `training.compile_blocks` warns without `static_shapes` and compiles each draft block in place before FSDP2 sharding so the composable hooks stay outside Dynamo. 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()`. diff --git a/specforge/training/assembly.py b/specforge/training/assembly.py index af0bb5aa9..2b47cad77 100644 --- a/specforge/training/assembly.py +++ b/specforge/training/assembly.py @@ -536,9 +536,22 @@ def _profiling_options(cfg: Config): ) +def shape_buckets(cfg: Config): + """Effective ``static_shape_buckets``: the configured lengths plus ``data.max_length`` as the last one.""" + t = cfg.training + if not t.static_shapes or not t.static_shape_buckets: + return None + return tuple(sorted(set(int(b) for b in t.static_shape_buckets) | {int(cfg.data.max_length)})) + + def _backend_options(cfg: Config) -> BackendOptions: """``training.*`` -> typed backend options, shared by every launch path.""" - return BackendOptions(compile_blocks=cfg.training.compile_blocks) + buckets = shape_buckets(cfg) + return BackendOptions( + compile_blocks=cfg.training.compile_blocks, + # Bucket padding alone is a data-path setting; only compiled blocks need to know the count. + compile_shape_buckets=len(buckets) if (buckets and cfg.training.compile_blocks) else 0, + ) def _common_launch_kwargs( @@ -563,6 +576,7 @@ def _common_launch_kwargs( fsdp_sharding=t.fsdp_sharding, backend_options=_backend_options(cfg), static_shapes=t.static_shapes, + static_shape_buckets=shape_buckets(cfg), 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 b31bab9f0..531f4da4d 100644 --- a/specforge/training/backend.py +++ b/specforge/training/backend.py @@ -157,6 +157,10 @@ class BackendOptions: #: ``torch.compile`` every draft block (or the EAGLE midlayer) in place before FSDP2 sharding. compile_blocks: bool = False + #: Number of static input lengths the compiled blocks will see (``training.static_shape_buckets``): + #: above 1, every block is compiled with ``dynamic=False`` and Dynamo's recompile limit is raised so + #: each bucket gets its own static graph instead of a symbolic-shape one. + compile_shape_buckets: int = 0 class TrainingBackend(abc.ABC): diff --git a/specforge/training/disaggregated.py b/specforge/training/disaggregated.py index 1676227fe..8ebc648d9 100644 --- a/specforge/training/disaggregated.py +++ b/specforge/training/disaggregated.py @@ -441,6 +441,7 @@ def _build_offline( _backend_options, _dataloader_num_workers, _profiling_options, + shape_buckets, ) manifest = _env("DISAGG_MANIFEST") @@ -543,6 +544,7 @@ def produce() -> int: fsdp_sharding=cfg.training.fsdp_sharding, backend_options=_backend_options(cfg), static_shapes=cfg.training.static_shapes, + static_shape_buckets=shape_buckets(cfg), run_id=cfg.run_id, output_dir=cfg.output_dir, ttt_length=cfg.training.ttt_length, @@ -611,6 +613,7 @@ def _build_online( _dataloader_num_workers, _load_input_tools, _profiling_options, + shape_buckets, ) modality = cfg.model.input_modality @@ -827,6 +830,7 @@ def produce() -> int: fsdp_sharding=cfg.training.fsdp_sharding, backend_options=_backend_options(cfg), static_shapes=cfg.training.static_shapes, + static_shape_buckets=shape_buckets(cfg), max_len=cfg.data.max_length, run_id=cfg.run_id, output_dir=cfg.output_dir, diff --git a/specforge/training/fsdp2.py b/specforge/training/fsdp2.py index b5a67a301..eefd97212 100644 --- a/specforge/training/fsdp2.py +++ b/specforge/training/fsdp2.py @@ -7,6 +7,26 @@ from specforge.training.backend import DistributedTrainingBackend +def _configure_bucketed_compile(buckets: int) -> None: + """Let Dynamo hold one static graph per length bucket. + + Each block is traced twice per bucket (the first block's input does not + require grad, the others' does), so the per-frame recompile limit must be at + least ``2 * buckets``; past the limit Dynamo would silently fall back to + eager for new shapes. Both the current and the pre-2.7 config names are set. + """ + import torch._dynamo + + cfg = torch._dynamo.config + needed = 2 * int(buckets) + 4 + for name in ("recompile_limit", "cache_size_limit"): + if hasattr(cfg, name) and int(getattr(cfg, name)) < needed: + setattr(cfg, name, needed) + for name in ("accumulated_recompile_limit", "accumulated_cache_size_limit"): + if hasattr(cfg, name) and int(getattr(cfg, name)) < 32 * needed: + setattr(cfg, name, 32 * needed) + + class FSDP2TrainingBackend(DistributedTrainingBackend): name = "fsdp2" compiled_blocks: int = 0 @@ -26,9 +46,16 @@ def _prepare_blocks(self, model, block_classes, optimizer_target) -> None: # FSDP2 hooks registered afterwards run inside the compiled call, but # Dynamo skips them (``torch._dynamo.config.skip_fsdp_hooks``), so they # execute eagerly around the compiled block body. Dynamo starts static - # and marks shapes dynamic only after a recompilation. + # and marks shapes dynamic only after a recompilation; with length + # buckets every bucket must stay a static graph instead. + buckets = int(getattr(self.options, "compile_shape_buckets", 0) or 0) + if buckets > 1: + _configure_bucketed_compile(buckets) for module in targets: - module.compile() + if buckets > 1: + module.compile(dynamic=False) + else: + module.compile() self.compiled_blocks = len(targets) def _shard_model(self, model, block_classes, ignored_frozen_modules): diff --git a/tests/test_runtime/test_shape_buckets.py b/tests/test_runtime/test_shape_buckets.py new file mode 100644 index 000000000..36e200532 --- /dev/null +++ b/tests/test_runtime/test_shape_buckets.py @@ -0,0 +1,182 @@ +"""``training.static_shape_buckets``: a few static lengths instead of one, one compiled graph per bucket.""" + +import os +import types +import unittest +from unittest import mock + +import torch + +from specforge.algorithms.builtin import builtin_algorithm_registry + +ALGORITHM = builtin_algorithm_registry().resolve("dflash") + + +def _sample(length, width=4): + return { + "input_ids": torch.arange(length).unsqueeze(0), + "loss_mask": torch.ones(1, length, dtype=torch.long), + "hidden_states": torch.randn(1, length, width), + } + + +class TestBucketResolution(unittest.TestCase): + def test_resolve_static_length(self): + from specforge.algorithms.common.collation import resolve_static_length + + self.assertEqual(resolve_static_length(300, None), 300) + self.assertEqual(resolve_static_length(300, 2048), 2048) + with self.assertRaisesRegex(ValueError, "static batch length"): + resolve_static_length(3000, 2048) + buckets = (512, 1024, 2048) + self.assertEqual(resolve_static_length(300, buckets), 512) + self.assertEqual(resolve_static_length(512, buckets), 512) + self.assertEqual(resolve_static_length(513, buckets), 1024) + self.assertEqual(resolve_static_length(2048, buckets), 2048) + with self.assertRaisesRegex(ValueError, "static batch length"): + resolve_static_length(2049, buckets) + # Unsorted input is tolerated. + self.assertEqual(resolve_static_length(700, [2048, 512, 1024]), 1024) + + def test_collators_pad_to_the_smallest_fitting_bucket(self): + from specforge.algorithms.common.hidden_states_data import build_collator + from specforge.data.utils import DataCollatorWithPadding + + collate = build_collator(pad_to=(8, 16)) + self.assertEqual(tuple(collate([_sample(5), _sample(3)])["input_ids"].shape), (2, 8)) + self.assertEqual(tuple(collate([_sample(9), _sample(3)])["hidden_states"].shape), (2, 16, 4)) + + def item(n): + ones = torch.ones(1, n, dtype=torch.long) + return {"input_ids": ones, "attention_mask": ones.clone(), "loss_mask": ones.clone()} + + eagle = DataCollatorWithPadding(pad_to=(8, 16)) + self.assertEqual(tuple(eagle([item(5), item(3)])["input_ids"].shape), (2, 8)) + self.assertEqual(tuple(eagle([item(9)])["input_ids"].shape), (1, 16)) + + def test_launch_resolves_buckets(self): + from specforge.launch import _offline_io, _static_pad_length, _streaming_collate + + self.assertEqual(_static_pad_length(True, 2048, [512, 1024]), (512, 1024, 2048)) + self.assertEqual(_static_pad_length(True, 2048, [512, 2048]), (512, 2048)) + self.assertEqual(_static_pad_length(True, 2048, None), 2048) + self.assertIsNone(_static_pad_length(False, 2048, [512])) + with self.assertRaisesRegex(ValueError, "exceed"): + _static_pad_length(True, 2048, [4096]) + collate = _streaming_collate(ALGORITHM, "text", None, pad_to=(8, 16)) + self.assertEqual(tuple(collate([_sample(9), _sample(2)])["input_ids"].shape), (2, 16)) + collate, _ = _offline_io( + ALGORITHM, "text", 16, ttt_length=7, use_usp_preprocess=False, + static_shapes=True, static_shape_buckets=[8], + ) + self.assertEqual(tuple(collate([_sample(5), _sample(3)])["input_ids"].shape), (2, 8)) + + +class TestBucketConfig(unittest.TestCase): + def test_training_config_validation(self): + from specforge.config.schema import TrainingConfig + + with self.assertRaisesRegex(ValueError, "requires training.static_shapes"): + TrainingConfig(static_shape_buckets=[512]) + with self.assertRaisesRegex(ValueError, "multiples of 128"): + TrainingConfig(static_shapes=True, static_shape_buckets=[500]) + with self.assertRaisesRegex(ValueError, "strictly increasing"): + TrainingConfig(static_shapes=True, static_shape_buckets=[1024, 512]) + with self.assertRaisesRegex(ValueError, "must not be empty"): + TrainingConfig(static_shapes=True, static_shape_buckets=[]) + cfg = TrainingConfig(static_shapes=True, static_shape_buckets=[512, 1024]) + self.assertEqual(cfg.static_shape_buckets, [512, 1024]) + + @staticmethod + def _config(buckets, max_length=2048, compile_blocks=True): + from specforge.config import Config + + training = { + "strategy": "dflash", "backend": "fsdp2", "static_shapes": True, + "static_shape_buckets": buckets, "compile_blocks": compile_blocks, "max_steps": 1, + } + return Config.model_validate( + { + "model": {"target_model_path": "t", "draft_model_config": "d"}, + "data": {"hidden_states_path": "features", "max_length": max_length}, + "training": training, + } + ) + + def test_buckets_must_fit_max_length_and_get_it_appended(self): + from specforge.training.assembly import _backend_options, shape_buckets + + with self.assertRaisesRegex(ValueError, "must not exceed data.max_length"): + self._config([512, 4096]) + cfg = self._config([512, 1024]) + self.assertEqual(shape_buckets(cfg), (512, 1024, 2048)) + self.assertEqual(_backend_options(cfg).compile_shape_buckets, 3) + self.assertTrue(_backend_options(cfg).compile_blocks) + # Padding buckets without compile_blocks leave the backend options untouched + # (FSDP1 rejects any set option, and buckets alone are a data-path setting). + loose = self._config([512, 1024], compile_blocks=False) + self.assertEqual(shape_buckets(loose), (512, 1024, 2048)) + self.assertEqual(_backend_options(loose).compile_shape_buckets, 0) + + def test_bucketed_compile_raises_the_recompile_limit(self): + import torch._dynamo + + from specforge.training.fsdp2 import _configure_bucketed_compile + + cfg = torch._dynamo.config + names = [n for n in ("recompile_limit", "cache_size_limit") if hasattr(cfg, n)] + self.assertTrue(names) + saved = {n: getattr(cfg, n) for n in names} + try: + for n in names: + setattr(cfg, n, 8) + _configure_bucketed_compile(4) + for n in names: + self.assertGreaterEqual(getattr(cfg, n), 12) + finally: + for n, v in saved.items(): + setattr(cfg, n, v) + + +class TestDisaggregatedLaunchForwardsBuckets(unittest.TestCase): + def test_online_consumer_receives_the_buckets(self): + from specforge.config import Config + from specforge.training.disaggregated import _build_online + + cfg = Config.model_validate( + { + "model": {"target_model_path": "t", "draft_model_config": "d"}, + "data": {"prompts_path": "prompts.jsonl", "max_length": 2048}, + "training": { + "strategy": "dflash", "role": "consumer", "max_steps": 1, "backend": "fsdp2", + "static_shapes": True, "static_shape_buckets": [512, 1024], "compile_blocks": True, + }, + "deployment": { + "mode": "disaggregated", + "disaggregated": {"control_dir": "/shared/buckets", "backend": "mooncake", "server_urls": ["http://capture:30000"]}, + }, + } + ) + bundle = types.SimpleNamespace(model=object(), target_head=None, strategy_kwargs={}) + + class _Trainer: + def fit(self): + return 1 + + with ( + mock.patch.dict(os.environ, {"DISAGG_REF_CHANNEL": "/shared/refs"}), + mock.patch("specforge.runtime.data_plane.streaming_ref_channel.StreamingRefChannel"), + mock.patch("specforge.training.disaggregated._mooncake_store", return_value=mock.Mock()), + mock.patch("specforge.launch.build_disagg_online_consumer", return_value=_Trainer()) as build, + ): + _build_online( + cfg, algorithm=ALGORITHM, build_model_bundle=lambda _cfg: bundle, + prepare_prompts=mock.Mock(), optimizer_factory=mock.Mock(), logger=None, + ) + kw = build.call_args.kwargs + self.assertEqual(kw["static_shape_buckets"], (512, 1024, 2048)) + self.assertEqual(kw["backend_options"].compile_shape_buckets, 3) + + +if __name__ == "__main__": + unittest.main()