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
3 changes: 2 additions & 1 deletion README.md
Original file line number Diff line number Diff line change
Expand Up @@ -66,6 +66,7 @@ That's it! Supported `torch.cuda.*` APIs are automatically redirected to `torch.
| C++ nvJPEG porting | nvJPEG source and build settings → MTJPEG |
| 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 |
| 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 |

Expand Down Expand Up @@ -392,7 +393,7 @@ See `src/torchada/_mappings/` for 400+ mapping rules grouped by API domain.

```
# pyproject.toml or requirements.txt
torchada>=0.1.86
torchada>=0.1.87
```

### Step 2: Conditional Import
Expand Down
3 changes: 2 additions & 1 deletion README_CN.md
Original file line number Diff line number Diff line change
Expand Up @@ -66,6 +66,7 @@ torch.cuda.synchronize()
| C++ nvJPEG 移植 | nvJPEG 源码及构建配置 → MTJPEG |
| ctypes 库加载 | `ctypes.CDLL` 使用 CUDA 函数名 → 自动转换为 MUSA |
| 统一加速器 API | `torch.accelerator.empty_cache()`、`memory_stats()`、`Stream`、`Event` 等 |
| MUSA float64 原地对数 | `torch_musa < 2.11.0.post2` 时,`Tensor.log_()` 复用受支持的非原地操作,同时保持原地操作契约 |
| Triton CUDA Extra | MUSA 上的 `tl.extra.cuda` → `tl.extra.musa` 兼容 |
| Triton 融合 MoE | 面向 vLLM 和 SGLang 的 Triton 3.2.0 MTT S5000 调优配置 |

Expand Down Expand Up @@ -376,7 +377,7 @@ if torchada.is_gpu_device(device): # 在 CUDA 和 MUSA 上都能工作

```
# pyproject.toml 或 requirements.txt
torchada>=0.1.86
torchada>=0.1.87
```

### 步骤 2:条件导入
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.86"
version = "0.1.87"
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.86"
__version__ = "0.1.87"

from . import cuda, utils

Expand Down
28 changes: 26 additions & 2 deletions src/torchada/_patch.py
Original file line number Diff line number Diff line change
Expand Up @@ -40,6 +40,7 @@

_patched = False
_original_init_process_group = None
_original_tensor_log_ = None

# Registry for patch functions
_patch_registry: List[Callable[[], None]] = []
Expand Down Expand Up @@ -160,6 +161,28 @@ def _patch_inductor_template_heuristics():
heuristic_cache.clear()


@patch_function
@requires_import("torch_musa")
def _patch_tensor_log_():
Comment thread
yingzhou-bjtu marked this conversation as resolved.
"""Backport MUSA float64 ``Tensor.log_`` for torch_musa < 2.11.0.post2."""
Comment thread
yingzhou-bjtu marked this conversation as resolved.
global _original_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__):
return

_original_tensor_log_ = torch.Tensor.log_

@functools.wraps(_original_tensor_log_)
def patched_log_(self):
if self.device.type == "musa" and self.dtype == torch.float64:
return self.copy_(torch.log(self))
return _original_tensor_log_(self)

torch.Tensor.log_ = patched_log_


# Cache for translated device strings - avoids repeated string operations
_device_str_cache = {}

Expand Down Expand Up @@ -1780,8 +1803,9 @@ def __getitem__(self, name: str):
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. Older
releases still need torchada to force those calls through torch.musa.
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.

Expand Down
96 changes: 96 additions & 0 deletions tests/test_cuda_patching.py
Original file line number Diff line number Diff line change
Expand Up @@ -3034,6 +3034,102 @@ def invalid_lse_result(
flash_attn._flash_attn_forward("q", "k", "v")


class TestTensorLogPatch:
"""CPU coverage for the MUSA float64 Tensor.log_ compatibility patch."""

def test_float64_log_patch_respects_torch_musa_version(self, monkeypatch):
import sys
from types import ModuleType, SimpleNamespace

import torch

from torchada import _patch

original_log = torch.Tensor.log_
monkeypatch.setattr(torch.Tensor, "log_", original_log)
monkeypatch.setitem(sys.modules, "torch_musa", ModuleType("torch_musa"))
monkeypatch.setattr(_patch, "is_musa_platform", lambda: True)
monkeypatch.setattr(_patch, "_original_tensor_log_", None)
monkeypatch.setattr(
torch,
"musa",
SimpleNamespace(__version__="2.11.0.post2"),
raising=False,
)

_patch._patch_tensor_log_()
assert torch.Tensor.log_ is original_log

torch.musa.__version__ = "2.11.0.post1+musa5.2.0"
_patch._patch_tensor_log_()
assert torch.Tensor.log_ is not original_log

def test_float64_log_patch_uses_out_of_place_copy_on_musa(self, monkeypatch):
import sys
from types import ModuleType, SimpleNamespace

import torch

from torchada import _patch

original_log = torch.Tensor.log_
monkeypatch.setattr(torch.Tensor, "log_", original_log)
monkeypatch.setitem(sys.modules, "torch_musa", ModuleType("torch_musa"))
monkeypatch.setattr(_patch, "is_musa_platform", lambda: True)
monkeypatch.setattr(_patch, "_original_tensor_log_", None)
monkeypatch.setattr(
torch,
"musa",
SimpleNamespace(__version__="2.11.0.post1+musa5.2.0"),
raising=False,
)

_patch._patch_tensor_log_()
patched_log = torch.Tensor.log_

log_args = []
logged = object()

def fake_torch_log(tensor):
log_args.append(tensor)
return logged

original_calls = []

def fake_original_log_(tensor):
original_calls.append(tensor)
return tensor

monkeypatch.setattr(torch, "log", fake_torch_log)
monkeypatch.setattr(_patch, "_original_tensor_log_", fake_original_log_)

class FakeTensor:
def __init__(self, device_type, dtype):
self.device = SimpleNamespace(type=device_type)
self.dtype = dtype
self.copied = None

def copy_(self, other):
self.copied = other
return self

musa_f64 = FakeTensor("musa", torch.float64)
result = patched_log(musa_f64)
assert result is musa_f64
assert musa_f64.copied is logged
assert log_args == [musa_f64]
assert original_calls == []

cpu_f64 = FakeTensor("cpu", torch.float64)
assert patched_log(cpu_f64) is cpu_f64
assert cpu_f64.copied is None

musa_f32 = FakeTensor("musa", torch.float32)
assert patched_log(musa_f32) is musa_f32
assert musa_f32.copied is None
assert original_calls == [cpu_f64, musa_f32]


class TestAcceleratorModuleWrapper:
"""Test the _AcceleratorModuleWrapper priority / fallback logic in isolation.

Expand Down
2 changes: 1 addition & 1 deletion tests/test_platform.py
Original file line number Diff line number Diff line change
Expand Up @@ -59,7 +59,7 @@ def test_get_version(self):

version = torchada.get_version()
assert version == torchada.__version__
assert version == "0.1.86"
assert version == "0.1.87"
assert isinstance(version, str)

def test_project_version_matches_runtime_version(self):
Expand Down