Repository navigation
Only disable the compiler on the offload path in AlignDevicesHook - #4327
Open
jiqing-feng wants to merge 6 commits into
Open
jiqing-feng wants to merge 6 commits into
jiqing-feng wants to merge 6 commits into
Conversation
jiqing-feng
marked this pull request as draft
September 24, 2026 05:56
jiqing-feng
marked this pull request as ready for review
September 24, 2026 09:58
jiqing-feng
force-pushed
the
fix-sharded-model-compile
branch
from
September 24, 2026 10:31
a796585 to
8750e27
Compare
SunMarc
self-requested a review
September 25, 2026 13:56
Member
|
Can you fix the conflict first ? |
`AlignDevicesHook.pre_forward`/`post_forward` and `SequentialHook.pre_forward`/ `post_forward` are unconditionally wrapped in `torch.compiler.disable()`. Since `dispatch_model` attaches these hooks to every submodule as soon as a model is spread over more than one device, any multi-device model becomes impossible to capture in a single graph and `torch.compile(..., fullgraph=True)` aborts with `Unsupported: Skip calling torch.compiler.disable()'d function`. The disable is only actually required for the offload path, where `set_module_tensor_to_device` swaps module weights at runtime. For plain multi-device sharding the hook body is just `send_to_device()` calls, which Dynamo traces fine. Move the offload logic into `_load_offloaded_weights`/`_offload_weights` and keep `@_compiler_disable` only on those, so sharded models can be compiled with `fullgraph=True` while offloaded models keep the existing behaviour.
dispatch_model only ever produces AlignDevicesHook, never a SequentialHook, in both the sharded and the offloaded case. A SequentialHook only appears when two hooks are stacked, e.g. device_map plus layerwise casting, and there fullgraph=True fails either way because LayerwiseCastingHook carries its own @_compiler_disable. Removing the decorator there changes no observable behaviour, so restore it and keep this PR minimal. Rename _load_offloaded_weights to _onload_weights to pair with _offload_weights and match the onloading terminology used in utils/modeling.py, and group the two public hook methods together.
This reverts commit 95979d8.
jiqing-feng
force-pushed
the
fix-sharded-model-compile
branch
from
October 3, 2026 07:39
8750e27 to
80548f7
Compare
Contributor
Author
Done. |
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.
What's the problem
AlignDevicesHook.pre_forward/post_forwardare unconditionally wrapped intorch.compiler.disable().dispatch_modelattaches the hook to every submodule as soon as a model spans more than one device, so any multi-device model gets a disabled call on every module boundary andtorch.compile(..., fullgraph=True)fails.The disable is only needed on the offload path, where
set_module_tensor_to_deviceswaps module weights at runtime. For plain multi-device sharding the hook body is justsend_to_device()calls, which Dynamo traces without trouble.Fix
Move the offload logic into
_onload_weights/_offload_weightsand keep@_compiler_disableonly on those. Offloaded models keep exactly the previous behaviour. Ignoring the reindent, the change is 16 insertions / 10 deletions in one file.Reproduction
Run it with at least two accelerators visible.
max_memoryforces the split so the repro does not need a model large enough to overflow a single device.Before:
After: the model is captured as a single graph.
Validation
4x A100 80GB, torch 2.14.0+cu130, transformers 5.17.0. Sharding forced with
max_memory, every model verified to really span 4 devices. Correctness is compared on prefill argmax rather than raw logits, since inductor reorders bf16 reductions.fullgraphbeforefullgraphafterOffload paths are unchanged.
cpu_offload,disk_offloadand adevice_mapwith a"cpu"entry all produce bit-identical logits to a single-device reference, and all parameters return to meta after the forward. An offloaded model still graph-breaks and still failsfullgraph=True, now pointing at_onload_weightsinstead ofpre_forward, which is the intended behaviour.tests/test_hooks.py,tests/test_big_modeling.pyandtests/test_modeling_utils.pypass (89 passed, 7 skipped).Validation script:
On performance
The broken case is not only
fullgraph=True. Plaintorch.compileon a sharded model is currently slower than eager, because Dynamo leaves and re-enters the compiled region at every hook and never gets a region large enough to be worth fusing.64-token greedy decode, median of 5-7.
eageris the reference and is not touched by this PR:torch.compilebeforetorch.compileafterRead the third column against the second: before this PR, compiling a sharded model made it 1.5x slower than not compiling it. After, it is 1.8-2.5x faster than eager.
The cause is the graph break count, which is what this PR actually changes:
Other compile modes, for completeness:
mode="reduce-overhead"fullgraph=Truemode="reduce-overhead"fullgraph=Truefullgraph=Truelands within noise of the default mode, as expected once the graph is whole either way, andreduce-overheadadds nothing on top.The benchmark has to be run twice, once per build, because
fullgraph=Truecannot execute onmainat all. Measuring the "before" column on a patched install is wrong: the patch is exactly what removes the graph breaks, so both columns end up measuring the same graph.Note on
SequentialHookI left
SequentialHook.pre_forward/post_forwardalone, sincedispatch_modelonly ever producesAlignDevicesHook, in both the sharded and the offloaded case. ASequentialHookonly shows up when two hooks are stacked, e.g.device_mapplus layerwise casting, and therefullgraph=Truefails either way becauseLayerwiseCastingHookcarries its own@_compiler_disable.