Skip to content

(fsdp2 dev. perf) Add training.fp8_linear: torchao Float8Linear with float8 all-gather - #923

Draft
yushengsu-thu wants to merge 9 commits into
sgl-project:mainfrom
yushengsu-thu:fsdp2/fp8-linear
Draft

yushengsu-thu wants to merge 9 commits into
sgl-project:mainfrom
yushengsu-thu:fsdp2/fp8-linear

Conversation

@yushengsu-thu

@yushengsu-thu yushengsu-thu commented Oct 4, 2026 •

Copy link
Copy Markdown

Stacked on #922 → #915 (fsdp2/fp8-linear on fsdp2/compile-blocks on codex/fsdp2-backend): the diff below also shows #922's and #915's commits until they merge. Increment-only view for review: yushengsu-thu#2

Conclusion

  • Size over (fsdp2 dev. perf) Add training.compile_blocks: per-block torch.compile before fully_shard #922 (excluding tests and docs): +125 / −13 lines of Python (fsdp2.py +101/−12, schema.py +18, assembly.py +4/−1, backend.py +2); tests +303, docs +1/−1. Over [verifying] Add optional FSDP2 training backend #915 the stack is +336 / −20.
  • FSDP1 → FSDP2 alone ([verifying] Add optional FSDP2 training backend #915, no option): speed — EAGLE3 faster, −132 ms/step (796 → 664 ms offline e2e); DFlash2 / DSpark no change. Memory — EAGLE3 smaller, −1782 MB peak; DFlash2 / DSpark no change.
  • The online rows are the ones to read (real conversations, variable lengths); the offline rows use synthetic samples and bracket the real case.
  • Online e2e (real ShareGPT conversations, variable lengths), FSDP2 + static_shapes + compile_blocks + fp8_linear vs fsdp2: DFlash2 — speed faster, −318 ms per step (1081 → 763 ms compute per step, −29.4%); memory smaller, large: −7440 MB (53982 → 46542 MB nvidia-smi max). DSpark — speed faster, −132 ms per step (719 → 587 ms compute per step, −18.3%); memory smaller, small: −3944 MB (39652 → 35708 MB nvidia-smi max).
  • Online e2e, the fp8 part alone (vs static_shapes + compile_blocks, i.e. vs (fsdp2 dev. perf) Add training.compile_blocks: per-block torch.compile before fully_shard #922): DFlash2 — speed no change (765 → 763 ms compute per step, −3 ms, −0.3%); memory smaller, large: −5890 MB (52432 → 46542 MB nvidia-smi max). DSpark — speed slower, +18 ms per step (569 → 587 ms compute per step, +3.2%); memory smaller, small: −3274 MB (38982 → 35708 MB nvidia-smi max).
  • Offline e2e, synthetic samples with variable loss masks, FSDP2 + static_shapes + compile_blocks + fp8_linear vs fsdp2: DFlash2 — speed faster, −955 ms/step (1750 → 795 ms, −54.6%); memory smaller, large: −8868 MB (37652 → 28783 MB). DSpark — speed faster, −176 ms/step (823 → 648 ms, −21.3%); memory smaller, large: −4359 MB (25156 → 20797 MB).
  • Offline e2e, variable masks, the fp8 part alone (vs (fsdp2 dev. perf) Add training.compile_blocks: per-block torch.compile before fully_shard #922): DFlash2 — speed faster, −43 ms/step (838 → 795 ms, −5.1%); memory smaller, large: −5319 MB (34103 → 28783 MB). DSpark — speed no change (641 → 648 ms, +6 ms); memory smaller, small: −1532 MB (22329 → 20797 MB).
  • Offline e2e, synthetic fixed-length inputs (the first measurement): speed vs fsdp2 — faster: DFlash2 −225 ms/step (995 → 770 ms, −22.6%), DSpark −88 ms (726 → 638), EAGLE3 −72 ms (664 → 592); vs compile alone — DFlash2 faster −43 ms, DSpark no change ±0 ms, EAGLE3 slower +23 ms. Memory — smaller, large: peak −8869 MB / −4262 MB / −1256 MB per rank.
  • FSDP2 + fp8_linear without compile: speed — slower, large: +135 ms / +64 ms / +104 ms per micro-step (eager dynamic scaling). Memory — small gain only (−961 / −921 / −415 MB). The config does not force the pairing, the recipe should enable both.
  • Constraints: float8 GEMMs need the token count of every micro-batch to be a multiple of 16, which pad-to-longest batches are not, so fp8_linear warns without static_shapes and, with it, enforces data.max_length % 16 == 0; when the data-parallel size does not divide a converted weight's rows (e.g. DP=3 for 4096 rows) the weights are all-gathered in bf16 and only the GEMMs run in float8 (logged once).
  • Numerics: DSpark loss +0.4%, grad norm −1.3% after two steps; DFlash2 / EAGLE3 unchanged to 5 digits. Acceptance length on a real run is the gate before a recipe adopts it.

Motivation

FP8 GEMMs for the draft blocks, with the weights all-gathered in float8 through FSDP2's extension point (fsdp_pre_all_gather), which only exists for the composable fully_shard from #915. Opt-in via training.fp8_linear.

Stacks on #922 (fsdp2/compile-blocks), which stacks on #915: fp8 is only worthwhile together with compile_blocks (eager float8 scaling is slower than bf16) and on real data it needs #922's static_shapes, see the constraints above.

Modifications

  • training/fsdp2.py: with fp8_linear, swap the trainable nn.Linear layers inside the draft blocks for torchao Float8Linear (Float8LinearConfig(enable_fsdp_float8_all_gather=...), dynamic tensorwise scaling) before compile_blocks compiles the block and before sharding (torchtitan's order), and call precompute_float8_dynamic_scale_for_fsdp after every optimizer step while the float8 all-gather is on. Linears with a dimension not divisible by 16, frozen linears and the composite's lm_head stay bf16. Uneven-shard guard: if the DP size does not divide a converted weight's rows the float8 all-gather is turned off (bf16 all-gather, float8 GEMMs) with a warning, because torchao builds the unsharded float8 view for the unpadded size.
  • config/schema.py: training.fp8_linear (requires backend: fsdp2; warns without static_shapes; with it data.max_length % 16 == 0 is enforced); plumbing through assembly._backend_options(), launch, trainer, both disaggregated call sites.
  • tests/test_runtime/test_fsdp2_fp8_linear.py: rejection/config tests (backend, static shapes, max_length), filter, DP-size helper, uneven-shard fallback, 2-GPU training smoke (requires torchao and sm_89+), disaggregated forwarding.

torchao is an optional dependency (pip install torchao; 0.18.0 was used).

End-to-end comparison — online (disaggregated)

Online e2e = the disaggregated path exactly as specforge train runs it (managed-local): one SGLang 0.5.18 capture server with the spec-capture patch on GPU 0 (Qwen3-8B, mem_fraction_static 0.5), Mooncake store, producers feeding real ShareGPT prompts (6000 conversations, max_length 2048, qwen chat template), and a DP=3 trainer on GPUs 1–3; batch 2 × accum 4, 40 optimizer steps, num_anchors 512; steady state = the controller's perf/* counters (averaged over the ranks that logged them) over the last 24 optimizer steps. samples/s = 24 samples per optimizer step (batch 2 × accum 4 × 3 trainer GPUs) divided by the mean step time of that window; train compute and data wait are per optimizer step. GPU memory is the nvidia-smi maximum over the run for the trainer GPUs.

DFlash2

variant steady samples/s (3 trainer GPUs) train compute per step Δ compute vs fsdp2 data wait per step trainer GPU memory max Δ memory vs fsdp2
fsdp2 22.15 1081 ms – 1 ms 53982 MB –
fsdp2 + static_shapes 24.84 963 ms −117 ms (−10.9%) 1 ms 57392 MB +3410 MB
fsdp2 + static_shapes + compile_blocks (#922) 31.24 765 ms −315 ms (−29.2%) 1 ms 52432 MB −1550 MB
fsdp2 + static_shapes + compile_blocks + fp8_linear 31.35 763 ms −318 ms (−29.4%) 1 ms 46542 MB −7440 MB

DSpark

variant steady samples/s (3 trainer GPUs) train compute per step Δ compute vs fsdp2 data wait per step trainer GPU memory max Δ memory vs fsdp2
fsdp2 33.30 719 ms – 1 ms 39652 MB –
fsdp2 + static_shapes 35.88 667 ms −52 ms (−7.2%) 1 ms 40950 MB +1298 MB
fsdp2 + static_shapes + compile_blocks (#922) 41.99 569 ms −150 ms (−20.9%) 1 ms 38982 MB −670 MB
fsdp2 + static_shapes + compile_blocks + fp8_linear 40.75 587 ms −132 ms (−18.3%) 1 ms 35708 MB −3944 MB

The first online attempt of this PR (first devbox, pad-to-longest, DP=3) failed before the first optimizer step with RuntimeError: Expected self.size(1) to be divisible by 16, but got self.size(1)=3858 (a block's K/V projection over a token count that is not a multiple of 16) and RuntimeError: setStorage: sizes [4096, 4096], strides [4096, 1], storage offset 0, and itemsize 1 requiring a storage size of 16777216 are out of bounds for storage of size 16760832 (float8 all-gather on padded uneven shards), plus the compile_blocks crash fixed in #922; static_shapes, the max_length check and the uneven-shard fallback above are the fixes.

Measured on a second 4 × B300 devbox (same type as the first) after the fixes described above; the fsdp2 baseline was re-run there. Online memory is nvidia-smi used on the trainer GPUs, i.e. what the caching allocator reserved, a coarse measure; the precise peak-allocated numbers are in the offline tables.

End-to-end comparison — offline with variable loss masks (Trainer.fit)

Offline e2e with variable loss masks = the same Trainer.fit() run as above, but every synthetic sample has a random 20–80% prefix masked out of the loss, so the number of valid anchors differs between micro-batches as it does with real conversations (the harness's --variable-mask); 256 samples × 6 epochs, seq 2048, batch 2 × accum 4, DP=4 on 4 × B300, steady state = last 24 steps.

DFlash2

variant steady step time Δ vs fsdp2 samples/s (4 GPUs) peak allocated per rank Δ memory
fsdp2 1750 ms – 18.3 37652 MB –
fsdp2 + static_shapes + compile_blocks (#922) 838 ms -912 ms (−52.1%) 38.2 34103 MB -3549 MB
fsdp2 + static_shapes + compile_blocks + fp8_linear 795 ms -955 ms (−54.6%) 40.2 28783 MB -8868 MB

DSpark

variant steady step time Δ vs fsdp2 samples/s (4 GPUs) peak allocated per rank Δ memory
fsdp2 823 ms – 38.9 25156 MB –
fsdp2 + static_shapes + compile_blocks (#922) 641 ms -182 ms (−22.1%) 49.9 22329 MB -2828 MB
fsdp2 + static_shapes + compile_blocks + fp8_linear 648 ms -176 ms (−21.3%) 49.4 20797 MB -4359 MB

End-to-end comparison — offline with fixed-length inputs (Trainer.fit)

Offline e2e = build_offline_runtime → Trainer.fit(): offline reader, FeatureDataLoader with 4 workers, TrainerController acks and logging, final checkpoint; 256 synthetic samples × 6 epochs = 48 optimizer steps, seq 2048, batch 2 × accum 4, DP=4 on 4 × B300 (torch 2.13.0+cu130); steady state = last 24 steps from the controller's log callbacks. Step-level = the real TrainerCore step on a resident batch, batch 2 × accum 8, seq 4096. Production-shaped Qwen3-8B drafts with random weights and synthetic features; deltas are per rank.

draft fsdp → fsdp2 step time fsdp2 → +compile → +compile+fp8 Δ vs fsdp2 Δ vs compile alone peak allocated: fsdp2 → +compile+fp8 Δ memory vs fsdp2 fit wall (compile+fp8)
DFlash2 978 → 995 ms (no change) 995 → 813 → 770 ms −225 ms (−22.6%) −43 ms 37652 → 28783 MB −8869 MB 80.6 s
DSpark 740 → 726 ms (no change) 726 → 638 → 638 ms −88 ms (−12.2%) ±0 ms 25058 → 20796 MB −4262 MB 66.4 s
EAGLE3 796 → 664 ms (−132 ms) 664 → 569 → 592 ms −72 ms (−11.0%) +23 ms 21168 → 19912 MB −1256 MB 144.1 s

Notes: the "+fp8 without compile" numbers in the conclusion come from the step-level harness (eager float8, one micro-step), not from an e2e run; everything else above is e2e. The memory reduction comes from float8 activations saved for backward and the float8 weight all-gather buffers.

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 sgl-project#915; the BackendOptions scaffolding is shared verbatim
with the other FSDP2 option PRs so they can land in any order.
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.
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.
@yushengsu-thu yushengsu-thu changed the title (fsdp2 dev.) Add training.fp8_linear: torchao Float8Linear with float8 all-gather (fsdp2 dev. perf) Add training.fp8_linear: torchao Float8Linear with float8 all-gather Oct 4, 2026
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.
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.
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.
With fp8_linear the FSDP2 backend swaps the trainable nn.Linear layers inside
the draft blocks for torchao Float8Linear before sharding (and before
compile_blocks compiles the block, so the graph holds the float8 GEMMs) and
precomputes the next step's float8 scales after every optimizer step.

Constraints learned from the online runs: float8 GEMMs need the token count of
every micro-batch to be a multiple of 16, so the option requires
training.static_shapes (data.max_length padded batches, fixed anchor count);
and torchao's float8 all-gather breaks on FSDP2's padded uneven shards, so when
the data-parallel size does not divide a converted weight's rows the weights
are all-gathered in bf16 and only the GEMMs run in float8 (logged once).

The single-process path and both disaggregated call sites forward the option;
tests cover rejection, config, the filter, the shard guard and a 2-GPU smoke.

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant