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
2 changes: 1 addition & 1 deletion README.md
Original file line number Diff line number Diff line change
Expand Up @@ -392,7 +392,7 @@ See `src/torchada/_mappings/` for 400+ mapping rules grouped by API domain.

```
# pyproject.toml or requirements.txt
torchada>=0.1.81
torchada>=0.1.82
```

### Step 2: Conditional Import
Expand Down
2 changes: 1 addition & 1 deletion README_CN.md
Original file line number Diff line number Diff line change
Expand Up @@ -375,7 +375,7 @@ if torchada.is_gpu_device(device): # 在 CUDA 和 MUSA 上都能工作

```
# pyproject.toml 或 requirements.txt
torchada>=0.1.81
torchada>=0.1.82
```

### 步骤 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.81",
"version": "0.1.82",
"date": "2026-01-29",
"platform": "MUSA",
"pytorch_version": "2.7.1",
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.81"
version = "0.1.82"
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.81"
__version__ = "0.1.82"

from . import cuda, utils

Expand Down
49 changes: 49 additions & 0 deletions src/torchada/_patch.py
Original file line number Diff line number Diff line change
Expand Up @@ -1587,6 +1587,55 @@ def _accepts_only_qv(func: Callable) -> bool:
continue
setattr(flash_attn_interface, _name, _drop_only_qv(_func))

# Ring Attention imports the private forward API for its softmax LSE.
# Adapt the public output+LSE contract only when the provider omits it.
if not hasattr(flash_attn_interface, "_flash_attn_forward"):
_public_flash_attn = getattr(flash_attn_interface, "flash_attn_func", None)
try:
_public_parameters = inspect.signature(_public_flash_attn).parameters
except (TypeError, ValueError):
_public_parameters = {}

if callable(_public_flash_attn) and "return_softmax_lse" in _public_parameters:

def _flash_attn_forward(
q,
k,
v,
*,
softmax_scale=None,
causal=False,
window_size_left=-1,
window_size_right=-1,
softcap=0.0,
):
"""Adapt public FA3 output+LSE to its low-level Ring contract."""
result = _public_flash_attn(
q,
k,
v,
softmax_scale=softmax_scale,
causal=causal,
window_size=(window_size_left, window_size_right),
softcap=softcap,
return_softmax_lse=True,
)
if not isinstance(result, (tuple, list)) or len(result) < 2:
raise RuntimeError(
"flash_attn_func(return_softmax_lse=True) must return "
"(output, softmax_lse)"
)
output, softmax_lse = result[:2]
# Inference-only Ring consumers use output and LSE; the public
# API does not expose the private dropout auxiliaries.
return output, softmax_lse, None, None

_flash_attn_forward.__name__ = "_flash_attn_forward"
_flash_attn_forward.__qualname__ = "_flash_attn_forward"
_flash_attn_forward.__module__ = flash_attn_interface.__name__
_flash_attn_forward._torchada_compat_shim = True
flash_attn_interface._flash_attn_forward = _flash_attn_forward


class _CDLLWrapper:
"""
Expand Down
122 changes: 122 additions & 0 deletions tests/test_cuda_patching.py
Original file line number Diff line number Diff line change
Expand Up @@ -2755,6 +2755,128 @@ def flash_attn_kwargs_func(q, **kwargs):
if mod is not None:
sys.modules[name] = mod

@staticmethod
def _install_fa3_test_provider(monkeypatch, flash_attn_func, native_forward=None):
import sys
from types import ModuleType

flash_attn = ModuleType("flash_attn_interface")
flash_attn.flash_attn_func = flash_attn_func
if native_forward is not None:
flash_attn._flash_attn_forward = native_forward
sgl_kernel = ModuleType("sgl_kernel")
sgl_kernel.__path__ = []
sgl_kernel.flash_attn = flash_attn
monkeypatch.setitem(sys.modules, "flash_attn_interface", flash_attn)
monkeypatch.setitem(sys.modules, "sgl_kernel", sgl_kernel)
monkeypatch.setitem(sys.modules, "sgl_kernel.flash_attn", flash_attn)
return flash_attn

def test_missing_fa3_private_forward_is_adapted(self, monkeypatch):
"""Adapt the public output+LSE API to Ring's private FA3 contract."""
from torchada._patch import _patch_flash_attn

calls = []

def flash_attn_func(
q,
k,
v,
*,
softmax_scale=None,
causal=False,
window_size=(-1, -1),
softcap=0.0,
return_softmax_lse=False,
):
calls.append(
{
"q": q,
"k": k,
"v": v,
"softmax_scale": softmax_scale,
"causal": causal,
"window_size": window_size,
"softcap": softcap,
"return_softmax_lse": return_softmax_lse,
}
)
return "output", "softmax_lse"

flash_attn = self._install_fa3_test_provider(monkeypatch, flash_attn_func)

_patch_flash_attn()
from flash_attn_interface import _flash_attn_forward

assert _flash_attn_forward is flash_attn._flash_attn_forward
result = flash_attn._flash_attn_forward(
"q",
"k",
"v",
softmax_scale=0.125,
causal=True,
window_size_left=32,
window_size_right=16,
softcap=4.0,
)

assert result == ("output", "softmax_lse", None, None)
assert calls == [
{
"q": "q",
"k": "k",
"v": "v",
"softmax_scale": 0.125,
"causal": True,
"window_size": (32, 16),
"softcap": 4.0,
"return_softmax_lse": True,
}
]
assert flash_attn._flash_attn_forward.__name__ == "_flash_attn_forward"
assert flash_attn._flash_attn_forward._torchada_compat_shim is True
first = flash_attn._flash_attn_forward
_patch_flash_attn()
assert flash_attn._flash_attn_forward is first

def native_forward(q, k, v, **kwargs):
return q, k, v, kwargs

native_provider = self._install_fa3_test_provider(
monkeypatch, flash_attn_func, native_forward
)
_patch_flash_attn()
assert native_provider._flash_attn_forward is native_forward

def test_fa3_private_forward_fails_closed_without_lse(self, monkeypatch):
"""Require both a public LSE parameter and a valid LSE result."""
from torchada._patch import _patch_flash_attn

def no_lse_support(q, k, v):
return q

flash_attn = self._install_fa3_test_provider(monkeypatch, no_lse_support)
_patch_flash_attn()
assert not hasattr(flash_attn, "_flash_attn_forward")

def invalid_lse_result(
q,
k,
v,
*,
softmax_scale=None,
causal=False,
window_size=(-1, -1),
softcap=0.0,
return_softmax_lse=False,
):
return q

flash_attn = self._install_fa3_test_provider(monkeypatch, invalid_lse_result)
_patch_flash_attn()
with pytest.raises(RuntimeError, match="must return"):
flash_attn._flash_attn_forward("q", "k", "v")


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.81"
assert version == "0.1.82"
assert isinstance(version, str)

def test_project_version_matches_runtime_version(self):
Expand Down