(fsdp2 dev. perf) Add training.fp8_linear: torchao Float8Linear with float8 all-gather - #923
Draft
yushengsu-thu wants to merge 9 commits into
Draft
yushengsu-thu wants to merge 9 commits into
yushengsu-thu wants to merge 9 commits into
Conversation
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.
yushengsu-thu
force-pushed
the
fsdp2/fp8-linear
branch
from
October 4, 2026 06:09
3591740 to
0b6b088
Compare
This was referenced Oct 4, 2026
Draft
Closed
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.
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.
yushengsu-thu
force-pushed
the
fsdp2/fp8-linear
branch
from
October 4, 2026 09:09
3844409 to
3109ed9
Compare
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.
yushengsu-thu
force-pushed
the
fsdp2/fp8-linear
branch
from
October 4, 2026 10:30
3109ed9 to
39b4cf4
Compare
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.
yushengsu-thu
force-pushed
the
fsdp2/fp8-linear
branch
from
October 4, 2026 12:47
39b4cf4 to
a063ccf
Compare
This branch has not been deployed
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
Conclusion
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.static_shapes+compile_blocks+fp8_linearvsfsdp2: 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).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).static_shapes+compile_blocks+fp8_linearvsfsdp2: 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).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.fp8_linearwithout 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.fp8_linearwarns withoutstatic_shapesand, with it, enforcesdata.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).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 composablefully_shardfrom #915. Opt-in viatraining.fp8_linear.Stacks on #922 (
fsdp2/compile-blocks), which stacks on #915: fp8 is only worthwhile together withcompile_blocks(eager float8 scaling is slower than bf16) and on real data it needs #922'sstatic_shapes, see the constraints above.Modifications
training/fsdp2.py: withfp8_linear, swap the trainablenn.Linearlayers inside the draft blocks for torchaoFloat8Linear(Float8LinearConfig(enable_fsdp_float8_all_gather=...), dynamic tensorwise scaling) beforecompile_blockscompiles the block and before sharding (torchtitan's order), and callprecompute_float8_dynamic_scale_for_fsdpafter every optimizer step while the float8 all-gather is on. Linears with a dimension not divisible by 16, frozen linears and the composite'slm_headstay 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(requiresbackend: fsdp2; warns withoutstatic_shapes; with itdata.max_length % 16 == 0is enforced); plumbing throughassembly._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 trainruns 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'sperf/*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 computeanddata waitare per optimizer step. GPU memory is the nvidia-smi maximum over the run for the trainer GPUs.DFlash2
fsdp2fsdp2fsdp2fsdp2+static_shapesfsdp2+static_shapes+compile_blocks(#922)fsdp2+static_shapes+compile_blocks+fp8_linearDSpark
fsdp2fsdp2fsdp2fsdp2+static_shapesfsdp2+static_shapes+compile_blocks(#922)fsdp2+static_shapes+compile_blocks+fp8_linearThe 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) andRuntimeError: 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 thecompile_blockscrash fixed in #922;static_shapes, themax_lengthcheck 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
fsdp2baseline was re-run there. Online memory is nvidia-smiusedon 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
fsdp2fsdp2fsdp2+static_shapes+compile_blocks(#922)fsdp2+static_shapes+compile_blocks+fp8_linearDSpark
fsdp2fsdp2fsdp2+static_shapes+compile_blocks(#922)fsdp2+static_shapes+compile_blocks+fp8_linearEnd-to-end comparison — offline with fixed-length inputs (
Trainer.fit)Offline e2e =
build_offline_runtime→Trainer.fit(): offline reader,FeatureDataLoaderwith 4 workers,TrainerControlleracks 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 realTrainerCorestep 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.fsdp→fsdp2step timefsdp2→+compile→+compile+fp8fsdp2fsdp2→+compile+fp8fsdp2Notes: 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.