Repository navigation
(fsdp2 dev.) Add training.static_shape_buckets: a few static lengths for compiled blocks - #6
Open
yushengsu-thu wants to merge 1 commit into
Open
yushengsu-thu wants to merge 1 commit into
yushengsu-thu wants to merge 1 commit into
Conversation
…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 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.
Increment-only view of sgl-project#929 (base
fsdp2/compile-blocks= sgl-project#922, which stacks on sgl-project#915).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 sgl-project/SpecForge#915. Why (fsdp2 dev. perf) Add training.compile_blocks: per-block torch.compile before fully_shard sgl-project/SpecForge#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 sgl-project/SpecForge#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 sgl-project/SpecForge#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 sgl-project/SpecForge#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(sgl-project#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 sgl-project#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(sgl-project#922)fsdp2+static_shape_bucketsfsdp2+static_shapes+compile_blocks(sgl-project#922)fsdp2+static_shape_buckets+compile_blocksDSpark
fsdp2-staticfsdp2-staticfsdp2+static_shapes(sgl-project#922)fsdp2+static_shape_bucketsfsdp2+static_shapes+compile_blocks(sgl-project#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(sgl-project#922)fsdp2+static_shape_buckets+compile_blocksDSpark
fsdp2fsdp2fsdp2(pad to longest)fsdp2+static_shapes+compile_blocks(sgl-project#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(sgl-project#922)fsdp2+static_shape_buckets+compile_blocksDSpark
fsdp2fsdp2(pad to longest)fsdp2+static_shapes+compile_blocks(sgl-project#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(sgl-project#922)fsdp2+static_shape_buckets+compile_blocksDSpark
fsdp2fsdp2(pad to longest)fsdp2+static_shapesfsdp2+static_shape_bucketsfsdp2+static_shapes+compile_blocks(sgl-project#922)fsdp2+static_shape_buckets+compile_blocks