Skip to content

(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
sgl-project:mainfrom
yushengsu-thu:fsdp2/shape-buckets
Draft

yushengsu-thu wants to merge 9 commits into
sgl-project:mainfrom
yushengsu-thu:fsdp2/shape-buckets

Conversation

@yushengsu-thu

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

Copy link
Copy Markdown

Stacked on #922 → #915 (fsdp2/shape-buckets 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#6

Conclusion

  • Size over (fsdp2 dev. perf) Add training.compile_blocks: per-block torch.compile before fully_shard #922 (excluding tests and docs): +122 / −23 lines of Python (launch.py, collation.py, assembly.py, fsdp2.py, schema.py, backend.py, data/utils.py, disaggregated.py); tests +182, docs +1/−1.
  • Stacked on (fsdp2 dev. perf) Add training.compile_blocks: per-block torch.compile before fully_shard #922 (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:
  • 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 at max_length 8192 (real ShareGPT, buckets 1024/2048/4096/8192), static_shape_buckets + compile_blocks vs static_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).
  • Online e2e at 8192, static_shapes + compile_blocks vs plain fsdp2 (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).
  • Online e2e at 8192, static_shape_buckets + compile_blocks vs plain fsdp2: 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).
  • Offline e2e at seq 8192, synthetic lengths uniform in 1024–8192 (mean ≈4.6k, like the Qwen3.8-27B regeneration corpus), buckets + compile vs static + compile: DFlash2 — speed faster, −578 ms/step (1640 → 1062 ms, −35.2%); memory no change (35733 → 35962 MB, +229 MB). DSpark — speed faster, −949 ms/step (1854 → 905 ms, −51.2%); memory no change (23813 → 23689 MB, −124 MB).
  • Offline e2e at 8192, static + compile vs plain 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).
  • Offline e2e at 8192, buckets + compile vs plain 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).
  • Read the synthetic 8192 rows as "buckets vs static" only. Their pad-to-longest baseline gets two discounts real data does not give it: short contexts and fewer valid anchors (a random 20–80% prefix is unsupervised, so the dynamic path often runs fewer than num_anchors anchor 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.
  • Online e2e (real ShareGPT conversations, max_length 2048, buckets 512/1024/1536/2048), static_shape_buckets + compile_blocks vs static_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).
  • Online e2e, bucket padding alone (static_shape_buckets vs static_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).
  • Online e2e, the whole stack vs this box's pad-to-longest baseline (static_shape_buckets + compile_blocks vs fsdp, 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).
  • Verdict: the buckets pay off in proportion to the padding share. At max_length 8192 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; at max_length 2048 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).
  • Offline e2e, synthetic samples with variable lengths (25–100% of 2048) and variable masks, static_shape_buckets + compile_blocks vs static_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).
  • Offline e2e, variable lengths, bucket padding alone (static_shape_buckets vs static_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).
  • Offline e2e, variable lengths, static_shape_buckets + compile_blocks vs plain fsdp2 (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).
  • Offline e2e, variable lengths, static_shapes alone vs plain fsdp2 (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).
  • The synthetic variable-length set is deliberately harsh on static shapes: lengths are uniform in 25–100% of max_length and a 20–80% prefix is unsupervised, so many samples have fewer than num_anchors valid anchors and the fixed anchor count (plus the max_length padding) 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.
  • What the buckets buy depends on the corpus: ShareGPT at max_length 2048 pads 1.23× the positions of batch-2 pad-to-longest under static_shapes; the Qwen3.8-27B regeneration corpus (mean 4–4.5k tokens, p90 at max_length 8192, 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 to data.max_length. Short conversations then do more work than they need: the padding grows with the gap between the typical conversation and max_length. training.static_shape_buckets keeps 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:

  • multiples of 128: flex_attention's block mask works in 128-token blocks, so aligned lengths never pay for a partial block;
  • with fp8_linear the token count must be a multiple of 16, which multiples of 128 satisfy;
  • no bucket above data.max_length, and max_length is always the last bucket (appended if missing);
  • the anchor count does not follow the buckets: it stays num_anchors (the draft-token dimension is already static).

Stacks on #922 (fsdp2/compile-blocks). The BackendOptions scaffolding 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] (requires static_shapes; positive multiples of 128, strictly increasing; a Config-level check rejects buckets above data.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_features and DataCollatorWithPadding accept a tuple of lengths as pad_to.
  • launch.py: _static_pad_length(static_shapes, max_len, buckets) returns the fixed length or the ascending bucket tuple (with max_len appended); every entry point and the eval loader take static_shape_buckets; assembly.shape_buckets(cfg) computes the effective tuple once and assembly._backend_options() passes the count to the backend only when compile_blocks is on (bucket padding alone works on either backend).
  • training/backend.py + training/fsdp2.py: BackendOptions.compile_shape_buckets; above 1 the blocks are compiled with dynamic=False and _configure_bucketed_compile() raises Dynamo's per-frame recompile limit to 2 × 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 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-static data wait per step trainer GPU memory max Δ memory vs fsdp2-static
fsdp2 + static_shapes (#922) 24.94 959 ms – 1 ms 58034 MB –
fsdp2 + static_shape_buckets 24.56 974 ms +15 ms (+1.6%) 1 ms 55590 MB −2444 MB
fsdp2 + static_shapes + compile_blocks (#922) 30.92 768 ms −192 ms (−20.0%) 1 ms 51164 MB −6870 MB
fsdp2 + static_shape_buckets + compile_blocks 30.78 765 ms −194 ms (−20.2%) 1 ms 53374 MB −4660 MB

DSpark

variant steady samples/s (3 trainer GPUs) train compute per step Δ compute vs fsdp2-static data wait per step trainer GPU memory max Δ memory vs fsdp2-static
fsdp2 + static_shapes (#922) 35.83 668 ms – 1 ms 41388 MB –
fsdp2 + static_shape_buckets 34.68 690 ms +22 ms (+3.3%) 1 ms 40454 MB −934 MB
fsdp2 + static_shapes + compile_blocks (#922) 42.03 569 ms −99 ms (−14.8%) 1 ms 38270 MB −3118 MB
fsdp2 + static_shape_buckets + compile_blocks 42.07 569 ms −99 ms (−14.8%) 1 ms 39194 MB −2194 MB

Measured 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_length 8192 — 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 (pad to longest) 22.14 1076 ms – 1 ms 53204 MB –
fsdp2 + static_shapes + compile_blocks (#922) 29.48 811 ms −265 ms (−24.6%) 1 ms 61156 MB +7952 MB
fsdp2 + static_shape_buckets + compile_blocks 31.05 770 ms −306 ms (−28.4%) 1 ms 55280 MB +2076 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 (pad to longest) 33.21 721 ms – 1 ms 39136 MB –
fsdp2 + static_shapes + compile_blocks (#922) (38 steps logged) 39.35 608 ms −113 ms (−15.7%) 1 ms 46264 MB +7128 MB
fsdp2 + static_shape_buckets + compile_blocks (38 steps logged) 41.48 577 ms −144 ms (−20.0%) 1 ms 41216 MB +2080 MB

Fourth devbox, data.max_length 8192 (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 under max_length 8192) 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

variant steady step time Δ vs fsdp2 samples/s (4 GPUs) peak allocated per rank Δ memory
fsdp2 (pad to longest) 1316 ms – 24.3 38094 MB –
fsdp2 + static_shapes + compile_blocks (#922) 1640 ms +324 ms (+24.7%) 19.5 35733 MB -2361 MB
fsdp2 + static_shape_buckets + compile_blocks 1062 ms -254 ms (−19.3%) 30.1 35962 MB -2132 MB

DSpark

variant steady step time Δ vs fsdp2 samples/s (4 GPUs) peak allocated per rank Δ memory
fsdp2 (pad to longest) 562 ms – 56.9 25458 MB –
fsdp2 + static_shapes + compile_blocks (#922) 1854 ms +1291 ms (+229.6%) 17.3 23813 MB -1645 MB
fsdp2 + static_shape_buckets + compile_blocks 905 ms +342 ms (+60.9%) 35.4 23689 MB -1769 MB

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

variant steady step time Δ vs fsdp2 samples/s (4 GPUs) peak allocated per rank Δ memory
fsdp2 (pad to longest) 980 ms – 32.6 37332 MB –
fsdp2 + static_shapes 935 ms -45 ms (−4.6%) 34.2 37651 MB +319 MB
fsdp2 + static_shape_buckets 902 ms -79 ms (−8.0%) 35.5 37515 MB +182 MB
fsdp2 + static_shapes + compile_blocks (#922) 749 ms -231 ms (−23.6%) 42.7 34103 MB -3230 MB
fsdp2 + static_shape_buckets + compile_blocks 768 ms -212 ms (−21.7%) 41.7 34083 MB -3249 MB

DSpark

variant steady step time Δ vs fsdp2 samples/s (4 GPUs) peak allocated per rank Δ memory
fsdp2 (pad to longest) 310 ms – 103.1 24930 MB –
fsdp2 + static_shapes 692 ms +382 ms (+123.0%) 46.2 25058 MB +127 MB
fsdp2 + static_shape_buckets 634 ms +324 ms (+104.3%) 50.5 25020 MB +89 MB
fsdp2 + static_shapes + compile_blocks (#922) 588 ms +278 ms (+89.7%) 54.4 22329 MB -2602 MB
fsdp2 + static_shape_buckets + compile_blocks 566 ms +256 ms (+82.5%) 56.5 22501 MB -2430 MB

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.
Copilot AI balanced review requested due to automatic review settings October 4, 2026 10:49

Copilot AI left a comment

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Copilot was unable to review this pull request because the user who requested the review has reached their quota limit.

@yushengsu-thu
yushengsu-thu marked this pull request as draft October 4, 2026 11:07
@yushengsu-thu yushengsu-thu changed the title (fsdp2 dev.) Add training.static_shape_buckets: a few static lengths for compiled blocks (fsdp2 dev.) - follow-up : Add training.static_shape_buckets: a few static lengths for compiled blocks Oct 4, 2026
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.

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.

2 participants