Skip to content

(fsdp2 dev. perf) Add training.compile_blocks: per-block torch.compile before fully_shard - #922

Draft
yushengsu-thu wants to merge 8 commits into
sgl-project:mainfrom
yushengsu-thu:fsdp2/compile-blocks
Draft

yushengsu-thu wants to merge 8 commits into
sgl-project:mainfrom
yushengsu-thu:fsdp2/compile-blocks

Conversation

@yushengsu-thu

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

Copy link
Copy Markdown

Stacked on #915 (codex/fsdp2-backend): the diff below also shows #915's commits until it merges. Increment-only view for review: yushengsu-thu#1

Conclusion

  • Size over [verifying] Add optional FSDP2 training backend #915 (excluding tests and docs): +224 / −20 lines of Python (launch.py +53/−3, backend.py +54/−2 of which 52 are the shared BackendOptions scaffolding, fsdp2.py +21, dflash_family_model.py +16/−2, schema.py +16, collation.py +11/−2, hidden_states_data.py +11/−6, data/utils.py +11/−1, trainer.py +10/−1, assembly.py +8, disaggregated.py +7, eagle3/data.py +5/−3, model_providers.py +1); tests +380, docs +2/−1.
  • FSDP1 → FSDP2 alone ([verifying] Add optional FSDP2 training backend #915, no option): online (real data, first devbox) DFlash2 1074 → 1058 ms compute per step, DSpark 720 → 720 ms, both no change; offline e2e EAGLE3 faster, −132 ms/step (796 → 664 ms), DFlash2 / DSpark no change. Memory — EAGLE3 smaller, −1782 MB peak; DFlash2 / DSpark no change.
  • The online rows are the ones to read: real conversations have variable lengths, which is exactly what static_shapes is for; the offline rows use synthetic samples (fixed-length, or fixed-length with variable loss masks) and bracket the real case.
  • Online e2e (real ShareGPT conversations, variable lengths), FSDP2 + static_shapes + compile_blocks vs fsdp2: DFlash2 — speed faster, −315 ms per step (1081 → 765 ms compute per step, −29.2%); memory no change (53982 → 52432 MB nvidia-smi max, −1550 MB, within the granularity of reserved memory). DSpark — speed faster, −150 ms per step (719 → 569 ms compute per step, −20.9%); memory no change (39652 → 38982 MB nvidia-smi max, −670 MB, within the granularity of reserved memory).
  • Online e2e, static_shapes alone vs fsdp2: DFlash2 — speed faster, −117 ms per step (1081 → 963 ms compute per step, −10.9%); memory larger, small: +3410 MB (53982 → 57392 MB nvidia-smi max). DSpark — speed faster, −52 ms per step (719 → 667 ms compute per step, −7.2%); memory no change (39652 → 40950 MB nvidia-smi max, +1298 MB, within the granularity of reserved memory).
  • Online e2e, the compile part alone (static_shapes + compile_blocks vs static_shapes): DFlash2 — speed faster, −198 ms per step (963 → 765 ms compute per step, −20.6%); memory smaller, large: −4960 MB (57392 → 52432 MB nvidia-smi max). DSpark — speed faster, −98 ms per step (667 → 569 ms compute per step, −14.7%); memory no change (40950 → 38982 MB nvidia-smi max, −1968 MB, within the granularity of reserved memory).
  • Offline e2e, synthetic samples with variable loss masks, FSDP2 + static_shapes + compile_blocks vs fsdp2: DFlash2 — speed faster, −912 ms/step (1750 → 838 ms, −52.1%); memory smaller, large: −3549 MB (37652 → 34103 MB). DSpark — speed faster, −182 ms/step (823 → 641 ms, −22.1%); memory smaller, large: −2828 MB (25156 → 22329 MB).
  • Offline e2e, variable masks, static_shapes alone vs fsdp2: DFlash2 — speed faster, −738 ms/step (1750 → 1012 ms, −42.2%); memory no change (37652 → 37444 MB, −208 MB). DSpark — speed faster, −89 ms/step (823 → 734 ms, −10.8%); memory no change (25156 → 25058 MB, −99 MB).
  • Offline e2e, variable masks, the compile part alone (vs static_shapes): DFlash2 — speed faster, −174 ms/step (1012 → 838 ms, −17.2%); memory smaller, large: −3341 MB (37444 → 34103 MB). DSpark — speed faster, −93 ms/step (734 → 641 ms, −12.7%); memory smaller, large: −2729 MB (25058 → 22329 MB).
  • Online e2e, static_shapes on FSDP1 (fsdp + static_shapes vs fsdp; the option is backend-independent): DFlash2 — speed faster, −128 ms per step (1083 → 955 ms compute per step, −11.9%); memory larger, small: +2838 MB (54812 → 57650 MB nvidia-smi max). DSpark — speed faster, −54 ms per step (721 → 667 ms compute per step, −7.5%); memory no change (41340 → 42594 MB nvidia-smi max, +1254 MB, within the granularity of reserved memory).
  • Online e2e, FSDP1 → FSDP2 with static_shapes on both (fsdp2 + static_shapes vs fsdp + static_shapes): DFlash2 — speed no change (955 → 959 ms compute per step, +5 ms, +0.5%); memory no change (57650 → 58034 MB nvidia-smi max, +384 MB, within the granularity of reserved memory). DSpark — speed no change (667 → 668 ms compute per step, +1 ms, +0.1%); memory no change (42594 → 41388 MB nvidia-smi max, −1206 MB, within the granularity of reserved memory).
  • Offline e2e, variable masks, static_shapes on FSDP1 (fsdp + static_shapes vs fsdp): DFlash2 — speed faster, −749 ms/step (1748 → 999 ms, −42.9%); memory no change (37648 → 37648 MB, +1 MB). DSpark — speed faster, −89 ms/step (825 → 736 ms, −10.8%); memory no change (25156 → 25056 MB, −100 MB).
  • Offline e2e, synthetic fixed-length inputs (the first measurement; shapes never change there, so it is the best case for compile): speed — faster: DFlash2 −182 ms/step (995 → 813 ms, −18.3%), DSpark −88 ms (726 → 638, −12.1%), EAGLE3 −95 ms (664 → 569, −14.3%). Memory — smaller: peak −3252 MB / −2730 MB / −784 MB per rank.
  • Cost: one-time compile warm-up ≈12 s (DFlash family) / ≈80 s (EAGLE3); static_shapes pads every micro-batch to data.max_length and every anchor set to num_anchors, so short conversations do more work per sample than with pad-to-longest (ShareGPT at max_length 2048: 1.23× the positions of batch-2 pad-to-longest; the variable-mask and online tables below include that cost). Numerics: loss and grad norm match fsdp2 to 5 digits with fixed shapes; with static_shapes the kept anchors are the same set, the padded slots are masked.

Motivation

The FSDP2 backend from #915 keeps each draft block an ordinary module (composable hooks, no wrapper), which is what makes per-block torch.compile practical: FSDP1's FullyShardedDataParallel wrappers sit between the blocks and Dynamo. This PR adds the opt-in training.compile_blocks. No torchtitan involved.

Compiled blocks need one input shape per run. With real data the padded context length changes from one micro-batch to the next (batches are padded to their longest sample) and the number of valid anchors changes with it; the first online run of this PR crashed at the first such change (InductorError: CantSplit: 8*s50 + 65536 not divisible by s50 + 8192: Dynamo recompiled with the context length symbolic and torch 2.13's Inductor could not lower flex_attention with it), and where the dynamic-shape graph does compile it gives back none of the fixed-shape gain. training.static_shapes therefore pads every micro-batch to data.max_length and keeps num_anchors anchor slots per sample. compile_blocks does not require it (fixed-shape inputs, e.g. long-form data truncated at max_length, run compiled blocks fine without the padding), but the validator warns when it is off and names the failure mode.

Stacks directly on #915 (base branch codex/fsdp2-backend). The BackendOptions scaffolding is shared verbatim with the other FSDP2 option PRs so they can land in any order.

Modifications

  • training/backend.py: BackendOptions (typed config → backend), _block_targets() (the _no_split_modules classes, or the EAGLE midlayer), _prepare_blocks() hook before sharding. FSDP1 rejects any option.
  • training/fsdp2.py: with compile_blocks, nn.Module.compile() every block in place before fully_shard, so the block keeps its class (the per-block FSDP boundary still matches) and its parameter FQNs (no _orig_mod. prefix in checkpoints/exports). The FSDP2 hooks registered afterwards run inside the compiled call, but Dynamo skips them (skip_fsdp_hooks), so they execute eagerly around the compiled block body.
  • training.static_shapes (config/schema.py; recommended with compile_blocks, the validator warns without it): the DFlash-family collators (algorithms/common/collation.py, hidden_states_data.py) and the EAGLE3 offline collator (data/utils.py) take pad_to; launch.py resolves them with pad_to=data.max_length on all four entry points (offline, disaggregated offline, online consumer, eval loader); algorithms/common/dflash_family_model.py samples exactly num_anchors slots (static_anchor_count, from model_providers.py), the slots past a row's valid anchors are masked exactly like today's short rows.
  • config/schema.py: training.compile_blocks (requires backend: fsdp2; warns without static_shapes); plumbing through assembly, launch, trainer.
  • training/disaggregated.py + assembly._backend_options(): the disaggregated runtime builds its trainers through its own call sites, which did not forward backend_options; the option was silently off in disaggregated (every online) runs until the online e2e of this series caught it. Both call sites now forward it, with tests for the offline and online consumer paths.
  • Tests: tests/test_runtime/test_fsdp2_compile_blocks.py (FSDP1 rejection, config validation, midlayer fallback, 2-GPU eager-vs-compiled parity, disaggregated forwarding) and tests/test_runtime/test_static_shapes.py (fixed-length collation for both collators, fixed anchor count with identical kept anchors, launch resolution, config).

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 31.24 765 ms −315 ms (−29.2%) 1 ms 52432 MB −1550 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 41.99 569 ms −150 ms (−20.9%) 1 ms 38982 MB −670 MB

The first online attempt of this PR (first devbox, before static_shapes) crashed at the first micro-batch whose padded length differed: InductorError: CantSplit: 8*s50 + 65536 not divisible by s50 + 8192. With static_shapes there is one graph per block for the whole run.

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.

FSDP1 vs FSDP2 once both have static_shapes (third devbox)

static_shapes lives in the data path and the DFlash-family model, not in the backend, so FSDP1 gets it too. These runs put FSDP1, FSDP1 + static_shapes and FSDP2 + static_shapes on one machine; the compile step (#922's other half) remains FSDP2-only.

Online (disaggregated, real ShareGPT prompts)

Same setup as the online tables above (DP=3 trainers, 40 steps, steady state = last 24 steps).

DFlash2

variant steady samples/s (3 trainer GPUs) train compute per step Δ compute vs fsdp data wait per step trainer GPU memory max Δ memory vs fsdp
fsdp (FSDP1) 22.10 1083 ms – 1 ms 54812 MB –
fsdp + static_shapes 25.06 955 ms −128 ms (−11.9%) 1 ms 57650 MB +2838 MB
fsdp2 + static_shapes 24.94 959 ms −124 ms (−11.4%) 1 ms 58034 MB +3222 MB

DSpark

variant steady samples/s (3 trainer GPUs) train compute per step Δ compute vs fsdp data wait per step trainer GPU memory max Δ memory vs fsdp
fsdp (FSDP1) 33.19 721 ms – 1 ms 41340 MB –
fsdp + static_shapes 35.87 667 ms −54 ms (−7.5%) 1 ms 42594 MB +1254 MB
fsdp2 + static_shapes 35.83 668 ms −53 ms (−7.4%) 1 ms 41388 MB +48 MB

Measured on a third 4 × B300 devbox so that FSDP1 and FSDP2 share one baseline; same configuration as the other online tables.

Offline, variable loss masks (synthetic, DP=4)

DFlash2

variant steady step time Δ vs fsdp samples/s (4 GPUs) peak allocated per rank Δ memory
fsdp (FSDP1) 1748 ms – 18.3 37648 MB –
fsdp + static_shapes 999 ms -749 ms (−42.9%) 32.0 37648 MB +1 MB

DSpark

variant steady step time Δ vs fsdp samples/s (4 GPUs) peak allocated per rank Δ memory
fsdp (FSDP1) 825 ms – 38.8 25156 MB –
fsdp + static_shapes 736 ms -89 ms (−10.8%) 43.5 25056 MB -100 MB

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.

With pad-to-longest and a varying anchor count the fsdp2 baseline itself is much slower than with fixed shapes (the flex block mask and the anchor tensors are rebuilt for every new shape); static_shapes alone removes that, and the compile gain comes back on top.

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 1012 ms -738 ms (−42.2%) 31.6 37444 MB -208 MB
fsdp2 + static_shapes + compile_blocks 838 ms -912 ms (−52.1%) 38.2 34103 MB -3549 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 734 ms -89 ms (−10.8%) 43.6 25058 MB -99 MB
fsdp2 + static_shapes + compile_blocks 641 ms -182 ms (−22.1%) 49.9 22329 MB -2828 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_blocks step time Δ time samples/s (4 GPUs) peak allocated per rank Δ memory fit wall (48 steps, incl. warm-up + final ckpt)
DFlash2 978 → 995 ms (+17 ms, no change) 995 → 813 ms −182 ms (−18.3%) 32.2 → 39.4 37652 → 34400 MB −3252 MB 62.5 → 54.8 s
DSpark 740 → 726 ms (−14 ms, no change) 726 → 638 ms −88 ms (−12.1%) 44.1 → 50.2 25058 → 22328 MB −2730 MB 48.5 → 62.1 s
EAGLE3 796 → 664 ms (−132 ms) 664 → 569 ms −95 ms (−14.3%) 48.2 → 56.2 21168 → 20384 MB −784 MB 42.7 → 122.3 s

Fixed-length synthetic samples never change shape, so this table shows the compile gain without static_shapes (the option changes nothing there).

Notes: with fixed shapes the step-level win survives the full loop and nothing else in the loop got slower; warm-up break-even ≈70 DFlash2 steps / ≈850 EAGLE3 steps. DFlash2's Triton grouped convolution turns itself off under torch.compiler.is_compiling() and runs eager inside the graph (keeping it as a custom op is a follow-up). The benchmark image has no Liger; recipes with use_liger_kernel: true should re-measure.

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.
@yushengsu-thu yushengsu-thu changed the title (fsdp2 dev.) Add training.compile_blocks: per-block torch.compile before fully_shard (fsdp2 dev. perf) Add training.compile_blocks: per-block torch.compile before fully_shard Oct 4, 2026
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.
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.

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.

1 participant