(fsdp2 dev.) - follow-up : Add training.static_shape_buckets: a few static lengths for compiled blocks - #929
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.
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.
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
requested review from
FlamingoPg,
FrankLeeeee,
shuaills,
sleepcoo and
zyksir
as code owners
October 4, 2026 10:49
yushengsu-thu
marked this pull request as draft
October 4, 2026 11:07
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.
…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.
yushengsu-thu
force-pushed
the
fsdp2/shape-buckets
branch
from
October 4, 2026 12:47
45aa746 to
5cb8627
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
launch.py,collation.py,assembly.py,fsdp2.py,schema.py,backend.py,data/utils.py,disaggregated.py); tests +182, docs +1/−1.fsdp2/compile-blocks), which stacks on [verifying] Add optional FSDP2 training backend #915. Why (fsdp2 dev. perf) Add training.compile_blocks: per-block torch.compile before fully_shard #922 and not elsewhere:static_shapesfrom one length to a few; the collatorpad_to, the fixed anchor count and the launch plumbing it relies on all live in (fsdp2 dev. perf) Add training.compile_blocks: per-block torch.compile before fully_shard #922;_prepare_blocks(one static graph per bucket:compile(dynamic=False)and a higher Dynamo recompile limit);max_length8192 (real ShareGPT, buckets 1024/2048/4096/8192),static_shape_buckets+compile_blocksvsstatic_shapes+compile_blocks: DFlash2 — speed faster, −41 ms per step (811 → 770 ms compute per step, −5.1%); memory smaller, large: −5876 MB (61156 → 55280 MB nvidia-smi max). DSpark — speed faster, −31 ms per step (608 → 577 ms compute per step, −5.1%); memory smaller, large: −5048 MB (46264 → 41216 MB nvidia-smi max).static_shapes+compile_blocksvs plainfsdp2(pad-to-longest): DFlash2 — speed faster, −265 ms per step (1076 → 811 ms compute per step, −24.6%); memory larger, large: +7952 MB (53204 → 61156 MB nvidia-smi max). DSpark — speed faster, −113 ms per step (721 → 608 ms compute per step, −15.7%); memory larger, large: +7128 MB (39136 → 46264 MB nvidia-smi max).static_shape_buckets+compile_blocksvs plainfsdp2: DFlash2 — speed faster, −306 ms per step (1076 → 770 ms compute per step, −28.4%); memory no change (53204 → 55280 MB nvidia-smi max, +2076 MB, within the granularity of reserved memory). DSpark — speed faster, −144 ms per step (721 → 577 ms compute per step, −20.0%); memory no change (39136 → 41216 MB nvidia-smi max, +2080 MB, within the granularity of reserved memory).fsdp2: DFlash2 — speed slower, +324 ms/step (1316 → 1640 ms, +24.7%); memory smaller, large: −2361 MB (38094 → 35733 MB). DSpark — speed slower, +1291 ms/step (562 → 1854 ms, +229.6%); memory smaller, small: −1645 MB (25458 → 23813 MB).fsdp2: DFlash2 — speed faster, −254 ms/step (1316 → 1062 ms, −19.3%); memory smaller, large: −2132 MB (38094 → 35962 MB). DSpark — speed slower, +342 ms/step (562 → 905 ms, +60.9%); memory smaller, small: −1769 MB (25458 → 23689 MB).num_anchorsanchor blocks while the static paths always run 512); the online 8192 rows above, on real conversations, put the static stack far ahead of pad-to-longest. The static and bucket rows come from 60-step runs so that the one-time compile of a rare bucket (two samples both ≤ 2048 happen in ≈2% of micro-batches) falls outside the measured window; even so the DFlash2 bucket row varies between runs with where that compile lands (1972 ms in one run, 1062 ms in the run shown, which logged 7 recompiles on rank 0 = 3 buckets × 2 grad variants + 1, no recompile-limit hit), so treat the synthetic bucket deltas as indicative and the online rows as the measurement.max_length2048, buckets 512/1024/1536/2048),static_shape_buckets+compile_blocksvsstatic_shapes+compile_blocks((fsdp2 dev. perf) Add training.compile_blocks: per-block torch.compile before fully_shard #922): DFlash2 — speed no change (768 → 765 ms compute per step, −2 ms, −0.3%); memory no change (51164 → 53374 MB nvidia-smi max, +2210 MB, within the granularity of reserved memory). DSpark — speed no change (569 → 569 ms compute per step, −0 ms, −0.0%); memory no change (38270 → 39194 MB nvidia-smi max, +924 MB, within the granularity of reserved memory).static_shape_bucketsvsstatic_shapes, no compile): DFlash2 — speed no change (959 → 974 ms compute per step, +15 ms, +1.6%); memory no change (58034 → 55590 MB nvidia-smi max, −2444 MB, within the granularity of reserved memory). DSpark — speed slower, +22 ms per step (668 → 690 ms compute per step, +3.3%); memory no change (41388 → 40454 MB nvidia-smi max, −934 MB, within the granularity of reserved memory).static_shape_buckets+compile_blocksvsfsdp, FSDP1; FSDP1 and FSDP2 are within 1% of each other without options): DFlash2 — speed faster, −318 ms per step (1083 → 765 ms compute per step, −29.3%); memory no change (54812 → 53374 MB nvidia-smi max, −1438 MB, within the granularity of reserved memory). DSpark — speed faster, −152 ms per step (721 → 569 ms compute per step, −21.1%); memory no change (41340 → 39194 MB nvidia-smi max, −2146 MB, within the granularity of reserved memory).max_length8192 on ShareGPT (median 1.3k tokens) they are faster and smaller on top of (fsdp2 dev. perf) Add training.compile_blocks: per-block torch.compile before fully_shard #922, see the 8192 bullets above; atmax_length2048 they are no change (and without compile DSpark pays +3% for switching between four shapes). Default off; pick the buckets from the corpus' length distribution (the Qwen3.8-27B regeneration corpus sits between the two cases: mean 4–4.5k under 8192).static_shape_buckets+compile_blocksvsstatic_shapes+compile_blocks: DFlash2 — speed slower, +19 ms/step (749 → 768 ms, +2.5%); memory no change (34103 → 34083 MB, −19 MB). DSpark — speed faster, −22 ms/step (588 → 566 ms, −3.8%); memory no change (22329 → 22501 MB, +172 MB).static_shape_bucketsvsstatic_shapes): DFlash2 — speed faster, −33 ms/step (935 → 902 ms, −3.6%); memory no change (37651 → 37515 MB, −137 MB). DSpark — speed faster, −58 ms/step (692 → 634 ms, −8.4%); memory no change (25058 → 25020 MB, −38 MB).static_shape_buckets+compile_blocksvs plainfsdp2(pad-to-longest): DFlash2 — speed faster, −212 ms/step (980 → 768 ms, −21.7%); memory smaller, large: −3249 MB (37332 → 34083 MB). DSpark — speed slower, +256 ms/step (310 → 566 ms, +82.5%); memory smaller, large: −2430 MB (24930 → 22501 MB).static_shapesalone vs plainfsdp2(the cost buckets exist to cut): DFlash2 — speed faster, −45 ms/step (980 → 935 ms, −4.6%); memory larger, small: +319 MB (37332 → 37651 MB). DSpark — speed slower, +382 ms/step (310 → 692 ms, +123.0%); memory no change (24930 → 25058 MB, +127 MB).max_lengthand a 20–80% prefix is unsupervised, so many samples have fewer thannum_anchorsvalid anchors and the fixed anchor count (plus themax_lengthpadding) does work the pad-to-longest path never does. Length buckets recover the context part of that; an anchor-count bucket would be the natural follow-up for corpora with few supervised tokens per sample. Real ShareGPT conversations (median 1.3k tokens, mostly supervised assistant turns) are much kinder, see the online rows.max_length2048 pads 1.23× the positions of batch-2 pad-to-longest understatic_shapes; the Qwen3.8-27B regeneration corpus (mean 4–4.5k tokens, p90 atmax_length8192, per the recipe doc) pads roughly 1.5×, which is what makes a handful of buckets worth their one-time compile warm-up (≈12 s per bucket for the DFlash family).Motivation
training.static_shapes(#922) gives compiled blocks one input shape by padding every micro-batch todata.max_length. Short conversations then do more work than they need: the padding grows with the gap between the typical conversation andmax_length.training.static_shape_bucketskeeps the shapes static but lets each micro-batch pad to the smallest of a few lengths that fits; compiled blocks hold one static graph per bucket. The constraints on the lengths:flex_attention's block mask works in 128-token blocks, so aligned lengths never pay for a partial block;fp8_linearthe token count must be a multiple of 16, which multiples of 128 satisfy;data.max_length, andmax_lengthis always the last bucket (appended if missing);num_anchors(the draft-token dimension is already static).Stacks on #922 (
fsdp2/compile-blocks). TheBackendOptionsscaffolding is shared verbatim with the other FSDP2 option PRs so they can land in any order.Modifications
config/schema.py:training.static_shape_buckets: list[int](requiresstatic_shapes; positive multiples of 128, strictly increasing; a Config-level check rejects buckets abovedata.max_length).algorithms/common/collation.py:resolve_static_length(longest, pad_to)picks the smallest bucket that fits (or the single fixed length);pad_and_concatenate_featuresandDataCollatorWithPaddingaccept a tuple of lengths aspad_to.launch.py:_static_pad_length(static_shapes, max_len, buckets)returns the fixed length or the ascending bucket tuple (withmax_lenappended); every entry point and the eval loader takestatic_shape_buckets;assembly.shape_buckets(cfg)computes the effective tuple once andassembly._backend_options()passes the count to the backend only whencompile_blocksis on (bucket padding alone works on either backend).training/backend.py+training/fsdp2.py:BackendOptions.compile_shape_buckets; above 1 the blocks are compiled withdynamic=Falseand_configure_bucketed_compile()raises Dynamo's per-frame recompile limit to2 × buckets + 4(each block is traced twice per bucket: the first block's input has no grad, the others' does) plus the accumulated limit, so a new bucket never turns a dimension symbolic and never falls back to eager.training/disaggregated.py: both call sites forward the buckets.tests/test_runtime/test_shape_buckets.py: bucket resolution, both collators, launch resolution, config validation, effective buckets and backend count, Dynamo limit, disaggregated forwarding.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
fsdp2-staticfsdp2-staticfsdp2+static_shapes(#922)fsdp2+static_shape_bucketsfsdp2+static_shapes+compile_blocks(#922)fsdp2+static_shape_buckets+compile_blocksDSpark
fsdp2-staticfsdp2-staticfsdp2+static_shapes(#922)fsdp2+static_shape_bucketsfsdp2+static_shapes+compile_blocks(#922)fsdp2+static_shape_buckets+compile_blocksMeasured on a third 4 × B300 devbox so that FSDP1 and FSDP2 share one baseline; same configuration as the other online tables.
End-to-end comparison at
max_length8192 — 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
fsdp2fsdp2fsdp2(pad to longest)fsdp2+static_shapes+compile_blocks(#922)fsdp2+static_shape_buckets+compile_blocksDSpark
fsdp2fsdp2fsdp2(pad to longest)fsdp2+static_shapes+compile_blocks(#922) (38 steps logged)fsdp2+static_shape_buckets+compile_blocks(38 steps logged)Fourth devbox,
data.max_length8192 (everything else as in the other online tables: real ShareGPT prompts, DP=3 trainers, 40 steps).End-to-end comparison at seq 8192 — offline, variable-length samples (
Trainer.fit)Offline e2e at seq 8192 =
Trainer.fit()on synthetic samples whose length is uniform in 1024–8192 tokens (mean ≈4.6k, the Qwen3.8-27B regeneration corpus averages 4–4.5k undermax_length8192) with a random 20–80% prefix unsupervised; 192 samples × 5 epochs = 30 steps, batch 2 × accum 4, DP=4 on 4 × B300, steady state = last 20 steps; buckets 1024/2048/4096/8192.DFlash2
fsdp2fsdp2(pad to longest)fsdp2+static_shapes+compile_blocks(#922)fsdp2+static_shape_buckets+compile_blocksDSpark
fsdp2fsdp2(pad to longest)fsdp2+static_shapes+compile_blocks(#922)fsdp2+static_shape_buckets+compile_blocksEnd-to-end comparison — offline with variable-length samples (
Trainer.fit)Offline e2e with variable lengths =
Trainer.fit()on synthetic samples whose length is uniform in 25–100% of 2048 tokens and whose loss mask hides a random 20–80% prefix (the harness's--variable-length --variable-mask), so pad-to-longest batches change shape every micro-batch like real conversations; 256 samples × 6 epochs, batch 2 × accum 4, DP=4 on 4 × B300, steady state = last 24 steps; buckets 512/1024/1536/2048.DFlash2
fsdp2fsdp2(pad to longest)fsdp2+static_shapesfsdp2+static_shape_bucketsfsdp2+static_shapes+compile_blocks(#922)fsdp2+static_shape_buckets+compile_blocksDSpark
fsdp2fsdp2(pad to longest)fsdp2+static_shapesfsdp2+static_shape_bucketsfsdp2+static_shapes+compile_blocks(#922)fsdp2+static_shape_buckets+compile_blocks