Compare commits
2 commits
main
...
studio-whe
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
ef5059614d | ||
|
|
741bebb36a |
2 changed files with 71 additions and 7 deletions
|
|
@ -120,12 +120,19 @@ def probe_torch_wheel_env(*, timeout: int | None = None) -> dict[str, str] | Non
|
|||
return env
|
||||
|
||||
|
||||
# torch 2.11 has no native prebuilt wheels for flash-attn / causal-conv1d / mamba
|
||||
# yet, but their torch 2.10 CUDA wheels load and pass the projects' own test suites
|
||||
# on torch 2.11 (verified on B200: FA2 fwd/bwd, causal-conv1d, and mamba selective
|
||||
# scan all match reference). Reuse the torch 2.10 wheels on torch 2.11 so a 2.11
|
||||
# install still gets these prebuilt accelerators instead of building from source.
|
||||
_PREBUILT_WHEEL_TORCH_MM = {"2.11": "2.10"}
|
||||
# torch 2.11 and 2.12 ship no native prebuilt wheels for flash-attn /
|
||||
# causal-conv1d / mamba-ssm, but the torch2.10 CUDA wheels load and pass each
|
||||
# project's own suite on both (B200, py3.12, torch 2.12.1+cu130: causal-conv1d
|
||||
# 9412 passed / 3888 skipped / 0 failed, mamba tests/ops 20 passed, flash-attn
|
||||
# splitkv+qkvpacked 848 passed; pass/fail/skip counts and the failing test-ID
|
||||
# sets are identical to a torch 2.10 control). Reuse them so a 2.11 / 2.12
|
||||
# install still gets prebuilt accelerators instead of building from source.
|
||||
#
|
||||
# The window is bounded, not open ended: torch broke extension ABI between 2.9
|
||||
# and 2.10, and the torch2.9 flash-attn .so raises "undefined symbol" on torch
|
||||
# 2.10 and on 2.12 alike. A wheel cannot skip a torch minor backwards, so every
|
||||
# new key here must be measured against the real wheels before it is added.
|
||||
_PREBUILT_WHEEL_TORCH_MM = {"2.11": "2.10", "2.12": "2.10"}
|
||||
|
||||
|
||||
def prebuilt_wheel_torch_mm(torch_mm: str) -> str:
|
||||
|
|
@ -155,6 +162,12 @@ def direct_wheel_url(
|
|||
|
||||
def flash_attn_package_version(torch_mm: str) -> str | None:
|
||||
if torch_mm == "2.10":
|
||||
# Newest flash-attn release still carrying the full torch2.10 asset
|
||||
# matrix (cu12 + cu13, cp312 + cp313, x86_64 + aarch64). Do not bump
|
||||
# this to "the latest release": v2.8.3 publishes only cu13/cp312 for
|
||||
# torch2.10 and v2.8.3.post1 dropped every torch2.10 asset, so both
|
||||
# 404 most users back to a source build, and post1's newest tag is
|
||||
# torch2.9, which will not load here at all.
|
||||
return "2.8.1"
|
||||
try:
|
||||
major, minor = (int(part) for part in torch_mm.split(".", 1))
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
{
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue