Skip to content

fix(musa): keep torch.isfinite asynchronous on MUSA floating tensors - #120

Merged
yeahdongcn merged 1 commit into
mainfrom
xd/musa-isfinite-async
Oct 8, 2026
Merged

yeahdongcn merged 1 commit into
mainfrom
xd/musa-isfinite-async

Conversation

@yeahdongcn

@yeahdongcn yeahdongcn commented Oct 5, 2026 •

Copy link
Copy Markdown
Collaborator

What / Why

On torch_musa 2.11.0.post2+musa5.2.0, every boolean mul blocks the host until the device queue drains. That covers boolbool, torch.mul, mul_ and boolscalar.

ATen computes isfinite of a floating tensor as (x == x) * (x.abs() != inf), so each torch.isfinite call on a MUSA tensor drains the queue too. With about 70 ms of GEMMs queued, isfinite blocked the host for 73-78 ms in float16, float32 and float64 alike. &, logical_and, integer mul and x.abs() < inf stayed asynchronous at 0.02-0.10 ms.

In vLLM-Omni's MAGI-2 Preview, the eager attention-sink correction calls isfinite once per attention layer, so every layer waited for the GPU to go idle.

Change

  • For MUSA float16, bfloat16, float32 and float64 tensors, torch.isfinite and Tensor.isfinite return x.abs() < inf, the same predicate: False for NaN and for both infinities.
  • Every other dtype and device keeps the original op:
    • integer and bool tensors are always finite;
    • a complex magnitude can overflow when both parts are finite;
    • fp8 types do not reliably implement abs and lt.
  • The wrapper is registered as a TorchScript alias of aten::isfinite.
  • torch.isfinite is added to the AOT autograd cache safelist, so compiled graphs that call it stay cacheable.
  • There is no version gate, because the predicate is exact on every stack. The shim can be deleted once torch_musa's boolean mul stops synchronizing.
  • Version 0.1.91.

Verification

tests/test_isfinite.py covers:

  • values against CPU in f16/bf16/f32/f64, for NaN, 卤inf, -0.0, a subnormal, 卤finfo.max and finfo.tiny, on contiguous, transposed and strided views;
  • 0-d and empty tensors;
  • which inputs take the new path, checked with a spy:
    • MUSA f16/bf16/f32/f64 tensors through both torch.isfinite and Tensor.isfinite;
    • and not int32/int64/bool/complex64, nor CPU tensors;
  • integer, bool and fp8 inputs returning the original op's result;
  • no autograd tracking;
  • work queued before the call still pending afterwards (checked with stream.query()), for f16/bf16/f32/f64 through both entry points. The test first asserts that the queue is busy, and skips under MUSA_LAUNCH_BLOCKING/CUDA_LAUNCH_BLOCKING;
  • TorchScript;
  • torch.compile(fullgraph=True) values and exactly one AOT autograd cache artifact, checked in a fresh interpreter.

On stock 0.1.90 all 8 queued-work cases (4 dtypes, 2 entry points) fail with "isfinite drained the device queue".

Results on MTT S5000:

  • pytest tests/: 631 passed, 19 skipped, 2 failed. The two failures are in TestInductorTemplateHeuristics (KeyError: ('triton::bmm', 'musa', None) and the same for triton::mm) and fail the same way on 0.1.90 without this change.

End to end, on MAGI-2 Preview on 8x MTT S5000 (SP4xCFG2, 100 steps):

  • on an earlier development tree, removing the stall cut a step from 2894 ms to 1137 ms;
  • with an earlier revision of this change (same routing for those fp32 tensors) and the model unchanged, a step is no slower than with the model's two isfinite calls rewritten (the mean step time differs by under 1 ms);
  • the generated video and audio are byte-identical (sha256) between the two.

Not covered / notes

  • Code that branches on the result on the host, such as if not torch.isfinite(x).all():, still waits for the device, by design.
  • TorchScript keeps calling aten::isfinite, so scripted code still takes the blocking path.
  • The boolean mul synchronization itself is a torch_musa issue and should be fixed there; this shim only takes isfinite off that path.

AI assistance: Claude Code drafted the change, the tests and this description and ran the MUSA validation listed above. I reviewed every change in this PR and re-ran the validation listed above.

On torch_musa 2.11.0.post2+musa5.2.0 every boolean mul blocks the host until
the device queue drains, and ATen computes isfinite as
(x == x) * (x.abs() != inf). Each torch.isfinite call on a MUSA tensor
therefore holds back later kernel launches until all queued work has
finished: 73-78 ms with about 70 ms of GEMMs queued, for float16, float32
and float64 alike. &, logical_and, integer mul and x.abs() < inf stay
asynchronous (0.02-0.10 ms).

torch.isfinite and Tensor.isfinite now return x.abs() < inf for MUSA
float16, bfloat16, float32 and float64 tensors; every other dtype and
device keeps the original op. The wrapper stays scriptable as
aten::isfinite and is registered in the AOT autograd cache safelist.

MAGI-2 Preview in vLLM-Omni calls isfinite in its eager attention-sink
correction once per layer. With an earlier revision of this change (same
routing for those fp32 tensors), a 100-step SP4xCFG2 run on 8x MTT S5000
is no slower than rewriting the model's isfinite calls (the mean step
time differs by under 1 ms), and the generated output is byte-identical.

Signed-off-by: Xiaodong Ye <xiaodong.ye@mthreads.com>
@yeahdongcn
yeahdongcn force-pushed the xd/musa-isfinite-async branch from 0bd3373 to 1c587e6 Compare October 5, 2026 09:19
@yeahdongcn
yeahdongcn marked this pull request as ready for review October 5, 2026 09:49
@yeahdongcn
yeahdongcn merged commit 80968f9 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