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] 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:

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. {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.",