Reuse the torch2.10 prebuilt accelerator wheels on torch 2.12

Studio already falls back to the torch2.10 flash-attn / causal-conv1d /
mamba-ssm wheels when it finds torch 2.11, because upstream publishes no
2.11-tagged builds. torch 2.12 is in exactly the same position, and the same
wheels work there, so a 2.12 install currently drops to a source build for no
reason.

Measured on a B200, python 3.12, fresh uv venv on torch 2.12.1+cu130, wheels
installed with --no-deps and torch verified unmoved afterwards, importing the
compiled .so directly rather than only the Python package:

  causal-conv1d 1.6.1   9412 passed / 3888 skipped / 0 failed
  mamba-ssm 2.3.1       tests/ops, 20 passed
  flash-attn 2.8.1      splitkv + qkvpacked subset, 848 passed

Against a torch 2.10 control the pass/fail/skip counts match and the failing
test-ID sets are byte identical.

The reuse window is bounded rather than open ended, so the comment now records
that. flash-attn v2.8.3.post1's torch2.9 wheel fails to import on torch 2.10
and on torch 2.12 alike, with an undefined symbol out of flash_attn_2_cuda:
torch broke extension ABI between 2.9 and 2.10 and has held it from 2.10
through 2.12. A wheel cannot skip a torch minor backwards, so torch 2.13 is
deliberately left out of the table until it is measured.

The torch2.10 flash-attn pin stays at 2.8.1. v2.8.3 looks like a free upgrade
but publishes only 2 of the 8 torch2.10 assets that v2.8.1 does, keeping just
cu13/cp312 for x86_64 and aarch64 and dropping every cu12 and every cp313
torch2.10 wheel, while v2.8.3.post1 dropped the torch2.10 assets entirely.
Bumping the pin would silently 404 most users back to a source build, so the
constant now carries that warning.

Tests cover the 2.12 mapping through both direct_wheel_url and the flash-attn
URL builder, that reuse only ever targets torch2.10, and that the selected
flash-attn version is never a .post release.
This commit is contained in:
Daniel Han 2026-07-27 10:42:37 +00:00
commit 741bebb36a
2 changed files with 71 additions and 7 deletions

View file

@ -20,10 +20,20 @@ class TestPrebuiltWheelTorchMapping:
def test_torch_211_maps_to_torch210(self):
assert wheel_utils.prebuilt_wheel_torch_mm("2.11") == "2.10"
def test_torch_212_maps_to_torch210(self):
assert wheel_utils.prebuilt_wheel_torch_mm("2.12") == "2.10"
def test_other_versions_pass_through(self):
for torch_mm in ("2.9", "2.10", "2.12"):
# 2.13 stays unmapped on purpose: a torch minor only joins the reuse
# table once its wheels have actually been measured.
for torch_mm in ("2.9", "2.10", "2.13"):
assert wheel_utils.prebuilt_wheel_torch_mm(torch_mm) == torch_mm
def test_reuse_never_targets_a_pre_210_wheel(self):
# torch broke extension ABI between 2.9 and 2.10, so the torch2.9 .so
# raises "undefined symbol" on 2.10+. Reuse may only point at torch2.10.
assert set(wheel_utils._PREBUILT_WHEEL_TORCH_MM.values()) == {"2.10"}
def test_direct_wheel_url_reuses_torch210_on_211(self):
# causal-conv1d / mamba go through direct_wheel_url; torch 2.11 reuses the
# torch2.10 wheel filename just like flash-attn does.
@ -43,11 +53,39 @@ class TestPrebuiltWheelTorchMapping:
assert url is not None
assert "causal_conv1d-1.6.1+cu13torch2.10cxx11abiTRUE-cp313-cp313-linux_x86_64.whl" in url
def test_direct_wheel_url_reuses_torch210_on_212(self):
url = wheel_utils.direct_wheel_url(
filename_prefix = "mamba_ssm",
package_version = "2.3.1",
release_tag = "v2.3.1",
release_base_url = "https://example.test/download",
env = {
"python_tag": "cp312",
"torch_mm": "2.12",
"cuda_major": "13",
"cxx11abi": "TRUE",
"platform_tag": "linux_x86_64",
},
)
assert url is not None
assert "mamba_ssm-2.3.1+cu13torch2.10cxx11abiTRUE-cp312-cp312-linux_x86_64.whl" in url
class TestFlashAttnWheelSelection:
def test_torch_210_maps_to_v281(self):
# v2.8.1 is the newest release still publishing the full torch2.10 asset
# matrix (cu12 + cu13, cp312 + cp313, x86_64 + aarch64).
assert ips._select_flash_attn_version("2.10") == "2.8.1"
def test_selected_version_is_never_a_post_release(self):
# The v2.8.3.post1 respin dropped every torch2.10 asset and stops at
# torch2.9, whose .so will not load on torch 2.10+. A future "just take
# the newest release" bump must fail here instead of shipping that.
for torch_mm in ("2.4", "2.7", "2.9", "2.10"):
version = ips._select_flash_attn_version(torch_mm)
assert version is not None
assert ".post" not in version
def test_torch_29_maps_to_v283(self):
assert ips._select_flash_attn_version("2.9") == "2.8.3"
@ -69,6 +107,19 @@ class TestFlashAttnWheelSelection:
assert url is not None
assert "flash_attn-2.8.1+cu13torch2.10cxx11abiTRUE-cp313-cp313-linux_x86_64.whl" in url
def test_torch_212_reuses_torch210_wheel(self):
url = ips._build_flash_attn_wheel_url(
{
"python_tag": "cp313",
"torch_mm": "2.12",
"cuda_major": "13",
"cxx11abi": "TRUE",
"platform_tag": "linux_x86_64",
}
)
assert url is not None
assert "flash_attn-2.8.1+cu13torch2.10cxx11abiTRUE-cp313-cp313-linux_x86_64.whl" in url
def test_exact_wheel_url_uses_full_env_tuple(self):
url = ips._build_flash_attn_wheel_url(
{