Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
16 changes: 15 additions & 1 deletion README.md
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down
11 changes: 10 additions & 1 deletion README_CN.md
Original file line number Diff line number Diff line change
Expand Up @@ -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;若承诺延期,则把上界抬到新的承诺版本。

## 示例

### 混合精度训练
Expand Down Expand Up @@ -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:条件导入
Expand Down
2 changes: 1 addition & 1 deletion benchmarks/benchmark_history.json
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand Down
2 changes: 1 addition & 1 deletion examples/extension_setup.py
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand Down
2 changes: 1 addition & 1 deletion pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -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"}
Expand Down
2 changes: 1 addition & 1 deletion src/torchada/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down
8 changes: 6 additions & 2 deletions src/torchada/_cpp_ops.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"

Expand Down
Loading