From 62a32ec71547e77f503f94ba6ac9bcf3615f93e7 Mon Sep 17 00:00:00 2001 From: Xiaodong Ye Date: Sun, 20 Sep 2026 15:57:27 +0800 Subject: [PATCH] fix(musa): backport torch.mm/bmm out_dtype semantics (torch_musa < 2.13.0) On the affected torch_musa stack the vendored `aten::mm.dtype` / `aten::bmm.dtype` overloads are registered but do not write their result: `torch.mm(..., out_dtype=)` returns a correctly shaped fp32 tensor that is silently all zeros, and `bmm` returns non-zero wrong values. Measured on 2.11.0.post1+musa5.2.0; plain and invalid arguments behave correctly, so only the promoted path is affected. The wrappers are armed below `2.13.0`, the release torch_musa committed to fix the overloads in - a vendor commitment, not a measurement of ours - and nothing is installed from 2.13.0 on. An unknown or unparsable __version__ ranks lowest in version_of and therefore stays armed, so a stack we cannot read is never assumed fixed. Once 2.13.0 is released and verified fixed here the shim is deleted; if the fix slips the bound moves to the newly committed release. The backport promotes the operands to fp32 and accumulates there, reusing the plain overloads, so it assumes nothing about the vendor kernel; every gate compares against an inline literal (no version constants), invalid dtype combinations are forwarded verbatim so the vendor error text stays byte-identical, and `_version.version_of` stays as the parse-failure-safe comparator. Correctness is decided per process by a three-valued, lazy probe: it runs only on the first promoted call, caches only deterministic verdicts, warns once when it cannot decide, and refuses to probe inside a CUDA/MUSA graph capture (the verdict then resolves on the next eager call, which is why a graph captured before the first eager call must be re-captured on a fixed stack). Verified on S5000: focused 307 passed/4 skipped, full 554 passed/19 skipped, gate table 2.10.0..2.12.0 armed and 2.13.0+ unarmed, capture safety, a no-shim control showing the all-zero/wrong-value defect, and the end-to-end A/B with byte-identical outputs (COMPARABLE, +0.044 ms). See the ticket for the full evidence list. Version unified on 0.1.89 as a true replace-all: every `0.1.x` literal in the tree now reads 0.1.89, including the C++ mirror in csrc/ops.h (was 0.1.0), the recorded benchmark history entry (was 0.1.86) and the example extension setup (was 0.1.0); a tree-wide grep for any other `0.1.x` literal returns nothing. --- README.md | 16 +- README_CN.md | 11 +- benchmarks/benchmark_history.json | 2 +- examples/extension_setup.py | 2 +- pyproject.toml | 2 +- src/torchada/__init__.py | 2 +- src/torchada/_cpp_ops.py | 8 +- src/torchada/_patch.py | 473 +++++++++++++++++--- src/torchada/_version.py | 147 +++++++ src/torchada/csrc/ops.h | 2 +- tests/test_cuda_patching.py | 14 +- tests/test_mm_out_dtype.py | 702 ++++++++++++++++++++++++++++++ tests/test_platform.py | 2 +- tests/test_version.py | 165 +++++++ 14 files changed, 1483 insertions(+), 65 deletions(-) create mode 100644 src/torchada/_version.py create mode 100644 tests/test_mm_out_dtype.py create mode 100644 tests/test_version.py diff --git a/README.md b/README.md index 68b8c44..dca8a68 100644 --- a/README.md +++ b/README.md @@ -67,9 +67,23 @@ That's it! Supported `torch.cuda.*` APIs are automatically redirected to `torch. | ctypes Libraries | `ctypes.CDLL` with CUDA function names → MUSA equivalents | | Unified Accelerator API | `torch.accelerator.empty_cache()`, `memory_stats()`, `Stream`, `Event`, ... | | MUSA float64 in-place log | On `torch_musa < 2.11.0.post2`, `Tensor.log_()` reuses the supported out-of-place operation while preserving the in-place contract | +| MUSA mm/bmm `out_dtype` | Below `torch_musa 2.13.0`, `torch.mm`/`torch.bmm` `out_dtype=` reuses the plain overloads and accumulates in fp32 while a runtime probe reports the overload broken. Measured broken on `2.11.0.post1+musa5.2.0` (`mm` writes zeros, `bmm` writes wrong values); **torch_musa committed to fix this in `2.13.0`, which we have not verified** — the wrappers are armed below that release, nothing is installed from it on, and the probe decides correctness per process | | Triton CUDA Extra | `tl.extra.cuda` → `tl.extra.musa` compatibility on MUSA | | Triton Fused MoE | Triton 3.2.0 MTT S5000 tuning configs for vLLM and SGLang | +**Not covered:** builds whose binding has no `*_Dtype` overload keep raising on `out_dtype=`. That is +the binding's contract, not a defect torchada repairs, so CUDA parity for that case is out of scope. + +The `out_dtype` backport is armed **below `torch_musa 2.13.0`**, the release torch_musa committed to +fix the overloads in. That is a vendor release commitment, not a measurement of ours, and trusting it +is the accepted risk: if `2.13.0` does not actually fix them, a `>= 2.13.0` stack installs no wrapper +and the silent all-zero result can come back. The gate only decides whether a Python wrapper sits in +front of `torch.mm`/`torch.bmm`; correctness is decided per process by the runtime probe, which +forwards to a healthy overload and emulates a broken one - so a fix backported into `2.12.x` is picked +up automatically. An unknown or unparsable `__version__` ranks lowest and therefore stays armed. Once +`2.13.0` is released and verified fixed here, the shim is deleted; if the fix slips, the bound moves to +the newly committed release. + ## Examples ### Mixed Precision Training @@ -393,7 +407,7 @@ See `src/torchada/_mappings/` for 400+ mapping rules grouped by API domain. ``` # pyproject.toml or requirements.txt -torchada>=0.1.88 +torchada>=0.1.89 ``` ### Step 2: Conditional Import diff --git a/README_CN.md b/README_CN.md index 89aa345..2607378 100644 --- a/README_CN.md +++ b/README_CN.md @@ -67,9 +67,18 @@ torch.cuda.synchronize() | ctypes 库加载 | `ctypes.CDLL` 使用 CUDA 函数名 → 自动转换为 MUSA | | 统一加速器 API | `torch.accelerator.empty_cache()`、`memory_stats()`、`Stream`、`Event` 等 | | MUSA float64 原地对数 | `torch_musa < 2.11.0.post2` 时,`Tensor.log_()` 复用受支持的非原地操作,同时保持原地操作契约 | +| MUSA mm/bmm `out_dtype` | 在 `torch_musa 2.13.0` 以下,`torch.mm`/`torch.bmm` 的 `out_dtype=` 复用普通重载并以 fp32 累加,同时运行期探针报告该重载损坏。实测损坏于 `2.11.0.post1+musa5.2.0`(`mm` 全零、`bmm` 错值);**torch_musa 承诺在 `2.13.0` 修复,我们尚未验证** —— 包装层在该版本以下武装、自该版本起不安装,正确性由探针按进程裁决 | | Triton CUDA Extra | MUSA 上的 `tl.extra.cuda` → `tl.extra.musa` 兼容 | | Triton 融合 MoE | 面向 vLLM 和 SGLang 的 Triton 3.2.0 MTT S5000 调优配置 | +**未覆盖**:binding 本身不含 `*_Dtype` 重载的构建,`out_dtype=` 仍会报错。那是 binding 的契约、不是 torchada 要修的缺陷, +为该情形做 CUDA 等价支持不在范围内。 + +`out_dtype` 兼容层在 **`torch_musa 2.13.0` 以下**武装(该版本是 torch_musa 承诺修复重载的发布)。那是 vendor 的发布承诺、不是我们的实测, +信任它就是已接受的风险:若 `2.13.0` 实际没修好,`>= 2.13.0` 的栈不安装包装层,静默全零会重新出现。门控只决定是否在 `torch.mm`/`torch.bmm` +前加一层 Python 包装;正确性由运行期探针按进程裁决(健康则直接转发、损坏才仿真),因此若修复被 backport 到 `2.12.x` 会自动生效。 +`__version__` 无法解析时 rank 最低 ⇒ 仍然武装。待 `2.13.0` 发布并在此**验证修好**后删除该 shim;若承诺延期,则把上界抬到新的承诺版本。 + ## 示例 ### 混合精度训练 @@ -377,7 +386,7 @@ if torchada.is_gpu_device(device): # 在 CUDA 和 MUSA 上都能工作 ``` # pyproject.toml 或 requirements.txt -torchada>=0.1.88 +torchada>=0.1.89 ``` ### 步骤 2:条件导入 diff --git a/benchmarks/benchmark_history.json b/benchmarks/benchmark_history.json index 2ef3d9a..2228eb7 100644 --- a/benchmarks/benchmark_history.json +++ b/benchmarks/benchmark_history.json @@ -3,7 +3,7 @@ "description": "Historical benchmark results for torchada performance tracking", "results": [ { - "version": "0.1.88", + "version": "0.1.89", "date": "2026-01-29", "platform": "MUSA", "pytorch_version": "2.7.1", diff --git a/examples/extension_setup.py b/examples/extension_setup.py index 96a64ec..3b173af 100755 --- a/examples/extension_setup.py +++ b/examples/extension_setup.py @@ -69,7 +69,7 @@ def get_extensions(): if extensions: setup( name="my_cuda_extension", - version="0.1.0", + version="0.1.89", ext_modules=extensions, cmdclass={"build_ext": BuildExtension.with_options(use_ninja=True)}, python_requires=">=3.8", diff --git a/pyproject.toml b/pyproject.toml index 3673d2d..eb8cc70 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta" [project] name = "torchada" -version = "0.1.88" +version = "0.1.89" description = "Adapter package for torch_musa to act exactly like PyTorch CUDA" readme = "README.md" license = {text = "MIT"} diff --git a/src/torchada/__init__.py b/src/torchada/__init__.py index 61f2182..09c1c16 100644 --- a/src/torchada/__init__.py +++ b/src/torchada/__init__.py @@ -24,7 +24,7 @@ from torch.utils.cpp_extension import CUDAExtension, BuildExtension, CUDA_HOME """ -__version__ = "0.1.88" +__version__ = "0.1.89" from . import cuda, utils diff --git a/src/torchada/_cpp_ops.py b/src/torchada/_cpp_ops.py index e2edb2a..54b6921 100644 --- a/src/torchada/_cpp_ops.py +++ b/src/torchada/_cpp_ops.py @@ -95,10 +95,14 @@ def load_cpp_ops(force_reload: bool = False) -> Optional[object]: import torch - from ._patch import _is_pre_torch_musa_2_11_0_post2 + from ._version import version_of musa_module = getattr(torch, "musa", None) - if not _is_pre_torch_musa_2_11_0_post2(getattr(musa_module, "__version__", None)): + # Legacy shim (compiled only below 2.11.0.post2): those releases lack the + # ``multinomial`` / ``log`` / ``log_`` kernels these overrides provide. From + # that release on the native implementations are used, so the overrides are + # switched off through their env flags. + if version_of(musa_module) >= "2.11.0.post2": for op_name in ("multinomial", "log", "log_"): os.environ[f"TORCHADA_DISABLE_OP_OVERRIDE_{op_name}"] = "1" diff --git a/src/torchada/_patch.py b/src/torchada/_patch.py index 038232b..a919c4e 100644 --- a/src/torchada/_patch.py +++ b/src/torchada/_patch.py @@ -26,6 +26,7 @@ import logging import os import sys +import threading import time import warnings from types import ModuleType, SimpleNamespace @@ -35,12 +36,15 @@ from ._cpp_ops import get_module from ._platform import is_musa_platform +from ._version import version_of logger = logging.getLogger(__name__) _patched = False _original_init_process_group = None _original_tensor_log_ = None +_original_torch_mm = None +_original_torch_bmm = None # Registry for patch functions _patch_registry: List[Callable[[], None]] = [] @@ -125,7 +129,20 @@ def _patch_inductor_template_heuristics(): return musa_module = getattr(torch, "musa", None) - if not _is_pre_torch_musa_2_11_0_post2(getattr(musa_module, "__version__", None)): + # 2.11.0.post2 is the release where four legacy shims are no longer needed, so + # every one of them is installed only *below* it, and each gate below carries + # its own reason: + # * inductor MUSA template heuristics - the release registers its own, so the + # missing registry keys no longer have to be filled in from CUDA; + # * MUSA float64 in-place ``Tensor.log_`` - the release implements it; + # * the ``torch.accelerator`` memory overrides - the release dispatches those + # APIs to the MUSA allocator itself, so the official ones must be kept; + # * the compiled ``multinomial`` / ``log`` / ``log_`` op overrides - the + # release provides those kernels. + # Evidence: the 2.11.0.post2 work the shims were written against (upstream + # torch_musa PR #113) and this repository's compatibility table, which lists + # each shim with its ``< 2.11.0.post2`` condition. + if version_of(musa_module) >= "2.11.0.post2": return from torch._inductor.codegen.common import init_backend_registration @@ -169,7 +186,9 @@ def _patch_tensor_log_(): if not is_musa_platform() or _original_tensor_log_ is not None: return - if not _is_pre_torch_musa_2_11_0_post2(torch.musa.__version__): + # Legacy shim (installed below 2.11.0.post2): from that release on, MUSA has a + # working float64 in-place ``log_``, so the out-of-place reuse is not needed. + if version_of(torch.musa) >= "2.11.0.post2": return _original_tensor_log_ = torch.Tensor.log_ @@ -183,6 +202,401 @@ def patched_log_(self): torch.Tensor.log_ = patched_log_ +# --------------------------------------------------------------------------- +# MUSA out_dtype backport for torch.mm / torch.bmm +# --------------------------------------------------------------------------- +# +# On the tested stack (torch_musa 2.11.0.post1+musa5.2.0) the +# ``aten::mm.dtype`` / ``aten::bmm.dtype`` overloads are registered, but their +# MUSA implementations do not write their result: the returned tensor is +# correctly shaped and typed and holds all zeros (``mm``) or garbage (``bmm``), +# while the plain overloads and the argument validation are correct. +# The backport therefore reuses the plain overloads, which makes it free of any +# vendor-kernel assumption, and it is only installed while the probe below +# observes the defect, so it removes itself once the vendor fixes the ops. +# +# Measured on one stack only: torch_musa 2.11.0.post1+musa5.2.0 (``mm`` all +# zeros, ``bmm`` non-zero wrong values). torch_musa committed to fix the overloads +# in **2.13.0** - a vendor release commitment, *not* verified here - so the +# wrappers are armed on stacks below 2.13.0 and nothing is installed from 2.13.0 on. +# The gate only decides whether a Python wrapper sits in front of +# ``torch.mm``/``torch.bmm``; correctness is decided per process by the probe +# below, which forwards to a healthy overload and emulates a broken one. +# +# Accepted risk of trusting that commitment: if 2.13.0 does not actually fix the +# overloads, a >= 2.13.0 stack installs nothing and the silent all-zero result can +# come back. Mitigation is the compatibility table plus a post-release measurement. +# +# Lifecycle: once 2.13.0 is released and *verified* fixed here, delete this shim +# entirely (do not keep the bound); if the fix slips, raise the bound to the newly +# committed release. +# +_MM_OUT_DTYPE_PROBE_LOCK = threading.Lock() +#: Only a *deterministic* ``{"mm": bool, "bmm": bool}`` verdict is ever stored +#: here. ``None`` means "not probed (yet)" or "the probe could not decide". +_mm_out_dtype_probe_cache: Optional[dict] = None +#: The probe failure warning is emitted at most once per process. +_mm_out_dtype_probe_warned = False + + +def _musa_devices_available() -> Optional[bool]: + """Return whether a usable MUSA device exists: ``None`` when unknowable. + + Explicit answers only: a platform without ``torch.musa`` or with zero devices + is "nothing to patch", while a question that cannot be answered at all is + reported as unknown so the caller stays conservative. + """ + musa_module = getattr(torch, "musa", None) + if musa_module is None: + return False + counter = getattr(musa_module, "device_count", None) + if not callable(counter): + return None + try: + return int(counter()) > 0 + except Exception: # pragma: no cover - depends on the torch build + return None + + +def _out_dtype_overload_is_missing(exc: BaseException) -> bool: + """Return whether ``exc`` says the kwarg/overload itself does not exist. + + Only an explicit "no such keyword" / "no matching overload" answer counts. + That is a torch_musa build without the ``*_Dtype`` family, where forwarding + the keyword is harmless because there is nothing to patch. Every other + failure - a device error, a muDNN kernel failure, a dtype validation error - + says nothing about the overload's correctness and must never be read as + health: on this stack the broken overload is exactly the one that raises + inside the device code. + """ + text = str(exc).lower() + if isinstance(exc, TypeError): + return "out_dtype" in text or "unexpected keyword" in text + if isinstance(exc, RuntimeError): + return "out_dtype" in text and any( + marker in text + for marker in ( + "unexpected keyword", + "unknown keyword", + "invalid keyword", + "no matching function", + "no matching overload", + "overload not found", + "does not accept", + ) + ) + return False + + +def _warn_probe_failure(exc: BaseException) -> None: + """Warn (once per process) that the probe could not decide, and why.""" + global _mm_out_dtype_probe_warned + + if _mm_out_dtype_probe_warned: + return + _mm_out_dtype_probe_warned = True + logger.warning( + "MUSA mm/bmm out_dtype probe could not decide (%s: %s); keeping the fp32 " + "backport, not caching a verdict, and retrying on the next out_dtype call", + type(exc).__name__, + exc, + ) + + +def _probe_op_out_dtype_broken(op: Callable, batched: bool) -> Optional[bool]: + """Return whether ``op`` writes a wrong result for ``out_dtype=float32``. + + Exact-integer operands make the expected matrix exactly representable, so a + conforming implementation (fp32 accumulation of the input-dtype values) + reproduces the plain overload bitwise whatever the summation order, and no + tolerance can hide a broken result. + + The result is deliberately three-valued: + + * ``True`` - the overload exists and writes a wrong result: patch it; + * ``False`` - the overload is *shown* to be correct, or an explicit + precondition says there is nothing to patch (no MUSA platform, no usable + device, a build whose binding does not accept ``out_dtype`` at all); + * ``None`` - the probe could not decide, because a device call failed for + any other reason. That case keeps the backport and is never cached: a + failure inside the device code is what the broken overload does, so + treating it as "healthy" would silently restore the all-zero result. + """ + if not is_musa_platform(): + return False + available = _musa_devices_available() + if available is False: + return False + if available is None: + _warn_probe_failure(RuntimeError("torch.musa.device_count() is not answerable")) + return None + + try: + a = torch.arange(1, 33, dtype=torch.float32, device="musa").reshape(4, 8) + b = torch.arange(1, 41, dtype=torch.float32, device="musa").reshape(8, 5) + if batched: + a, b = a.unsqueeze(0), b.unsqueeze(0) + reference = op(a, b) + # The explicit overload-existence question, in the one pair every build + # with the ``*_Dtype`` family accepts: only this call may answer "the + # keyword does not exist", and every other failure is unknown. + op(a, b, out_dtype=torch.float32) + except Exception as exc: + if _out_dtype_overload_is_missing(exc): + logger.debug("MUSA %s has no out_dtype overload: %s", op, exc) + return False + _warn_probe_failure(exc) + return None + + try: + for dtype in (torch.float32, torch.bfloat16): + requested = op(a.to(dtype), b.to(dtype), out_dtype=torch.float32) + if not torch.equal(requested, reference): + return True + except Exception as exc: + _warn_probe_failure(exc) + return None + return False + + +def _probe_out_dtype_broken_ops() -> dict: + """Return ``{"mm": Optional[bool], "bmm": Optional[bool]}`` for the overloads. + + Thread-safe, and it caches **only** a pair of deterministic verdicts. When a + probe could not decide, this call reports ``None`` for that op and stores + nothing, so the next promoted call probes again instead of inheriting a + verdict nobody reached. Only calls that really pass ``out_dtype`` reach this + function, so a retry costs one probe per such call until it decides. + """ + global _mm_out_dtype_probe_cache + + cached = _mm_out_dtype_probe_cache + if cached is not None: + return cached + + with _MM_OUT_DTYPE_PROBE_LOCK: + cached = _mm_out_dtype_probe_cache + if cached is not None: + return cached + verdicts = { + "mm": _probe_op_out_dtype_broken( + _original_torch_mm if _original_torch_mm is not None else torch.mm, + batched=False, + ), + "bmm": _probe_op_out_dtype_broken( + _original_torch_bmm if _original_torch_bmm is not None else torch.bmm, + batched=True, + ), + } + if None in verdicts.values(): + return verdicts + _mm_out_dtype_probe_cache = verdicts + return _mm_out_dtype_probe_cache + + +def _call_musa_op(op: Callable, input, mat2, args, out, out_dtype, kwargs) -> Any: + """Call the vendor op, forwarding only what the caller actually passed.""" + if out_dtype is not None: + kwargs["out_dtype"] = out_dtype + if out is not None: + kwargs["out"] = out + return op(input, mat2, *args, **kwargs) + + +def _in_device_capture() -> bool: + """Return whether the active device stream is being captured into a graph. + + An exception, or a torch build with no capture API to ask, is reported as + "capturing". The wrapper must not run a synchronizing probe without a + trustworthy negative answer, and an unresolved verdict is already treated as + broken inside a capture; answering ``False`` here instead would do the + opposite and run the probe inside a capture - the failure mode that makes a + capture die *and* latch a wrong verdict (see the ticket's D2 evidence). + """ + for owner in (getattr(torch, "musa", None), getattr(torch, "cuda", None)): + checker = getattr(owner, "is_current_stream_capturing", None) + if callable(checker): + try: + return bool(checker()) + except Exception as exc: # pragma: no cover - depends on the build + logger.debug( + "MUSA capture check failed (%s: %s); assuming a capture", + type(exc).__name__, + exc, + ) + return True + return True + + +def _wrap_mm_out_dtype(original: Callable, op_name: str) -> Callable: + """Wrap ``torch.mm`` / ``torch.bmm`` so ``out_dtype`` matches CUDA. + + Both the keyword and the documented positional form of ``out_dtype`` are + backported, and the call is touched as little as possible: + + * no ``out_dtype`` at all is forwarded verbatim, without inspecting the + operands - the plain call is the hot path; + * ``out_dtype`` equal to both input dtypes (which includes fp32 in / fp32 + out) is the plain overload, so the keyword is *dropped* and the vendor's + dtype overload, which discards the result on this stack, is bypassed + entirely; + * ``out_dtype=torch.float32`` with fp16/bf16 inputs is the only pair that + needs emulation: the operands are promoted and multiplied in fp32, giving + an fp32 *accumulation* of the input-dtype values rather than a + compute-in-input-dtype-then-cast; + * every other combination - over-long argument lists, non-tensor operands, + mismatched input dtypes, invalid dtype pairs - goes to the vendor entry + point with the keyword intact, so its argument validation and error + messages stay authoritative. + + ``op_name`` is the probe key (``"mm"`` or ``"bmm"``). Only the fp16/bf16 + with ``out_dtype=torch.float32`` pair consults it, the verdict is cached, + and a vendor overload that honours ``out_dtype`` is delegated to forever. + """ + checked = False + broken = False + + def vendor_honours_out_dtype() -> bool: + """Resolve (once) whether the vendor overload can be trusted. + + The probe runs the op on the device and therefore synchronizes, so it is + never resolved inside a graph capture. There an unresolved verdict counts + as *untrusted* instead, which means the capture bakes the fp32 emulation + path into that graph: a graph is static, so it keeps using the emulation + for its whole lifetime. The verdict is resolved by the next *eager* call. + + Consequence on a stack whose vendor overload already honours + ``out_dtype``: if the first ``out_dtype`` call happens inside a capture, + every graph captured before the first eager call has to be **captured + again** to get the vendor path back. The values stay correct within + tolerance either way; only the fast path needs the re-capture. + + Only a deterministic verdict is remembered. A probe that could not decide + (``None``) leaves the wrapper unresolved: the backport stays in use for + this call and the next promoted call probes again, because "the device + raised while probing" is exactly what the broken overload does and must + never be remembered as "the vendor is healthy". + """ + nonlocal checked, broken + if not checked: + if _in_device_capture(): + return False + verdict = _probe_out_dtype_broken_ops().get(op_name) + if verdict is None: + return False + checked = True + broken = bool(verdict) + if broken: + logger.info( + "MUSA %s ignores out_dtype: backport emulation in use", + f"torch.{op_name}", + ) + else: + logger.info( + "MUSA %s honours out_dtype: vendor implementation kept", + f"torch.{op_name}", + ) + return not broken + + @functools.wraps(original) + def wrapped(input, mat2, *args, out=None, out_dtype=None, **kwargs): + if out_dtype is None and out is None and not args: + return original(input, mat2, **kwargs) + + if args: + # ``mm(input, mat2, out_dtype, *, out=None)`` is the documented + # positional form of the same keyword. Anything else stays with the + # vendor binding, which raises its own TypeError. + if len(args) > 1 or out_dtype is not None: + return _call_musa_op(original, input, mat2, args, out, out_dtype, kwargs) + out_dtype, args = args[0], () + + if out_dtype is None: + # Only ``out=`` (or extra positional arguments) is in play: forward + # the call unchanged. + return _call_musa_op(original, input, mat2, args, out, out_dtype, kwargs) + + if ( + not isinstance(input, torch.Tensor) + or not isinstance(mat2, torch.Tensor) + or input.dtype != mat2.dtype + ): + # Not a pair this API defines: the vendor binding decides. + return _call_musa_op(original, input, mat2, args, out, out_dtype, kwargs) + + if out_dtype == input.dtype: + # ``out_dtype`` equal to both input dtypes (fp32 in / fp32 out + # included) is the plain overload: bitwise identical, and no probe + # is needed to know that. + result = original(input, mat2) + elif out_dtype == torch.float32 and input.dtype in (torch.float16, torch.bfloat16): + # The only combination that depends on the installed vendor: use a + # vendor overload that honours ``out_dtype``, otherwise promote the + # operands and multiply in fp32. + if vendor_honours_out_dtype(): + return _call_musa_op(original, input, mat2, args, out, out_dtype, kwargs) + result = original(input.to(torch.float32), mat2.to(torch.float32)) + else: + # Invalid dtype pair: keep the vendor's own argument validation. + return _call_musa_op(original, input, mat2, args, out, out_dtype, kwargs) + + if out is None: + return result + if out.dtype != result.dtype: + return _call_musa_op(original, input, mat2, args, out, out_dtype, kwargs) + return out.copy_(result) + + return wrapped + + +@patch_function +@requires_import("torch_musa") +def _patch_mm_out_dtype(): + """Backport CUDA ``out_dtype=`` semantics for MUSA ``torch.mm``/``torch.bmm``. + + The wrappers are armed on MUSA stacks below the committed fix release, but the + defect probe + itself is deferred to the first call that actually passes ``out_dtype``. + Probing means running ``mm``/``bmm`` on the device, and that is observable + from the outside: it initialises the vendor libraries before the host + process gets to its own warm-up, which moves the memory-profiling peak of a + downstream vLLM server, changes the KV cache budget it derives from it + (746,446 vs 731,482 tokens on the Nemotron-3.5 MTP6 config) and shifted + measured decode TPOT by 7-9% even though every device kernel was identical. + Installing a compatibility backport must not perturb a process that never + asks for a promoted ``out_dtype``, so nothing touches the device until the + first promoted call needs the verdict. Plain calls, and calls whose dtype + pair is invalid or mismatched, still go straight to the vendor op. + """ + global _original_torch_mm, _original_torch_bmm + + if not is_musa_platform() or _original_torch_mm is not None: + return + musa_module = getattr(torch, "musa", None) + # Armed only below the release torch_musa committed to fix these overloads in + # (2.13.0 - a vendor commitment, not a measurement of ours; see the section note + # above and the README). Nothing at all is installed from 2.13.0 on. + # An unknown or unparsable version ranks lowest in ``version_of``, so it + # compares below 2.13.0 and keeps the wrappers armed - a version we cannot read + # is never assumed to be fixed. + if version_of(getattr(musa_module, "__version__", None)) >= "2.13.0": + return + + original_mm, original_bmm = torch.mm, torch.bmm + _original_torch_mm, _original_torch_bmm = original_mm, original_bmm + + wrapped_mm = _wrap_mm_out_dtype(original_mm, "mm") + _register_jit_builtin_alias(original_mm, wrapped_mm) + torch.mm = wrapped_mm + wrapped_bmm = _wrap_mm_out_dtype(original_bmm, "bmm") + _register_jit_builtin_alias(original_bmm, wrapped_bmm) + torch.bmm = wrapped_bmm + logger.info( + "MUSA out_dtype backport armed for %s (probe deferred to the first out_dtype call)", + ", ".join(f"torch.{name}" for name in ("mm", "bmm")), + ) + + # Cache for translated device strings - avoids repeated string operations _device_str_cache = {} @@ -1797,49 +2211,6 @@ def __getitem__(self, name: str): _original_ctypes_CDLL = None -_TORCH_MUSA_POST2_VERSION = "2.11.0.post2" - - -def _is_pre_torch_musa_2_11_0_post2(version) -> bool: - """Return whether the torch_musa version predates 2.11.0.post2. - - torch_musa 2.11.0.post2 fixes the unified accelerator memory APIs and the - float64 in-place ``Tensor.log_``. Older releases still need torchada to - force those memory calls through torch.musa and to backport the log path. - Ignore the local version suffix (for example ``+musa5.2.0``), because it - identifies the MUSA stack build rather than the torch_musa fix level. - - If a torch_musa build does not expose a version, retain the compatibility - overrides rather than risking the known runtime failure. - """ - if version is None: - return True - - public_version = str(version).split("+", 1)[0] - try: - # Use PyTorch's vendored PEP 440 parser so post releases compare - # semantically (post10 > post2) without adding a torchada dependency. - from torch._vendor.packaging.version import InvalidVersion, Version - except ImportError: - # If the parser is unavailable, keep the workaround enabled: disabling - # it could re-expose the failure this gate fixes. - logger.warning( - "Unable to parse torch_musa version %r; retaining compatibility patches", - version, - ) - return True - try: - return Version(public_version) < Version(_TORCH_MUSA_POST2_VERSION) - except InvalidVersion: - # An unknown or malformed version must keep the workaround enabled: - # disabling it could re-expose the failure this gate fixes. - logger.warning( - "Unable to parse torch_musa version %r; retaining compatibility patches", - version, - ) - return True - - class _AcceleratorModuleWrapper(ModuleType): """ Wrapper module that extends torch.accelerator with fallbacks to torch.musa. @@ -1916,11 +2287,11 @@ def __init__(self, original_accel, musa_module): self._musa_module = musa_module self._overrides = {} - # torch_musa versions before 2.11.0.post2 route these APIs through - # torch._C._accelerator_* without dispatching to the MUSA allocator. - # Newer versions provide working unified accelerator implementations, - # so preserve those instead of forcing the torch.musa compatibility path. - if _is_pre_torch_musa_2_11_0_post2(getattr(musa_module, "__version__", None)): + # Legacy shim (installed below 2.11.0.post2): those releases route these + # APIs through torch._C._accelerator_* without dispatching to the MUSA + # allocator. From that release on, the unified accelerator implementations + # work, so preserve them instead of forcing the torch.musa path. + if version_of(musa_module) < "2.11.0.post2": for name in self._MUSA_OVERRIDES: musa_name = self._REMAP_ATTRS.get(name, name) if hasattr(original_accel, name) and hasattr(musa_module, musa_name): @@ -2198,6 +2569,8 @@ def apply_patches(): - optional CUDA graph debug dumps via TORCHADA_CUDA_GRAPH_DEBUG_DUMP_PATH - torch.cuda.nccl -> torch.musa.mccl - torch.amp.autocast(device_type='cuda') -> 'musa' + - torch.mm / torch.bmm out_dtype= backport while the MUSA dtype overload + does not write its result - torch.utils.cpp_extension (CUDAExtension, BuildExtension) -> MUSA versions - CUDA_VISIBLE_DEVICES -> MUSA_VISIBLE_DEVICES environment fallback - torch._inductor.autotune_process.CUDA_VISIBLE_DEVICES -> MUSA_VISIBLE_DEVICES diff --git a/src/torchada/_version.py b/src/torchada/_version.py new file mode 100644 index 0000000..74bbde6 --- /dev/null +++ b/src/torchada/_version.py @@ -0,0 +1,147 @@ +"""PEP 440 version comparisons with infix operators. + +torchada gates a number of patches on the installed ``torch_musa`` version. +Spelling those gates as dedicated predicates (``_version_lt``/``_is_pre_*``) +does not scale: every new gate needs a new helper, the bound is buried in a +function body, and the call site cannot say what it is comparing against. This +module provides a small proxy instead, so a gate reads like the relationship it +encodes:: + + from ._version import version_of + + if version_of(getattr(torch.musa, "__version__", None)) < "2.11.0.post2": + ... # workaround for everything older than the fix + +``version_of`` accepts a version string, ``None``, an object exposing +``__version__`` (a module, for example), or another proxy, and the proxy +supports ``<``, ``<=``, ``>``, ``>=``, ``==`` and ``!=`` against any of those, +in either operand order. + +Two policies are implemented here so that they do not have to be repeated at +every call site: + +* **The local version segment is ignored.** ``2.11.0.post1+musa5.2.0`` and + ``2.11.0.post1`` compare equal, because the ``+musa*`` suffix identifies the + MUSA stack build rather than the level of the fix being gated. Comparisons + therefore use the *public* version. +* **An unknown or unparsable version compares as the lowest possible version** + (equivalent to ``0``). A gate written as an upper bound - the shape used by + torchada's existing gates - is then ``True`` for an unknown version, which is + exactly the "keep the workaround when the version cannot be trusted" + behaviour those gates document. A gate that must *skip* work for an old + release has to say so explicitly, hence ``is_known``. + +Gates spell their bound inline (``version_of(module) >= "2.11.0.post2"``), and this +proxy stays the comparator for that shape because it is parse-failure safe: a +malformed ``__version__`` must not raise (``packaging.version.parse`` would abort +at import time and take the patch with it) and must rank below every bound, so the +workaround stays armed instead of silently disappearing. +""" + +from __future__ import annotations + +from typing import Any + +__all__ = ["VersionComparison", "version_of"] + +_LOWEST = "0" + + +def _public_version(raw: Any): + """Import PyTorch's vendored PEP 440 parser and parse ``raw``. + + Returns ``None`` when the version is missing, malformed, or when the parser + itself is unavailable (a torch build without the vendored copy). + """ + if raw is None: + return None + public = str(raw).split("+", 1)[0].strip() + if not public: + return None + try: + from torch._vendor.packaging.version import InvalidVersion, Version + except ImportError: # pragma: no cover - depends on the torch build + return None + try: + return Version(public) + except InvalidVersion: + return None + + +def _raw_version(source: Any) -> Any: + """Return the version string carried by ``source``, if any.""" + if isinstance(source, VersionComparison): + return source.raw + if source is None or isinstance(source, str): + return source + return getattr(source, "__version__", None) + + +class VersionComparison: + """A comparable view of a package version, driven by infix operators.""" + + __slots__ = ("_raw", "_parsed") + + def __init__(self, source: Any): + self._raw = _raw_version(source) + self._parsed = _public_version(self._raw) + + @property + def raw(self) -> Any: + """The version string this proxy was built from (``None`` if unknown).""" + return self._raw + + @property + def is_known(self) -> bool: + """Whether the version parsed into a comparable PEP 440 version.""" + return self._parsed is not None + + def _rank(self) -> Any: + """The parsed version, or the lowest version for anything unknown.""" + return self._parsed if self._parsed is not None else _public_version(_LOWEST) + + def _compare(self, other: Any) -> int: + mine = self._rank() + theirs = _public_version(_raw_version(other)) + if theirs is None: + theirs = _public_version(_LOWEST) + if mine == theirs: + return 0 + return -1 if mine < theirs else 1 + + def __lt__(self, other: Any) -> bool: + return self._compare(other) < 0 + + def __le__(self, other: Any) -> bool: + return self._compare(other) <= 0 + + def __gt__(self, other: Any) -> bool: + return self._compare(other) > 0 + + def __ge__(self, other: Any) -> bool: + return self._compare(other) >= 0 + + def __eq__(self, other: Any) -> bool: + return self._compare(other) == 0 + + def __ne__(self, other: Any) -> bool: + return self._compare(other) != 0 + + def __hash__(self) -> int: + return hash(self._rank()) + + def __repr__(self) -> str: + shown = repr(self._raw) if self.is_known else f"unknown ({self._raw!r})" + return f"VersionComparison({shown})" + + +def version_of(source: Any) -> VersionComparison: + """Return a comparable version for ``source``. + + ``source`` may be a version string, ``None``, an object with a + ``__version__`` attribute (for example ``torch.musa`` or a module), or an + existing :class:`VersionComparison`. + """ + if isinstance(source, VersionComparison): + return source + return VersionComparison(source) diff --git a/src/torchada/csrc/ops.h b/src/torchada/csrc/ops.h index 28788d0..9b2669d 100644 --- a/src/torchada/csrc/ops.h +++ b/src/torchada/csrc/ops.h @@ -40,7 +40,7 @@ namespace torchada { // Version information -constexpr const char* VERSION = "0.1.0"; +constexpr const char* VERSION = "0.1.89"; // Check if operator override is enabled via environment variable inline bool is_override_enabled(const char* op_name) { diff --git a/tests/test_cuda_patching.py b/tests/test_cuda_patching.py index f3a8324..efbfd2f 100644 --- a/tests/test_cuda_patching.py +++ b/tests/test_cuda_patching.py @@ -3158,9 +3158,12 @@ class TestAcceleratorModuleWrapper: ), ) def test_torch_musa_version_boundary(self, musa_version, expected): - from torchada._patch import _is_pre_torch_musa_2_11_0_post2 + from torchada._version import version_of - assert _is_pre_torch_musa_2_11_0_post2(musa_version) is expected + # Same table as before the version proxy landed: the legacy shims are + # installed below this release, and an unknown or unparsable version + # ranks below the bound, so the workaround stays on. + assert (version_of(musa_version) < "2.11.0.post2") is expected def _make_wrapper( self, @@ -3403,11 +3406,12 @@ def test_empty_cache_uses_version_appropriate_implementation(self): if not torchada.is_musa_platform(): pytest.skip("Only applicable on MUSA platform") - from torchada._patch import _is_pre_torch_musa_2_11_0_post2 + from torchada._version import version_of - # Must not raise on either side of the torch_musa post2 boundary. + # Must not raise on either side of 2.11.0.post2, where the accelerator + # shim is installed below it and skipped from it on. torch.accelerator.empty_cache() - if _is_pre_torch_musa_2_11_0_post2(getattr(torch.musa, "__version__", None)): + if version_of(torch.musa) < "2.11.0.post2": assert torch.accelerator.empty_cache.__module__.startswith("torch_musa") else: assert torch.accelerator.empty_cache is torch.accelerator._original_accel.empty_cache diff --git a/tests/test_mm_out_dtype.py b/tests/test_mm_out_dtype.py new file mode 100644 index 0000000..d884921 --- /dev/null +++ b/tests/test_mm_out_dtype.py @@ -0,0 +1,702 @@ +"""Tests for the MUSA ``out_dtype`` backport of ``torch.mm`` / ``torch.bmm``. + +torch_musa accepts ``out_dtype`` on both ops but its implementation never +writes the result (all zeros for ``mm``, non-zero garbage for ``bmm``), while +the plain overloads and the argument validation are correct. torchada arms a +wrapper on the affected stack only, and the defect probe that decides whether +the vendor overload is broken runs on the first call that actually passes +``out_dtype`` - probing touches the device and was measured to perturb an +unrelated serving process. The wrapper tests below need no GPU (a recording op +stands in for the vendor entry point), the contract tests need MUSA hardware. +""" + +import logging + +import pytest +import torch + +from torchada import _patch + +ATOL = 1e-2 +RTOL = 1.6e-2 + + +def _musa_available() -> bool: + musa = getattr(torch, "musa", None) + return musa is not None and musa.is_available() + + +def _require_musa() -> None: + if not _musa_available(): + pytest.skip("MUSA device required") + + +def _recording_op(calls, result=None): + """Stand-in for the vendor op that records how the wrapper called it.""" + + def op(input, mat2, *args, **kwargs): + calls.append((input, mat2, args, kwargs)) + return result + + return op + + +def _computing_op(calls): + """Stand-in for the vendor op that records the call and computes the product.""" + + def op(input, mat2, *args, **kwargs): + calls.append((input, mat2, args, kwargs)) + return input @ mat2 + + return op + + +class TestMMOutDtypeGating: + """The backport must stay a no-op unless the probe reports a broken op.""" + + @staticmethod + def _install_patch(monkeypatch, broken, version="2.11.0.post1+musa5.2.0"): + """Run ``_patch_mm_out_dtype`` as if it were applied on MUSA.""" + import sys + from types import ModuleType, SimpleNamespace + + monkeypatch.setitem(sys.modules, "torch_musa", ModuleType("torch_musa")) + monkeypatch.setattr(_patch, "is_musa_platform", lambda: True) + monkeypatch.setattr(_patch, "_original_torch_mm", None) + monkeypatch.setattr(_patch, "_original_torch_bmm", None) + monkeypatch.setattr(_patch, "_mm_out_dtype_probe_cache", None) + monkeypatch.setattr(_patch, "_mm_out_dtype_probe_warned", False) + # These tests are about the eager path, so the capture state is pinned + # instead of being read from a host whose checker may raise. + monkeypatch.setattr(_patch, "_in_device_capture", lambda: False) + monkeypatch.setattr(torch, "musa", SimpleNamespace(__version__=version), raising=False) + monkeypatch.setattr(_patch, "_probe_out_dtype_broken_ops", lambda: dict(broken)) + + def test_arming_does_not_touch_the_device(self, monkeypatch): + """The probe must not run while the patch is applied: it perturbs serving. + + Probing means running ``mm``/``bmm`` on the device, which initialises the + vendor libraries ahead of the host process' own warm-up and changed the + KV cache budget and decode TPOT of an unrelated vLLM server. + """ + probes = [] + self._install_patch(monkeypatch, {"mm": True, "bmm": True}) + monkeypatch.setattr( + _patch, "_probe_out_dtype_broken_ops", lambda: probes.append(1) or {"mm": True, "bmm": True} + ) + + _patch._patch_mm_out_dtype() + + assert probes == [] + assert _patch._mm_out_dtype_probe_cache is None + + def test_healthy_vendor_op_is_delegated_to(self, monkeypatch): + """A healthy vendor op must not be reimplemented (no GPU needed).""" + calls = [] + vendor_mm = _recording_op(calls, result="vendor") + monkeypatch.setattr(torch, "mm", vendor_mm) + self._install_patch(monkeypatch, {"mm": False, "bmm": False}) + + _patch._patch_mm_out_dtype() + result = torch.mm(torch.randn(2, 3, dtype=torch.bfloat16), torch.randn(3, 2, dtype=torch.bfloat16), out_dtype=torch.float32) + + assert result == "vendor" + assert len(calls) == 1 and calls[0][3] == {"out_dtype": torch.float32} + + def test_plain_calls_never_probe(self, monkeypatch): + """Only a call that asks for ``out_dtype`` may trigger the probe.""" + probes = [] + calls = [] + monkeypatch.setattr(torch, "mm", _recording_op(calls, result="vendor")) + self._install_patch(monkeypatch, {"mm": True, "bmm": True}) + monkeypatch.setattr( + _patch, "_probe_out_dtype_broken_ops", lambda: probes.append(1) or {"mm": True, "bmm": True} + ) + + _patch._patch_mm_out_dtype() + a = torch.randn(2, 3, dtype=torch.bfloat16) + b = torch.randn(3, 2, dtype=torch.bfloat16) + torch.mm(a, b) + torch.mm(a, b, out=torch.empty(2, 2, dtype=torch.bfloat16)) + + assert probes == [] + assert len(calls) == 2 + + def test_same_dtype_calls_never_probe(self, monkeypatch): + """The two same-dtype fast paths are decided without touching the GPU.""" + probes = [] + calls = [] + monkeypatch.setattr(torch, "mm", _recording_op(calls, result="vendor")) + self._install_patch(monkeypatch, {"mm": True, "bmm": True}) + monkeypatch.setattr( + _patch, "_probe_out_dtype_broken_ops", lambda: probes.append(1) or {"mm": True, "bmm": True} + ) + _patch._patch_mm_out_dtype() + a = torch.randn(2, 3, dtype=torch.bfloat16) + b = torch.randn(3, 2, dtype=torch.bfloat16) + torch.mm(a, b, out_dtype=torch.bfloat16) + torch.mm(a.float(), b.float(), out_dtype=torch.float32) + + assert probes == [] + assert len(calls) == 2 + assert all("out_dtype" not in call[3] for call in calls) + + def test_probe_runs_once_and_is_cached(self, monkeypatch): + probes = [] + monkeypatch.setattr(torch, "mm", _recording_op([], result="vendor")) + self._install_patch(monkeypatch, {"mm": True, "bmm": True}) + monkeypatch.setattr( + _patch, "_probe_out_dtype_broken_ops", lambda: probes.append(1) or {"mm": True, "bmm": True} + ) + + _patch._patch_mm_out_dtype() + a = torch.randn(2, 3, dtype=torch.bfloat16) + b = torch.randn(3, 2, dtype=torch.bfloat16) + for _ in range(3): + torch.mm(a, b, out_dtype=torch.float32) + + assert len(probes) == 1 + + def test_a_failing_probe_keeps_the_backport_and_is_not_cached( + self, monkeypatch, caplog + ): + """A probe that cannot decide must never be remembered as "healthy". + + The defect this locks: a probe exception used to be read as "nothing to + patch" and cached for the whole process, so a device error while probing + would silently disable the backport and let ``out_dtype=float32`` return + the vendor's zeros again. Only explicit conditions (no MUSA platform, no + usable device, a binding that does not accept the keyword) may report + health; every other failure keeps the backport and stays uncached. + """ + calls = [] + real_probe = _patch._probe_out_dtype_broken_ops + monkeypatch.setattr(torch, "mm", _computing_op(calls)) + self._install_patch(monkeypatch, {"mm": True, "bmm": True}) + # Use the real probe, not the helper's stub: the whole point is what the + # probe does with a device error. + monkeypatch.setattr(_patch, "_probe_out_dtype_broken_ops", real_probe) + monkeypatch.setattr(_patch, "_musa_devices_available", lambda: True) + _patch._patch_mm_out_dtype() + + def exploding_op(*args, **kwargs): + # The way the broken vendor overload fails: a device/kernel error, not + # a missing keyword. + raise RuntimeError("Run MUDNN failed in: mudnnMatmulGetWorkspaceSize") + + monkeypatch.setattr(_patch, "_original_torch_mm", exploding_op) + monkeypatch.setattr(_patch, "_original_torch_bmm", exploding_op) + + a = torch.randn(2, 3, dtype=torch.bfloat16) + b = torch.randn(3, 2, dtype=torch.bfloat16) + expected = a.to(torch.float32) @ b.to(torch.float32) + with caplog.at_level(logging.WARNING, logger=_patch.logger.name): + results = [torch.mm(a, b, out_dtype=torch.float32) for _ in range(2)] + + for result in results: + assert torch.allclose(result, expected, atol=1e-5, rtol=1e-5) + assert [call[3] for call in calls] == [{}, {}], "the vendor kwarg must not be forwarded" + assert all(call[0].dtype == torch.float32 for call in calls), "the promoted path is used" + assert _patch._mm_out_dtype_probe_cache is None, ( + "an undecidable probe must not be stored as a verdict" + ) + assert caplog.text.count("probe could not decide") == 1, "the warning is throttled" + + def test_a_failing_capture_check_never_probes(self, monkeypatch): + """A capture check that raises answers "capturing", so nothing is probed. + + Locks the direction of the default: an exception here used to answer "not + capturing" and ran the probe *inside* a capture - the failure mode that + both kills the capture and latches a wrong "healthy" verdict. The wrapper + must keep using the emulation and resolve later instead. + """ + probes = [] + calls = [] + real_capture_check = _patch._in_device_capture + monkeypatch.setattr(torch, "mm", _computing_op(calls)) + self._install_patch(monkeypatch, {"mm": True, "bmm": True}) + # Undo the helper's pin so the real capture check (with the exploding + # checker installed below) is what the wrapper consults. + monkeypatch.setattr(_patch, "_in_device_capture", real_capture_check) + monkeypatch.setattr( + _patch, + "_probe_out_dtype_broken_ops", + lambda: probes.append(1) or {"mm": True, "bmm": True}, + ) + + def exploding_checker(): + raise RuntimeError("no device to query") + + for owner in (getattr(torch, "musa", None), torch.cuda): + if owner is not None: + monkeypatch.setattr( + owner, "is_current_stream_capturing", exploding_checker, raising=False + ) + + assert _patch._in_device_capture() is True + + _patch._patch_mm_out_dtype() + a = torch.randn(2, 3, dtype=torch.bfloat16) + b = torch.randn(3, 2, dtype=torch.bfloat16) + result = torch.mm(a, b, out_dtype=torch.float32) + + assert probes == [], "no probe may run when the capture state is unknown" + assert [call[3] for call in calls] == [{}], "the vendor kwarg must not be forwarded" + assert all(call[0].dtype == torch.float32 for call in calls) + assert torch.allclose(result, a.to(torch.float32) @ b.to(torch.float32), atol=1e-5, rtol=1e-5) + + def test_healthy_bmm_probe_is_resolved_on_its_own_op(self, monkeypatch): + """A broken ``mm`` must not make a healthy ``bmm`` take the backport.""" + calls = [] + monkeypatch.setattr(torch, "bmm", _recording_op(calls, result="vendor")) + self._install_patch(monkeypatch, {"mm": True, "bmm": False}) + + _patch._patch_mm_out_dtype() + a = torch.randn(1, 2, 3, dtype=torch.bfloat16) + b = torch.randn(1, 3, 2, dtype=torch.bfloat16) + result = torch.bmm(a, b, out_dtype=torch.float32) + + assert result == "vendor" + assert len(calls) == 1 + + def test_unknown_torch_musa_version_keeps_the_patch_enabled(self, monkeypatch): + """An unparsable torch_musa version must not disable the backport.""" + original_mm = _recording_op([]) + monkeypatch.setattr(torch, "mm", original_mm) + self._install_patch(monkeypatch, {"mm": True, "bmm": True}, version="not-a-version") + + _patch._patch_mm_out_dtype() + + assert torch.mm is not original_mm + + def test_below_the_committed_fix_release_is_armed(self, monkeypatch): + """Everything below 2.13.0 is armed, including releases older than 2.11.0.""" + original_mm = _recording_op([]) + monkeypatch.setattr(torch, "mm", original_mm) + self._install_patch(monkeypatch, {"mm": True, "bmm": True}, version="2.7.1+musa4.3.0") + + _patch._patch_mm_out_dtype() + + assert torch.mm is not original_mm + + +class TestMMOutDtypeWrapper: + """Wrapper semantics on CPU: a recording op stands in for the MUSA op.""" + + @pytest.fixture(autouse=True) + def _broken_probe(self, monkeypatch): + """Wrapper tests exercise the broken stack, i.e. the backport path.""" + monkeypatch.setattr(_patch, "_mm_out_dtype_probe_cache", {"mm": True, "bmm": True}) + + @staticmethod + def _wrap(calls, result=None, op_name="mm"): + return _patch._wrap_mm_out_dtype(_recording_op(calls, result), op_name) + + def test_plain_call_is_forwarded_unchanged(self): + """No ``out_dtype``: the call reaches the original untouched.""" + calls = [] + wrapped = self._wrap(calls) + a = torch.randn(2, 3, dtype=torch.bfloat16) + b = torch.randn(3, 2, dtype=torch.bfloat16) + + wrapped(a, b) + + assert calls == [(a, b, (), {})] + + def test_out_dtype_float32_promotes_both_operands(self): + calls = [] + wrapped = self._wrap(calls) + a = torch.randn(2, 3, dtype=torch.bfloat16) + b = torch.randn(3, 2, dtype=torch.bfloat16) + + wrapped(a, b, out_dtype=torch.float32) + + called_a, called_b, args, kwargs = calls[0] + assert called_a.dtype == torch.float32 + assert called_b.dtype == torch.float32 + assert args == () + assert "out_dtype" not in kwargs + + def test_fp16_inputs_are_promoted_too(self): + calls = [] + wrapped = self._wrap(calls) + a = torch.randn(2, 3, dtype=torch.float16) + b = torch.randn(3, 2, dtype=torch.float16) + + wrapped(a, b, out_dtype=torch.float32) + + called_a, called_b, _, kwargs = calls[0] + assert called_a.dtype == torch.float32 + assert called_b.dtype == torch.float32 + assert kwargs == {} + + def test_same_dtype_out_dtype_drops_the_keyword(self): + """``out_dtype == input dtype`` is the plain op: the spy must see no kwarg.""" + calls = [] + wrapped = self._wrap(calls) + a = torch.randn(2, 3, dtype=torch.bfloat16) + b = torch.randn(3, 2, dtype=torch.bfloat16) + + wrapped(a, b, out_dtype=torch.bfloat16) + + assert len(calls) == 1 + called_a, called_b, args, kwargs = calls[0] + assert called_a is a + assert called_b is b + assert args == () + assert kwargs == {} + assert "out_dtype" not in kwargs + + def test_fp32_out_dtype_on_fp32_inputs_drops_the_keyword(self): + """fp32 in / fp32 out is the same no-op fast path as same-dtype.""" + calls = [] + wrapped = self._wrap(calls) + a = torch.randn(2, 3, dtype=torch.float32) + b = torch.randn(3, 2, dtype=torch.float32) + + wrapped(a, b, out_dtype=torch.float32) + + assert len(calls) == 1 + called_a, called_b, args, kwargs = calls[0] + assert called_a is a + assert called_b is b + assert args == () + assert kwargs == {} + assert "out_dtype" not in kwargs + + def test_the_probe_is_not_resolved_inside_a_capture(self, monkeypatch): + """No probe and no device synchronization while a graph is capturing.""" + calls = [] + probes = [] + + def _probe(): + probes.append(True) + return {"mm": True, "bmm": True} + + monkeypatch.setattr(_patch, "_mm_out_dtype_probe_cache", None) + monkeypatch.setattr(_patch, "_probe_out_dtype_broken_ops", _probe) + monkeypatch.setattr(_patch, "_in_device_capture", lambda: True) + wrapped = self._wrap(calls) + a = torch.randn(2, 3, dtype=torch.bfloat16) + b = torch.randn(3, 2, dtype=torch.bfloat16) + + wrapped(a, b, out_dtype=torch.float32) + + assert probes == [] + called_a, called_b, _, kwargs = calls[0] + assert called_a.dtype == torch.float32 + assert called_b.dtype == torch.float32 + assert "out_dtype" not in kwargs + + def test_illegal_dtype_pair_keeps_the_vendor_validation(self): + calls = [] + wrapped = self._wrap(calls) + a = torch.randn(2, 3, dtype=torch.bfloat16) + b = torch.randn(3, 2, dtype=torch.bfloat16) + + wrapped(a, b, out_dtype=torch.float16) + + called_a, called_b, args, kwargs = calls[0] + assert called_a is a + assert called_b is b + assert args == () + assert kwargs == {"out_dtype": torch.float16} + + def test_mismatched_input_dtypes_keep_the_vendor_validation(self): + calls = [] + wrapped = self._wrap(calls) + a = torch.randn(2, 3, dtype=torch.bfloat16) + b = torch.randn(3, 2, dtype=torch.float32) + + wrapped(a, b, out_dtype=torch.float32) + + called_a, called_b, _, kwargs = calls[0] + assert called_a is a + assert called_b is b + assert kwargs == {"out_dtype": torch.float32} + + def test_positional_out_dtype_is_backported(self): + calls = [] + wrapped = self._wrap(calls) + a = torch.randn(2, 3, dtype=torch.bfloat16) + b = torch.randn(3, 2, dtype=torch.bfloat16) + + wrapped(a, b, torch.float32) + + called_a, called_b, args, kwargs = calls[0] + assert called_a.dtype == torch.float32 + assert called_b.dtype == torch.float32 + assert args == () + assert kwargs == {} + + def test_over_long_argument_lists_are_forwarded(self): + calls = [] + wrapped = self._wrap(calls) + a = torch.randn(2, 3, dtype=torch.bfloat16) + b = torch.randn(3, 2, dtype=torch.bfloat16) + + wrapped(a, b, torch.float32, torch.float32) + + called_a, called_b, args, kwargs = calls[0] + assert called_a is a + assert called_b is b + assert args == (torch.float32, torch.float32) + assert kwargs == {} + + def test_positional_and_keyword_out_dtype_are_forwarded(self): + calls = [] + wrapped = self._wrap(calls) + a = torch.randn(2, 3, dtype=torch.bfloat16) + b = torch.randn(3, 2, dtype=torch.bfloat16) + + wrapped(a, b, torch.float32, out_dtype=torch.float32) + + called_a, called_b, args, kwargs = calls[0] + assert called_a is a + assert called_b is b + assert args == (torch.float32,) + assert kwargs == {"out_dtype": torch.float32} + + def test_non_tensor_operands_are_forwarded(self): + calls = [] + wrapped = self._wrap(calls) + + wrapped([[1.0, 2.0]], [[3.0], [4.0]], out_dtype=torch.float32) + + assert calls[0][0] == [[1.0, 2.0]] + assert calls[0][3] == {"out_dtype": torch.float32} + + def test_out_tensor_receives_the_result(self): + result = torch.arange(4, dtype=torch.float32).reshape(2, 2) + calls = [] + wrapped = self._wrap(calls, result=result) + a = torch.randn(2, 3, dtype=torch.bfloat16) + b = torch.randn(3, 2, dtype=torch.bfloat16) + out = torch.empty(2, 2, dtype=torch.float32) + + returned = wrapped(a, b, out=out, out_dtype=torch.float32) + + assert returned is out + assert torch.equal(out, result) + + def test_out_tensor_of_the_wrong_dtype_keeps_the_vendor_validation(self): + calls = [] + wrapped = self._wrap(calls, result=torch.zeros(2, 2, dtype=torch.float32)) + a = torch.randn(2, 3, dtype=torch.bfloat16) + b = torch.randn(3, 2, dtype=torch.bfloat16) + out = torch.empty(2, 2, dtype=torch.bfloat16) + + wrapped(a, b, out=out, out_dtype=torch.float32) + + # The fp32 promotion runs first, then the dtype mismatch is handed back + # to the vendor op with both keywords intact. + assert calls[0][0].dtype == torch.float32 + assert calls[-1][3] == {"out_dtype": torch.float32, "out": out} + + def test_plain_call_with_out_tensor_is_forwarded(self): + calls = [] + wrapped = self._wrap(calls) + a = torch.randn(2, 3, dtype=torch.float32) + b = torch.randn(3, 2, dtype=torch.float32) + out = torch.empty(2, 2, dtype=torch.float32) + + wrapped(a, b, out=out) + + assert calls[0][3] == {"out": out} + + +@pytest.mark.musa +class TestMMOutDtypeContract: + """Hardware contract: what CUDA promises must hold on MUSA.""" + + @staticmethod + def _pair(shape, dtype): + a = torch.randn(*shape, device="musa", dtype=dtype) + b = torch.randn(shape[-1], 5, device="musa", dtype=dtype) + return a, b + + @pytest.mark.parametrize("dtype", [torch.float32, torch.float16, torch.bfloat16]) + def test_mm_out_dtype_float32_matches_fp32_reference(self, dtype): + _require_musa() + a, b = self._pair((7, 256), dtype) + reference = torch.mm(a.to(torch.float32), b.to(torch.float32)) + + result = torch.mm(a, b, out_dtype=torch.float32) + + assert result.dtype == torch.float32 + assert result.shape == reference.shape + assert result.abs().max().item() > 0.0, "out_dtype result is all zeros" + assert torch.allclose(result, reference, atol=ATOL, rtol=RTOL) + + @pytest.mark.parametrize("dtype", [torch.float32, torch.float16, torch.bfloat16]) + def test_bmm_out_dtype_float32_matches_fp32_reference(self, dtype): + _require_musa() + a, b = self._pair((7, 256), dtype) + a, b = a.unsqueeze(0), b.unsqueeze(0) + reference = torch.bmm(a.to(torch.float32), b.to(torch.float32)) + + result = torch.bmm(a, b, out_dtype=torch.float32) + + assert result.dtype == torch.float32 + assert result.abs().max().item() > 0.0, "out_dtype result is all zeros" + assert torch.allclose(result, reference, atol=ATOL, rtol=RTOL) + + def test_positional_out_dtype_is_backported(self): + _require_musa() + a, b = self._pair((7, 256), torch.bfloat16) + reference = torch.mm(a.to(torch.float32), b.to(torch.float32)) + + result = torch.mm(a, b, torch.float32) + + assert result.dtype == torch.float32 + assert torch.allclose(result, reference, atol=ATOL, rtol=RTOL) + + a, b = a.unsqueeze(0), b.unsqueeze(0) + assert torch.allclose( + torch.bmm(a, b, torch.float32), + torch.bmm(a.to(torch.float32), b.to(torch.float32)), + atol=ATOL, + rtol=RTOL, + ) + + @pytest.mark.parametrize("dtype", [torch.float32, torch.float16, torch.bfloat16]) + def test_same_dtype_out_dtype_matches_the_plain_call(self, dtype): + _require_musa() + a, b = self._pair((7, 256), dtype) + + assert torch.equal(torch.mm(a, b, out_dtype=dtype), torch.mm(a, b)) + + a, b = a.unsqueeze(0), b.unsqueeze(0) + assert torch.equal(torch.bmm(a, b, out_dtype=dtype), torch.bmm(a, b)) + + def test_fp32_out_dtype_is_bitwise_equal_to_the_plain_call(self): + _require_musa() + a, b = self._pair((7, 256), torch.float32) + + assert torch.equal(torch.mm(a, b, out_dtype=torch.float32), torch.mm(a, b)) + + @pytest.mark.parametrize( + "in_dtype,out_dtype", + [ + (torch.bfloat16, torch.float16), + (torch.float32, torch.float16), + (torch.float32, torch.bfloat16), + ], + ) + def test_illegal_out_dtype_still_raises(self, in_dtype, out_dtype): + _require_musa() + a, b = self._pair((7, 256), in_dtype) + + with pytest.raises(RuntimeError, match="out_dtype must be the same as input dtype"): + torch.mm(a, b, out_dtype=out_dtype) + + # bmm validates one stage later than mm on MUSA (the rejection comes from + # the kernel rather than from the binding), so assert only that the patch + # keeps the vendor's own error instead of substituting one. + a, b = a.unsqueeze(0), b.unsqueeze(0) + vendor_bmm = _patch._original_torch_bmm or torch.bmm + with pytest.raises(RuntimeError) as vendor_error: + vendor_bmm(a, b, out_dtype=out_dtype) + with pytest.raises(RuntimeError) as patched_error: + torch.bmm(a, b, out_dtype=out_dtype) + assert str(patched_error.value) == str(vendor_error.value) + + def test_fp16_out_dtype_on_fp32_inputs_still_raises(self): + _require_musa() + a, b = self._pair((7, 256), torch.float32) + + with pytest.raises(RuntimeError, match="out_dtype must be the same as input dtype"): + torch.mm(a, b, out_dtype=torch.float16) + + def test_capture_before_the_first_eager_call_uses_the_emulation_path( + self, monkeypatch, caplog + ): + """A capture must not resolve the verdict, so it records the emulation path. + + If the first ``out_dtype`` call happens inside a MUSA graph capture, the + wrapper treats the unresolved verdict as untrusted: the capture contains + the fp32 emulation, no probe runs, and nothing synchronizes (a graph is + static, so that graph keeps the emulation until it is re-captured). One + eager call afterwards resolves the real verdict and logs it, which is what + restores the vendor path for later captures. + + This is a necessity, not an optimization. Measured on MUSA with the guard + disabled and a fresh wrapper whose first ``out_dtype`` call happens inside + a capture: the capture dies with ``RuntimeError: MUSA error: operation + failed due to a previous error during capture`` (the probe runs the vendor + overloads and a device readback on a stream that must not block), and - + worse than the failure - the probe swallows it and caches + ``{'mm': False, 'bmm': False}``, i.e. "the vendor overload is healthy", so + every later call would be sent to the broken overload and return zeros. + The guard removes that failure mode: the capture records the emulation and + the verdict is resolved later, from eager code. + + A fresh wrapper is built here because the verdict is remembered per + wrapper instance: the module-level ``torch.mm`` may already be resolved by + an earlier test in the session. + """ + probes = [] + original = ( + _patch._original_torch_mm if _patch._original_torch_mm is not None else torch.mm + ) + monkeypatch.setattr(_patch, "_mm_out_dtype_probe_cache", None) + monkeypatch.setattr( + _patch, + "_probe_out_dtype_broken_ops", + lambda: probes.append(True) or {"mm": True, "bmm": True}, + ) + wrapped = _patch._wrap_mm_out_dtype(original, "mm") + a = torch.randn(32, 32, dtype=torch.bfloat16, device="musa") + b = torch.randn(32, 32, dtype=torch.bfloat16, device="musa") + reference = torch.mm(a.float(), b.float()) + + graph = torch.musa.MUSAGraph() + with torch.musa.graph(graph): + captured = wrapped(a, b, out_dtype=torch.float32) + torch.musa.synchronize() + + assert probes == [], "the verdict must not be resolved inside a capture" + + graph.replay() + torch.musa.synchronize() + assert torch.allclose(captured, reference, atol=1e-2, rtol=1e-2) + + with caplog.at_level(logging.INFO, logger=_patch.logger.name): + eager = wrapped(a, b, out_dtype=torch.float32) + + assert probes == [True], "the next eager call resolves the verdict exactly once" + assert torch.allclose(eager, reference, atol=1e-2, rtol=1e-2) + assert "ignores out_dtype" in caplog.text + assert "backport emulation in use" in caplog.text + + def test_backport_is_armed_and_the_probe_resolves_lazily(self): + """A qualifying stack is armed at import; the verdict comes on first use.""" + _require_musa() + from torchada._version import version_of + + gated = version_of(torch.musa.__version__) < "2.13.0" + + assert (getattr(torch.mm, "__wrapped__", None) is not None) == gated + if not gated: + return + + a = torch.randn(2, 3, dtype=torch.bfloat16, device="musa") + b = torch.randn(3, 2, dtype=torch.bfloat16, device="musa") + result = torch.mm(a, b, out_dtype=torch.float32) + + assert set(_patch._mm_out_dtype_probe_cache) == {"mm", "bmm"}, "the first call resolves the verdict" + if _patch._mm_out_dtype_probe_cache["mm"]: + reference = a.to(torch.float32) @ b.to(torch.float32) + assert torch.allclose(result.cpu(), reference.cpu(), atol=ATOL, rtol=RTOL) + else: + assert result.dtype == torch.float32 + + def test_probe_verdict_matches_a_direct_measurement(self): + """A reported-broken op must really be broken, and vice versa.""" + _require_musa() + vendor_mm = _patch._original_torch_mm or torch.mm + a = torch.arange(1, 33, dtype=torch.float32, device="musa").reshape(4, 8) + b = torch.arange(1, 41, dtype=torch.float32, device="musa").reshape(8, 5) + measured = not torch.equal(vendor_mm(a, b, out_dtype=torch.float32), vendor_mm(a, b)) + + assert _patch._probe_out_dtype_broken_ops()["mm"] == measured diff --git a/tests/test_platform.py b/tests/test_platform.py index 531a59d..52de101 100644 --- a/tests/test_platform.py +++ b/tests/test_platform.py @@ -59,7 +59,7 @@ def test_get_version(self): version = torchada.get_version() assert version == torchada.__version__ - assert version == "0.1.88" + assert version == "0.1.89" assert isinstance(version, str) def test_project_version_matches_runtime_version(self): diff --git a/tests/test_version.py b/tests/test_version.py new file mode 100644 index 0000000..d8bc377 --- /dev/null +++ b/tests/test_version.py @@ -0,0 +1,165 @@ +"""Tests for the infix version comparisons in ``torchada._version``. + +No GPU and no torch_musa build is required: the proxy is pure string handling. +""" + +from __future__ import annotations + +import pytest + +from types import SimpleNamespace + +from torchada._version import VersionComparison, version_of + + +class TestVersionOf: + """The accepted shapes of a version source.""" + + def test_accepts_a_string(self): + assert version_of("2.11.0") == "2.11.0" + + def test_accepts_none(self): + assert isinstance(version_of(None), VersionComparison) + assert not version_of(None).is_known + + def test_accepts_an_object_with_a_version_attribute(self): + class _Module: + __version__ = "2.12.0+musa6.0.0" + + assert version_of(_Module()) >= "2.11.0.post2" + assert version_of(_Module).is_known # a module/class works too + + def test_accepts_another_proxy(self): + proxy = version_of("2.11.0.post1") + assert version_of(proxy) is proxy + + def test_missing_attribute_is_unknown(self): + class _Module: + pass + + assert not version_of(_Module()).is_known + + +class TestComparisonSemantics: + """The two documented policies: public version, unknown ranks lowest.""" + + def test_local_segment_is_ignored(self): + # ``+musa5.2.0`` identifies the MUSA stack build, not the fix level. + assert version_of("2.11.0.post1+musa5.2.0") < "2.11.0.post2" + assert not (version_of("2.11.0.post1+musa5.2.0") >= "2.11.0.post2") + assert version_of("2.11.0.post1+musa5.2.0") == version_of("2.11.0.post1") + assert version_of("2.11.0.post2+musa5.2.0") >= "2.11.0.post2" + + def test_post_releases_compare_semantically(self): + assert version_of("2.11.0.post10+musa5.2.0") >= "2.11.0.post2" + assert version_of("2.11.0.post1") < "2.11.0.post10" + + def test_unparsable_version_ranks_lowest(self): + # The "unknown ⇒ keep the workaround" gates rely on this: any upper + # bound is satisfied, so the patch stays enabled. + for unknown in ("not-a-version", "not-a-version+musa5.2.0", "", None): + proxy = version_of(unknown) + assert not proxy.is_known + assert proxy < "2.11.0.post2" + assert proxy < "0.0.1" + assert not (proxy >= "2.11.0.post2") + + def test_lower_bound_gate_skips_only_known_versions(self): + # The mm/bmm gate is a lower bound: an unknown version must NOT be read + # as "newer than the line" (that would disable the backport). + minimum = "2.11.0" + for unknown in ("not-a-version", None): + proxy = version_of(unknown) + assert not (proxy.is_known and proxy < minimum) + assert version_of("2.10.0") < minimum + assert not (version_of("2.11.0") < minimum) + + def test_operators_work_in_both_operand_orders(self): + proxy = version_of("2.11.0.post1") + assert proxy < "2.11.0.post2" + assert "2.11.0.post2" > proxy + assert proxy <= "2.11.0.post1" + assert "2.11.0.post1" >= proxy + assert proxy != "2.11.0.post2" + assert not (proxy == "2.11.0.post2") + + def test_hashing_and_repr(self): + assert len({version_of("2.11.0"), version_of("2.11.0+musa5.2.0")}) == 1 + assert "2.11.0" in repr(version_of("2.11.0")) + assert "unknown" in repr(version_of("nonsense")) + + +class TestTorchMusaGate: + """The gate helper now built on the proxy keeps its old behaviour.""" + + @pytest.mark.parametrize( + ("version", "expected"), + ( + ("2.10.0", True), + ("2.11.0", True), + ("2.11.0.post1+musa5.2.0", True), + ("2.12.0+musa6.0.0", True), + ("2.13.0", False), + ("2.13.0.post1", False), + ("2.14.0", False), + ("not-a-version+musa5.2.0", True), + (None, True), + ), + ) + def test_mm_out_dtype_bound_is_a_single_upper_bound(self, version, expected): + """The arming gate is one inline comparison against the committed fix release. + + Below 2.13.0 the wrappers are armed, from 2.13.0 on nothing is installed, + and an unknown or unparsable version ranks lowest in ``version_of``, so it + stays armed. + """ + assert (version_of(version) < "2.13.0") is expected + + +class TestOutDtypeArming: + """The arming decision, driven through the real patch function (no device).""" + + @pytest.mark.parametrize( + "version,expected", + ( + ("2.10.0", True), + ("2.11.0.post1+musa5.2.0", True), + ("2.12.9+musa6.0.0", True), + ("2.13.0", False), + ("2.13.0.post1+musa6.1.0", False), + ("not-a-version", True), + (None, True), + ), + ) + def test_arming_is_a_single_upper_bound(self, monkeypatch, version, expected): + import sys + from types import ModuleType + + import torch + + from torchada import _patch + + musa = ModuleType("fake_torch_musa") + if version is not None: + musa.__version__ = version + installed = [] + # The patch function is guarded by ``requires_import("torch_musa")``, and this + # test needs the real body to run, so stand in for the import. + monkeypatch.setitem(sys.modules, "torch_musa", ModuleType("torch_musa")) + monkeypatch.setattr(_patch, "is_musa_platform", lambda: True) + monkeypatch.setattr(torch, "musa", musa, raising=False) + monkeypatch.setattr(_patch, "_original_torch_mm", None) + monkeypatch.setattr(_patch, "_original_torch_bmm", None) + monkeypatch.setattr( + _patch, "_register_jit_builtin_alias", lambda original, wrapper: installed.append(wrapper) + ) + original_mm, original_bmm = torch.mm, torch.bmm + try: + _patch._patch_mm_out_dtype() + finally: + torch.mm, torch.bmm = original_mm, original_bmm + _patch._original_torch_mm = None + _patch._original_torch_bmm = None + + assert bool(installed) is expected + pass