Skip to content

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

Draft
yushengsu-thu wants to merge 1 commit into
fsdp2/compile-blocksfrom
fsdp2/fp8-linear
Draft

yushengsu-thu wants to merge 1 commit into
fsdp2/compile-blocksfrom
fsdp2/fp8-linear

Conversation

@yushengsu-thu

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

Copy link
Copy Markdown
Owner

Increment-only view of sgl-project#923 (base fsdp2/compile-blocks = sgl-project#922, which stacks on sgl-project#915).

Conclusion

  • Size over (fsdp2 dev. perf) Add training.compile_blocks: per-block torch.compile before fully_shard sgl-project/SpecForge#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 sgl-project/SpecForge#915 the stack is +336 / −20.
  • FSDP1 → FSDP2 alone ([verifying] Add optional FSDP2 training backend sgl-project/SpecForge#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 sgl-project/SpecForge#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 sgl-project/SpecForge#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 sgl-project#915. Opt-in via training.fp8_linear.

Stacks on sgl-project#922 (fsdp2/compile-blocks), which stacks on sgl-project#915: fp8 is only worthwhile together with compile_blocks (eager float8 scaling is slower than bf16) and on real data it needs sgl-project#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 (sgl-project#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 (sgl-project#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 sgl-project#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 (sgl-project#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 (sgl-project#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.

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.
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