Add flash-linear-attention package for hardware that supports it.

This commit is contained in:
Jaret Burkett
2026-07-29 09:14:37 -06:00
parent 1e22732db7
commit 3bd2119c04
2 changed files with 43 additions and 14 deletions

View File

@@ -441,15 +441,23 @@ def _stale_git_pins(spec):
return [name for name, line in current.items() if stored.get(name) != line]
def _optional_import_name(pkg):
"""Importable module name for an optional package spec or wheel URL."""
# optional packages whose import name differs from the distribution name
_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:
name = os.path.basename(pkg).split("-")[0]
else:
name = pkg
for sep in ("==", ">=", "<=", "<", ">", "["):
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):
@@ -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
_verify_torch(spec, dry_run=dry_run)
# 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):
label = os.path.basename(pkg) if "://" in pkg else pkg
info("Installing optional accelerator: %s" % label)
code = _pip_install(
[pkg] + find_links, dry_run=dry_run, upgrade=True, check=False
)
code = _pip_install([pkg] + find_links, dry_run=dry_run, check=False)
if code != 0:
warn("Optional package failed to install (continuing): %s" % label)
continue
@@ -524,13 +534,14 @@ def ensure_requirements(spec, dry_run=False, force=False):
# prebuilt accelerator wheels are sometimes built against a torch
# nightly and fail to load against the release ABI — verify the
# import and roll back rather than leaving a broken wheel installed
name = _optional_import_name(pkg)
if not _venv_import_ok(name):
dist_name, import_name = _optional_names(pkg)
if not _venv_import_ok(import_name):
warn(
"%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
# which work fine without it (e.g. tensorboard's grpcio on win_arm64)
for pkg in spec.no_deps_packages:

View File

@@ -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.
- 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.
- 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
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
@@ -59,6 +64,10 @@ TRITON_WINDOWS = "triton-windows>=3.7,<3.8"
NATTEN_VERSION = "0.21.7"
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"
_FA_BASE = (
"https://github.com/mjun0812/flash-attention-prebuild-wheels/releases/download/"
@@ -309,7 +318,7 @@ def _cuda_spec(detection):
]
extras = [TORCHCODEC]
optional = []
optional = [FLA]
find_links = []
# 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",
],
# 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],
notes=[
"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"]
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":
notes.append("Intel Mac detected — training will be extremely slow.")
return EnvSpec(
@@ -443,6 +459,8 @@ def _build_spec(detection, allow_cpu=False):
TORCH,
torch_index=PYTORCH_INDEX + "rocm7.1",
extra_packages=[TORCHCODEC],
# runs on ROCm via the triton bundled with rocm torch
optional_packages=[FLA],
notes=[
"AMD ROCm support is experimental and largely untested.",
"flash-attn / NATTEN prebuilt wheels are unavailable for ROCm.",