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]
|
||||
|
||||
|
||||
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:
|
||||
|
||||
@@ -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.",
|
||||
|
||||
Reference in New Issue
Block a user