Add flash-linear-attention package for hardware that supports it.
This commit is contained in:
@@ -441,15 +441,23 @@ def _stale_git_pins(spec):
|
|||||||
return [name for name, line in current.items() if stored.get(name) != line]
|
return [name for name, line in current.items() if stored.get(name) != line]
|
||||||
|
|
||||||
|
|
||||||
def _optional_import_name(pkg):
|
# optional packages whose import name differs from the distribution name
|
||||||
"""Importable module name for an optional package spec or wheel URL."""
|
_IMPORT_ALIASES = {"flash_linear_attention": "fla"}
|
||||||
|
# companion dists to remove on rollback (fla-core provides the `fla` module
|
||||||
|
# itself; leaving it behind would keep the broken import resolvable)
|
||||||
|
_ROLLBACK_EXTRAS = {"flash_linear_attention": ["fla-core"]}
|
||||||
|
|
||||||
|
|
||||||
|
def _optional_names(pkg):
|
||||||
|
"""(distribution, import) names for an optional package spec or wheel URL."""
|
||||||
if "://" in pkg:
|
if "://" in pkg:
|
||||||
name = os.path.basename(pkg).split("-")[0]
|
name = os.path.basename(pkg).split("-")[0]
|
||||||
else:
|
else:
|
||||||
name = pkg
|
name = pkg
|
||||||
for sep in ("==", ">=", "<=", "<", ">", "["):
|
for sep in ("==", ">=", "<=", "<", ">", "["):
|
||||||
name = name.split(sep)[0]
|
name = name.split(sep)[0]
|
||||||
return name.strip().replace("-", "_")
|
name = name.strip().replace("-", "_")
|
||||||
|
return name, _IMPORT_ALIASES.get(name, name)
|
||||||
|
|
||||||
|
|
||||||
def _venv_import_ok(module_name):
|
def _venv_import_ok(module_name):
|
||||||
@@ -509,13 +517,15 @@ def ensure_requirements(spec, dry_run=False, force=False):
|
|||||||
# actually intend to ship, so repair it first if something got through
|
# actually intend to ship, so repair it first if something got through
|
||||||
_verify_torch(spec, dry_run=dry_run)
|
_verify_torch(spec, dry_run=dry_run)
|
||||||
# accelerators (flash-attn, NATTEN, ...): install one-by-one, warn on
|
# accelerators (flash-attn, NATTEN, ...): install one-by-one, warn on
|
||||||
# failure — training works without them, so never fail the whole install
|
# failure — training works without them, so never fail the whole install.
|
||||||
|
# No --upgrade here: every optional spec is an exact pin or wheel URL, and
|
||||||
|
# uv's -U eagerly upgrades the whole dependency closure, blowing past
|
||||||
|
# requirements.txt pins (numpy/transformers) that only the requirements
|
||||||
|
# pass enforces.
|
||||||
for pkg in _filter_extras(spec.optional_packages):
|
for pkg in _filter_extras(spec.optional_packages):
|
||||||
label = os.path.basename(pkg) if "://" in pkg else pkg
|
label = os.path.basename(pkg) if "://" in pkg else pkg
|
||||||
info("Installing optional accelerator: %s" % label)
|
info("Installing optional accelerator: %s" % label)
|
||||||
code = _pip_install(
|
code = _pip_install([pkg] + find_links, dry_run=dry_run, check=False)
|
||||||
[pkg] + find_links, dry_run=dry_run, upgrade=True, check=False
|
|
||||||
)
|
|
||||||
if code != 0:
|
if code != 0:
|
||||||
warn("Optional package failed to install (continuing): %s" % label)
|
warn("Optional package failed to install (continuing): %s" % label)
|
||||||
continue
|
continue
|
||||||
@@ -524,13 +534,14 @@ def ensure_requirements(spec, dry_run=False, force=False):
|
|||||||
# prebuilt accelerator wheels are sometimes built against a torch
|
# prebuilt accelerator wheels are sometimes built against a torch
|
||||||
# nightly and fail to load against the release ABI — verify the
|
# nightly and fail to load against the release ABI — verify the
|
||||||
# import and roll back rather than leaving a broken wheel installed
|
# import and roll back rather than leaving a broken wheel installed
|
||||||
name = _optional_import_name(pkg)
|
dist_name, import_name = _optional_names(pkg)
|
||||||
if not _venv_import_ok(name):
|
if not _venv_import_ok(import_name):
|
||||||
warn(
|
warn(
|
||||||
"%s installed but fails to import against this torch build — "
|
"%s installed but fails to import against this torch build — "
|
||||||
"removing it (training falls back to native attention)." % name
|
"removing it (training falls back to native attention)."
|
||||||
|
% import_name
|
||||||
)
|
)
|
||||||
_pip_uninstall([name])
|
_pip_uninstall([dist_name] + _ROLLBACK_EXTRAS.get(dist_name, []))
|
||||||
# packages whose dependency metadata is unsatisfiable on this platform but
|
# packages whose dependency metadata is unsatisfiable on this platform but
|
||||||
# which work fine without it (e.g. tensorboard's grpcio on win_arm64)
|
# which work fine without it (e.g. tensorboard's grpcio on win_arm64)
|
||||||
for pkg in spec.no_deps_packages:
|
for pkg in spec.no_deps_packages:
|
||||||
|
|||||||
@@ -21,6 +21,11 @@ kernels smoke-tested on an RTX 5090 / sm120 with torch 2.13.0+cu130):
|
|||||||
{cp310..cp314} x {linux x86_64, linux aarch64}. No Windows/mac wheels.
|
{cp310..cp314} x {linux x86_64, linux aarch64}. No Windows/mac wheels.
|
||||||
- triton: bundled with torch on Linux (incl. aarch64; torch 2.13 bundles
|
- triton: bundled with torch on Linux (incl. aarch64; torch 2.13 bundles
|
||||||
triton 3.7.1); triton-windows 3.7.x matches on Windows; nothing for MPS.
|
triton 3.7.1); triton-windows 3.7.x matches on Windows; nothing for MPS.
|
||||||
|
- flash-linear-attention 0.5.2: pure-Python (py3-none-any) Triton kernels —
|
||||||
|
installs anywhere, needs triton>=3.3 at runtime. Installed bare (no backend
|
||||||
|
extra) so it never pulls its own torch/triton over our pinned stack; usable
|
||||||
|
on every platform with a working triton (CUDA linux/windows, Spark, ROCm),
|
||||||
|
not on MPS/CPU.
|
||||||
- Windows-on-ARM (verified 2026-07): there are NO win_arm64 wheels for any of
|
- Windows-on-ARM (verified 2026-07): there are NO win_arm64 wheels for any of
|
||||||
the CUDA stack — torch cu130, triton-windows, flash-attn and Prisma's node
|
the CUDA stack — torch cu130, triton-windows, flash-attn and Prisma's node
|
||||||
engine are all x64-only. The supported configuration is therefore the x64
|
engine are all x64-only. The supported configuration is therefore the x64
|
||||||
@@ -59,6 +64,10 @@ TRITON_WINDOWS = "triton-windows>=3.7,<3.8"
|
|||||||
NATTEN_VERSION = "0.21.7"
|
NATTEN_VERSION = "0.21.7"
|
||||||
NATTEN_FIND_LINKS = "https://whl.natten.org"
|
NATTEN_FIND_LINKS = "https://whl.natten.org"
|
||||||
|
|
||||||
|
# pure-Python triton kernels; bare install (no [cuda]/[rocm] extra) on purpose —
|
||||||
|
# the extras only add torch/triton pins we already manage per-platform
|
||||||
|
FLA = "flash-linear-attention==0.5.2"
|
||||||
|
|
||||||
FLASH_ATTN_VERSION = "2.8.3"
|
FLASH_ATTN_VERSION = "2.8.3"
|
||||||
_FA_BASE = (
|
_FA_BASE = (
|
||||||
"https://github.com/mjun0812/flash-attention-prebuild-wheels/releases/download/"
|
"https://github.com/mjun0812/flash-attention-prebuild-wheels/releases/download/"
|
||||||
@@ -309,7 +318,7 @@ def _cuda_spec(detection):
|
|||||||
]
|
]
|
||||||
|
|
||||||
extras = [TORCHCODEC]
|
extras = [TORCHCODEC]
|
||||||
optional = []
|
optional = [FLA]
|
||||||
find_links = []
|
find_links = []
|
||||||
|
|
||||||
# Windows-on-ARM runs the x64 wheel stack (see module docstring), so wheel
|
# Windows-on-ARM runs the x64 wheel stack (see module docstring), so wheel
|
||||||
@@ -370,7 +379,11 @@ def _make_spark_spec(detection, wheels_source, dll_dirs):
|
|||||||
"triton==3.8.0+git8743423b",
|
"triton==3.8.0+git8743423b",
|
||||||
],
|
],
|
||||||
# import-verified with rollback, like accelerators on other platforms
|
# import-verified with rollback, like accelerators on other platforms
|
||||||
optional_packages=["flash-attn==2.8.3+cu134torch2.14", "natten==0.21.7"],
|
optional_packages=[
|
||||||
|
"flash-attn==2.8.3+cu134torch2.14",
|
||||||
|
"natten==0.21.7",
|
||||||
|
FLA,
|
||||||
|
],
|
||||||
find_links=[wheels_source],
|
find_links=[wheels_source],
|
||||||
notes=[
|
notes=[
|
||||||
"RTX Spark native mode: win_arm64 CUDA %s stack from the "
|
"RTX Spark native mode: win_arm64 CUDA %s stack from the "
|
||||||
@@ -423,7 +436,10 @@ def _build_spec(detection, allow_cpu=False):
|
|||||||
os_name = detection["os"]
|
os_name = detection["os"]
|
||||||
|
|
||||||
if os_name == "mac":
|
if os_name == "mac":
|
||||||
notes = ["flash-attn / NATTEN / triton are unavailable on macOS."]
|
notes = [
|
||||||
|
"flash-attn / NATTEN / triton / flash-linear-attention are "
|
||||||
|
"unavailable on macOS."
|
||||||
|
]
|
||||||
if detection["backend"] != "mps":
|
if detection["backend"] != "mps":
|
||||||
notes.append("Intel Mac detected — training will be extremely slow.")
|
notes.append("Intel Mac detected — training will be extremely slow.")
|
||||||
return EnvSpec(
|
return EnvSpec(
|
||||||
@@ -443,6 +459,8 @@ def _build_spec(detection, allow_cpu=False):
|
|||||||
TORCH,
|
TORCH,
|
||||||
torch_index=PYTORCH_INDEX + "rocm7.1",
|
torch_index=PYTORCH_INDEX + "rocm7.1",
|
||||||
extra_packages=[TORCHCODEC],
|
extra_packages=[TORCHCODEC],
|
||||||
|
# runs on ROCm via the triton bundled with rocm torch
|
||||||
|
optional_packages=[FLA],
|
||||||
notes=[
|
notes=[
|
||||||
"AMD ROCm support is experimental and largely untested.",
|
"AMD ROCm support is experimental and largely untested.",
|
||||||
"flash-attn / NATTEN prebuilt wheels are unavailable for ROCm.",
|
"flash-attn / NATTEN prebuilt wheels are unavailable for ROCm.",
|
||||||
|
|||||||
Reference in New Issue
Block a user