Skip to content

fix(device): register the device factory as the FX device builtin - #124

Merged
yeahdongcn merged 2 commits into
mainfrom
xd/fx-device-builtin
Oct 8, 2026
Merged

yeahdongcn merged 2 commits into
mainfrom
xd/fx-device-builtin

Conversation

@yeahdongcn

Copy link
Copy Markdown
Collaborator

What / Why

torchada replaces torch.device with DeviceFactoryWrapper. torch.fx prints device constants in generated code as device(type=...) and binds the name device in the code's globals from its registered device builtin, which is still the original torch.device class. CodeGen._gen_python_code.add_global exempts torch.device from the qualified-name path with obj != torch.device; after the patch that comparison no longer matches the registered class, so device is never bound and any GraphModule holding a device constant fails with NameError: name 'device' is not defined when its code runs.

This affects torch.compile with the eager and aot_eager backends and plain fx.symbolic_trace output. Inductor does not execute the FX code, so Inductor compiles are not affected. vLLM-Omni's MAGI-2 test_native_compile*.py tests, which compile with backend="eager", fail this way on MUSA.

Change

_patch_torch_device registers DeviceFactoryWrapper as the FX device builtin through torch.fx.graph._register_custom_builtin. Updating _custom_builtins alone is not enough: _Namespace.create_name also checks _illegal_names, which still maps device to the original class, and renames the global to device_1. The registration updates both tables.

The generated source is unchanged, and the GraphModule import block now contains from torch import device, as it does without the device patch. Because that block is part of the FX graph cache key, a warm cache built with an earlier torchada misses once. torch.device identity, isinstance and equality semantics are unchanged.

Verification

New tests in tests/test_fx_device_builtin.py: fx.symbolic_trace and torch.compile with the eager and aot_eager backends on graphs holding a device constant, on CPU and on MUSA; generated code and import block identical to the unpatched ones; a device factory call inside a compiled region; and the import-order invariant that Dynamo's common_constant_types keeps the original torch.device.

Results on MTT S5000 (torch 2.11.0.post2, torch_musa 2.11.0.post2+musa5.2.0), on top of #120 and #121 (the torchada tree used for the vLLM-Omni runs):

  • tests/test_fx_device_builtin.py: 22 passed; on the unfixed source 19 fail with the NameError and 3 pass.
  • A probe that compiles graphs holding device constants with backend="eager" and "aot_eager", on CPU and on MUSA, and a device factory call inside a compiled region: every case passes with the fix and every case fails with the NameError without it. Inductor compiles give the same results with and without the fix.
  • vLLM-Omni MAGI-2 test_native_compile.py, test_native_compile_cuda.py and test_native_compile_distributed.py, on vLLM-Omni main c1e84ce94 and on [Perf] MAGI-2: drop the Ulysses send/receive buffer copies聽vllm-project/vllm-omni#8595: the three tests that compile with backend="eager" and failed with the NameError now pass on both. test_compiled_regions_match_eager_under_dlo_and_hsdp no longer reaches the NameError but fails later, in torch_musa's FSDP patch (it creates MUSA streams for an FSDP group on a CPU mesh), which is independent of this change; the two setup_compile tests need network access the test machine did not have.
  • pytest tests/ on CPU (torch 2.11.0+cpu): the same 11 failures as main, none in the new file. The full suite was not run on MUSA.

Not covered / notes

  • Only torch 2.11 was run.
  • _register_custom_builtin is a private torch.fx helper; if it is missing the registration is skipped with a debug log, as for the TorchScript builtin above it.
  • The version is not bumped here.

AI assistance: Claude Code drafted the change, the tests and this description and ran the MUSA validation listed above.

@yeahdongcn
yeahdongcn marked this pull request as draft October 8, 2026 06:53
FX prints torch.device constants as ``device(type=...)`` and binds that
name in the generated code's globals from its registered ``device``
builtin, which is the original torch.device class. Once torch.device is
replaced by DeviceFactoryWrapper, add_global's torch.device exemption no
longer matches the registered class, so ``device`` is never bound and any
GraphModule holding a device constant raises
``NameError: name 'device' is not defined``. This breaks torch.compile
with the eager and aot_eager backends on MUSA (Inductor does not execute
that code) as well as plain fx.symbolic_trace output.

Register the wrapper through torch.fx.graph._register_custom_builtin.
Updating only ``_custom_builtins`` is not enough: _Namespace.create_name
also checks ``_illegal_names``, which still maps ``device`` to the
original class, and renames the global to ``device_1``.

The generated source is unchanged, and the GraphModule import block
(part of the FX graph cache key) matches the one produced without the
device patch. torch.device identity, isinstance and equality semantics
are unchanged.

Signed-off-by: Xiaodong Ye <xiaodong.ye@mthreads.com>
Signed-off-by: Xiaodong Ye <xiaodong.ye@mthreads.com>
@yeahdongcn
yeahdongcn force-pushed the xd/fx-device-builtin branch from 8bcc469 to c878612 Compare October 8, 2026 09:18
@yeahdongcn
yeahdongcn marked this pull request as ready for review October 8, 2026 09:18
@yeahdongcn
yeahdongcn merged commit e07bd19 into main Oct 8, 2026
yeahdongcn added a commit that referenced this pull request Oct 10, 2026
The English and Chinese READMEs have not kept up with the May-October work.
This documents what merged, moves the version-gated shims into one table, and
refreshes the measured numbers that had gone stale.

Feature table
- CUDA memory-pool APIs, `torch.cuda.streams`, CUDA-graph executable rotation,
  `torch.cuda._get_device_index`, `get_memory_info()`, and the FlashAttention
  provider shims (#61, #98, #103, #106, #108, #115)
- the "What Works" table goes back to one line per feature; the paragraph-sized
  `log_` / `isfinite` / `out_dtype` cells move into the new section below

New "torch_musa Compatibility" section
- one table of every version-gated shim with the release it is installed on:
  the four `< 2.11.0.post2` patches (#106, #113, #124), the `< 2.13.0`
  `mm`/`bmm` `out_dtype=` backport (#116), the stable-ABI header backport (#86,
  #96), asynchronous `isfinite` (#120), and `torch.cuda.streams` (#98)

New "Environment Variables" section
- the graph-rotation knobs (#72), `TORCHADA_PLATFORM`, the C++ operator-override
  switches (#61, #128), and the two variables that were already documented

Corrected and extended details
- torch.compile: FX `device` builtin (#124), Dynamo's device-index helper (#108),
  `MUSA_VISIBLE_DEVICES` mirroring (#106)
- C++ extensions: nested `<torch/cuda.h>` porting (#95), stable-ABI
  `STABLE_TORCH_LIBRARY_IMPL` rekeying and stream helpers (#100), torch 2.6+
  `include_paths`/`library_paths` signatures (#121), stale JIT build locks (#128)
- MoE tables are generated from checked-in recipes (#115)
- unsupported CUDA runtime APIs as no-ops (#65)
- Performance: replace the 0.1.94 / torch_musa 2.7.1 numbers with the checked-in
  0.1.95 / 2.11.0.post2 entry, and stop claiming every fast path is under 200ns
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