Skip to content

Only disable the compiler on the offload path in AlignDevicesHook - #4327

Open
jiqing-feng wants to merge 6 commits into
huggingface:mainfrom
jiqing-feng:fix-sharded-model-compile
Open

jiqing-feng wants to merge 6 commits into
huggingface:mainfrom
jiqing-feng:fix-sharded-model-compile

Conversation

@jiqing-feng

@jiqing-feng jiqing-feng commented Sep 24, 2026 •

Copy link
Copy Markdown
Contributor

What's the problem

AlignDevicesHook.pre_forward / post_forward are unconditionally wrapped in torch.compiler.disable(). dispatch_model attaches 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 and torch.compile(..., fullgraph=True) fails.

The disable is only needed on 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 without trouble.

Fix

Move the offload logic into _onload_weights / _offload_weights and keep @_compiler_disable only on those. Offloaded models keep exactly the previous behaviour. Ignoring the reindent, the change is 16 insertions / 10 deletions in one file.

Reproduction

import torch
from transformers import AutoModelForCausalLM, AutoTokenizer

MODEL = "Qwen/Qwen3-1.7B"

tokenizer = AutoTokenizer.from_pretrained(MODEL)
model = AutoModelForCausalLM.from_pretrained(
    MODEL, device_map="auto", max_memory={i: "1GiB" for i in range(torch.cuda.device_count())}, dtype=torch.bfloat16
).eval()
print("weight devices:", sorted({str(p.device) for p in model.parameters()}))

inputs = tokenizer("The capital of France is", return_tensors="pt").to(model.device)
model.forward = torch.compile(model.forward, fullgraph=True)

with torch.no_grad():
    print(model(**inputs).logits.shape)

Run it with at least two accelerators visible. max_memory forces the split so the repro does not need a model large enough to overflow a single device.

Before:

weight devices: ['cuda:0', 'cuda:1', 'cuda:2', 'cuda:3']
torch._dynamo.exc.Unsupported: Skip calling `torch.compiler.disable()`d function
  Explanation: Skip calling function `<function AlignDevicesHook.pre_forward ...>`

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.

model size fullgraph before fullgraph after prefill argmax
Qwen/Qwen2.5-14B-Instruct 27.51 GiB fails ok match
Qwen/Qwen3-8B 15.26 GiB fails ok match
HuggingFaceTB/SmolLM2-135M-Instruct 0.30 GiB fails ok match

Offload paths are unchanged. cpu_offload, disk_offload and a device_map with 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 fails fullgraph=True, now pointing at _onload_weights instead of pre_forward, which is the intended behaviour.

tests/test_hooks.py, tests/test_big_modeling.py and tests/test_modeling_utils.py pass (89 passed, 7 skipped).

Validation script:

import gc, sys, torch
from accelerate import init_empty_weights
from accelerate.utils.modeling import compute_module_sizes
from transformers import AutoConfig, AutoModelForCausalLM, AutoTokenizer

n = torch.cuda.device_count()

for model_id in sys.argv[1:]:
    with init_empty_weights():
        meta = AutoModelForCausalLM.from_config(AutoConfig.from_pretrained(model_id), dtype=torch.bfloat16)
    total = compute_module_sizes(meta)[""]
    del meta
    gc.collect()

    # Cap each device below the model size so the weights are guaranteed to span several.
    cap = min(int(total * 0.3), int(min(torch.cuda.mem_get_info(i)[0] for i in range(n)) * 0.8))
    model = AutoModelForCausalLM.from_pretrained(
        model_id, device_map="auto", max_memory={i: cap for i in range(n)}, dtype=torch.bfloat16
    ).eval()
    devices = sorted({str(p.device) for p in model.parameters()})
    assert len(devices) > 1, devices

    tok = AutoTokenizer.from_pretrained(model_id)
    inputs = tok("The capital of France is", return_tensors="pt").to(model.device)
    with torch.no_grad():
        eager = model(**inputs).logits.float().cpu()
        compiled = torch.compile(model.forward, fullgraph=True)(**inputs).logits.float().cpu()

    print(f"{model_id}: {total / 2**30:.2f} GiB over {devices}, fullgraph OK")
    print(f"  argmax match {torch.equal(compiled.argmax(-1), eager.argmax(-1))}")

    del model
    gc.collect()
    torch.cuda.empty_cache()
    torch._dynamo.reset()

On performance

The broken case is not only fullgraph=True. Plain torch.compile on 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. eager is the reference and is not touched by this PR:

model eager torch.compile before torch.compile after speedup
Qwen/Qwen2.5-14B-Instruct 3095 ms 4682 ms 1700 ms 2.75x
Qwen/Qwen3-8B 2722 ms 4134 ms 1071 ms 3.86x

Read 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:

model before after
Qwen/Qwen2.5-14B-Instruct 11 0
Qwen/Qwen3-8B 13 0

Other compile modes, for completeness:

model mode before after
Qwen/Qwen2.5-14B-Instruct mode="reduce-overhead" 4689 ms 1718 ms
Qwen/Qwen2.5-14B-Instruct fullgraph=True raises 1707 ms
Qwen/Qwen3-8B mode="reduce-overhead" 4143 ms 1090 ms
Qwen/Qwen3-8B fullgraph=True raises 1082 ms

fullgraph=True lands within noise of the default mode, as expected once the graph is whole either way, and reduce-overhead adds nothing on top.

The benchmark has to be run twice, once per build, because fullgraph=True cannot execute on main at 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.

import gc, statistics, sys, time, torch
from transformers import AutoModelForCausalLM, AutoTokenizer

n = torch.cuda.device_count()

def timed(fn, iters=5):
    for _ in range(3):
        fn()
    torch.cuda.synchronize()
    s = []
    for _ in range(iters):
        t0 = time.perf_counter()
        fn()
        torch.cuda.synchronize()
        s.append(time.perf_counter() - t0)
    return statistics.median(s) * 1e3

for model_id in sys.argv[1:]:
    # Same max_memory capping as the validation script above.
    model = AutoModelForCausalLM.from_pretrained(
        model_id, device_map="auto", max_memory={i: cap for i in range(n)}, dtype=torch.bfloat16
    ).eval()
    tok = AutoTokenizer.from_pretrained(model_id)
    inputs = tok("The capital of France is", return_tensors="pt").to(model.device)
    gk = dict(max_new_tokens=64, min_new_tokens=64, do_sample=False)
    eager_forward = model.forward

    def run():
        with torch.no_grad():
            model.generate(**inputs, **gk)

    print(f"=== {model_id}")
    print(f"  eager         {timed(run):9.1f} ms")
    print(f"  graph breaks  {torch._dynamo.explain(eager_forward)(**inputs).graph_break_count}")

    for tag, kw in [("compile", {}), ("reduce-overhead", {"mode": "reduce-overhead"}), ("fullgraph", {"fullgraph": True})]:
        torch._dynamo.reset()
        model.forward = torch.compile(eager_forward, **kw)
        try:
            print(f"  {tag:15s} {timed(run):9.1f} ms")
        except Exception as exc:
            print(f"  {tag:15s} FAILED: {str(exc).splitlines()[0]}")

    del model
    gc.collect()
    torch.cuda.empty_cache()
    torch._dynamo.reset()

Note on SequentialHook

I left SequentialHook.pre_forward / post_forward alone, since dispatch_model only ever produces AlignDevicesHook, in both the sharded and the offloaded case. A SequentialHook only shows up 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.

@jiqing-feng
jiqing-feng marked this pull request as draft September 24, 2026 05:56
@jiqing-feng
jiqing-feng marked this pull request as ready for review September 24, 2026 09:58
@jiqing-feng
jiqing-feng force-pushed the fix-sharded-model-compile branch from a796585 to 8750e27 Compare September 24, 2026 10:31
@SunMarc
SunMarc self-requested a review September 25, 2026 13:56
@SunMarc

SunMarc commented Sep 25, 2026

Copy link
Copy Markdown
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.
@jiqing-feng
jiqing-feng force-pushed the fix-sharded-model-compile branch from 8750e27 to 80548f7 Compare October 3, 2026 07:39
@jiqing-feng

Copy link
Copy Markdown
Contributor Author

Can you fix the conflict first ?

Done.

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