(fsdp2 dev. perf) Add training.compile_blocks: per-block torch.compile before fully_shard - #922
Draft
yushengsu-thu wants to merge 8 commits into
Draft
yushengsu-thu wants to merge 8 commits into
yushengsu-thu wants to merge 8 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/compile-blocks
branch
from
October 4, 2026 06:09
5db77ab to
177dc64
Compare
This was referenced Oct 4, 2026
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/compile-blocks
branch
from
October 4, 2026 09:05
e958320 to
7815e57
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.
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
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+53/−3,backend.py+54/−2 of which 52 are the sharedBackendOptionsscaffolding,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.static_shapesis for; the offline rows use synthetic samples (fixed-length, or fixed-length with variable loss masks) and bracket the real case.static_shapes+compile_blocksvsfsdp2: 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).static_shapesalone vsfsdp2: 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).static_shapes+compile_blocksvsstatic_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).static_shapes+compile_blocksvsfsdp2: 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).static_shapesalone vsfsdp2: 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).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).static_shapeson FSDP1 (fsdp+static_shapesvsfsdp; 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).static_shapeson both (fsdp2+static_shapesvsfsdp+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).static_shapeson FSDP1 (fsdp+static_shapesvsfsdp): 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).static_shapespads every micro-batch todata.max_lengthand every anchor set tonum_anchors, so short conversations do more work per sample than with pad-to-longest (ShareGPT atmax_length2048: 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 matchfsdp2to 5 digits with fixed shapes; withstatic_shapesthe 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.compilepractical: FSDP1'sFullyShardedDataParallelwrappers sit between the blocks and Dynamo. This PR adds the opt-intraining.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 lowerflex_attentionwith it), and where the dynamic-shape graph does compile it gives back none of the fixed-shape gain.training.static_shapestherefore pads every micro-batch todata.max_lengthand keepsnum_anchorsanchor slots per sample.compile_blocksdoes not require it (fixed-shape inputs, e.g. long-form data truncated atmax_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). TheBackendOptionsscaffolding 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_modulesclasses, or the EAGLEmidlayer),_prepare_blocks()hook before sharding. FSDP1 rejects any option.training/fsdp2.py: withcompile_blocks,nn.Module.compile()every block in place beforefully_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 withcompile_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) takepad_to;launch.pyresolves them withpad_to=data.max_lengthon all four entry points (offline, disaggregated offline, online consumer, eval loader);algorithms/common/dflash_family_model.pysamples exactlynum_anchorsslots (static_anchor_count, frommodel_providers.py), the slots past a row's valid anchors are masked exactly like today's short rows.config/schema.py:training.compile_blocks(requiresbackend: fsdp2; warns withoutstatic_shapes); plumbing throughassembly,launch,trainer.training/disaggregated.py+assembly._backend_options(): the disaggregated runtime builds its trainers through its own call sites, which did not forwardbackend_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/test_runtime/test_fsdp2_compile_blocks.py(FSDP1 rejection, config validation, midlayer fallback, 2-GPU eager-vs-compiled parity, disaggregated forwarding) andtests/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 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_blocksDSpark
fsdp2fsdp2fsdp2fsdp2+static_shapesfsdp2+static_shapes+compile_blocksThe 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. Withstatic_shapesthere 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
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.FSDP1 vs FSDP2 once both have
static_shapes(third devbox)static_shapeslives in the data path and the DFlash-family model, not in the backend, so FSDP1 gets it too. These runs put FSDP1, FSDP1 +static_shapesand FSDP2 +static_shapeson 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
fsdpfsdpfsdp(FSDP1)fsdp+static_shapesfsdp2+static_shapesDSpark
fsdpfsdpfsdp(FSDP1)fsdp+static_shapesfsdp2+static_shapesMeasured 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
fsdpfsdp(FSDP1)fsdp+static_shapesDSpark
fsdpfsdp(FSDP1)fsdp+static_shapesEnd-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
fsdp2baseline itself is much slower than with fixed shapes (the flex block mask and the anchor tensors are rebuilt for every new shape);static_shapesalone removes that, and the compile gain comes back on top.DFlash2
fsdp2fsdp2fsdp2+static_shapesfsdp2+static_shapes+compile_blocksDSpark
fsdp2fsdp2fsdp2+static_shapesfsdp2+static_shapes+compile_blocksEnd-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_blocksstep timeFixed-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 withuse_liger_kernel: trueshould re-measure.