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