Skip to content
Draft
32 changes: 32 additions & 0 deletions docs/sections/basic_usage/training.md
Original file line number Diff line number Diff line change
Expand Up @@ -487,6 +487,30 @@ 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
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.

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
Expand Down Expand Up @@ -658,6 +682,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.

Expand Down
1 change: 1 addition & 0 deletions examples/configs/README.md
Original file line number Diff line number Diff line change
Expand Up @@ -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. |
Expand Down
13 changes: 11 additions & 2 deletions specforge/algorithms/common/collation.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,7 +2,7 @@

from __future__ import annotations

from typing import Mapping, Sequence
from typing import Mapping, Optional, Sequence


def concatenate_features(features):
Expand All @@ -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:
Expand All @@ -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

Expand Down
27 changes: 23 additions & 4 deletions specforge/algorithms/common/dflash_family_model.py
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -461,20 +466,30 @@ 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)
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,
Expand Down Expand Up @@ -1830,6 +1845,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,
Expand All @@ -1842,6 +1858,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",
Expand Down Expand Up @@ -2117,6 +2134,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,
Expand All @@ -2131,6 +2149,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",
Expand Down
17 changes: 11 additions & 6 deletions specforge/algorithms/common/hidden_states_data.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)}
Expand All @@ -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,
)


Expand Down Expand Up @@ -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__ = [
Expand Down
8 changes: 5 additions & 3 deletions specforge/algorithms/eagle3/data.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
1 change: 1 addition & 0 deletions specforge/algorithms/model_providers.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
}
Expand Down
35 changes: 35 additions & 0 deletions specforge/config/schema.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -894,6 +895,13 @@ 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
#: 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
#: Run the trainable linears inside the draft blocks as torchao ``Float8Linear`` with float8 FSDP2 all-gather. Requires ``backend: fsdp2``.
fp8_linear: 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)
Expand Down Expand Up @@ -1000,6 +1008,28 @@ 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")
if self.compile_blocks and not self.static_shapes:
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,
)
if self.fp8_linear and self.backend != "fsdp2":
raise ValueError("training.fp8_linear requires training.backend=fsdp2")
if self.fp8_linear and not self.static_shapes:
warnings.warn(
"training.fp8_linear without training.static_shapes needs every "
"micro-batch's token count to be a multiple of 16 (float8 GEMMs) and a "
"fixed shape for compile_blocks; with pad-to-longest batches it fails at "
"the first step. 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":
if self.batch_size != 1:
Expand Down Expand Up @@ -1113,6 +1143,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.fp8_linear and self.training.static_shapes and self.data.max_length % 16:
raise ValueError(
"training.fp8_linear with training.static_shapes needs data.max_length "
"to be a multiple of 16 (float8 GEMMs over the padded context positions)"
)
deployment = self.deployment.mode
role = self.training.role

Expand Down
12 changes: 11 additions & 1 deletion specforge/data/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down Expand Up @@ -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 = (
Expand Down
Loading