studio: auto-discover FLA-using model types from installed transformers

Drop the hand-maintained `_TILELANG_MODEL_SUBSTRINGS` tuple
(qwen3.5 / qwen3_5 / qwen3.6 / qwen3_6 / qwen3-next / qwen3_next)
and derive the allowlist by scanning the installed
`transformers/models/*/modeling_*.py` for `from fla.` imports.

A model "wants tilelang" iff its modeling file imports an FLA op,
which is the same signal `is_flash_linear_attention_available()` is
the runtime test for. The scan happens once per worker subprocess
and is cached for the process lifetime; an empty result (eg
transformers not importable) means "no tilelang pre-install" --
the FLA runtime hook still drives the install via the gate when
the loaded model actually probes it.

Verified against the live installed transformers, the auto-derived
set is {qwen3_5, qwen3_5_moe, qwen3_next}, with `_model_wants_tilelang`
matching the HF Hub names `unsloth/Qwen3.5-2B`, `Qwen/Qwen3.5-MoE-A3B`,
`mlx-community/qwen3-next-80b`, and correctly rejecting Llama,
Mistral, Nemotron-H, Falcon-H1, etc. Future GDN models (Qwen3.7,
OLMo-Hybrid-FA, ...) are picked up automatically once they ship in
transformers; no further worker edits needed.

Also trim docstrings / comments through the FLA / tilelang / HIP /
hook block: constants get 1-line trailing comments, function
docstrings collapse to 1-3 lines, and the fast-path-hooks banner
shrinks from a 27-line block to 4 lines. The file drops from 2847
to 2630 lines without losing the load-bearing WHY notes
(--no-deps protects torch; `__dict__.get` avoids lazy-module
__getattr__; two-step tvm-ffi repair keeps torch off the dep
graph; HIP setdefault disables FLA's TileLang dispatch even with
tilelang already installed).

7 new tests (50 -> 57 total): discovery returns only FLA-using
model_types; discovery cache reuse; missing transformers handled;
OSError on a modeling file is non-fatal; `_model_wants_tilelang`
matches real HF repo names across separator variants; empty
discovery -> always False; normalization across `-`, `.`, `/`,
space.
This commit is contained in:
danielhanchen 2026-05-17 11:18:03 +00:00
commit c358b05734
2 changed files with 256 additions and 317 deletions

View file

@ -52,42 +52,22 @@ _MAMBA_SSM_RELEASE_TAG = "v2.3.1"
_MAMBA_SSM_PACKAGE_VERSION = "2.3.1"
_FLASH_ATTN_RUNTIME_MIN_SEQ_LEN = 32768
_FLASH_ATTN_SKIP_ENV = "UNSLOTH_STUDIO_SKIP_FLASHATTN_INSTALL"
# tilelang 0.1.9+ pairs with apache-tvm-ffi >=0.1.10 by default, but
# apache-tvm-ffi 0.1.10/0.1.11 has an alignment regression that crashes
# subsequent Triton kernels with "CUDA: misaligned address" on sm_100
# (Blackwell). 0.1.9 is the last known-good. mamba_ssm 2.3.2 also pins
# apache-tvm-ffi<=0.1.9, which is the original source of this pin.
# apache-tvm-ffi 0.1.10/0.1.11 crash Triton with "CUDA: misaligned address" on sm_100.
_TILELANG_PACKAGE_VERSION = "0.1.8"
_APACHE_TVM_FFI_PACKAGE_VERSION = "0.1.9"
_TILELANG_SKIP_ENV = "UNSLOTH_STUDIO_SKIP_TILELANG_INSTALL"
# fla-core 0.5.0 requires torch>=2.7.0; pin both so plain pip never
# upgrades torch underneath the Studio venv.
# Pin both so plain pip cannot silently upgrade torch under the worker (fla-core needs torch>=2.7).
_FLA_PACKAGE_VERSION = "0.5.0"
_FLA_CORE_PACKAGE_VERSION = "0.5.0"
_FLA_SKIP_ENV = "UNSLOTH_STUDIO_SKIP_FLA_INSTALL"
# fla-core declares `einops` in its METADATA but `fla/utils.py`
# also imports `packaging` at module load; that one is NOT declared
# upstream (an FLA bug). triton is a torch dep but we list it
# defensively because some torch wheel builds skip it. With --no-deps
# we have to bring these in ourselves, otherwise `import fla.modules`
# raises ModuleNotFoundError at startup.
# `--no-deps` saves torch but loses fla-core's transitive deps; `packaging` is also undeclared upstream.
_FLA_RUNTIME_DEPS = ("einops", "packaging", "triton")
# Studio installer permits torch>=2.4,<2.11.0 but fla-core 0.5.0
# declares torch>=2.7.0; skip FLA on older torch to keep the
# fallback path clean.
_FLA_MIN_TORCH = (2, 7)
# flash-linear-attention and tilelang both require Python >=3.10.
_FLA_MIN_PYTHON = (3, 10)
# tilelang 0.1.8 wheels: Linux x86_64 / aarch64 and macOS arm64.
# We never want to fall back to its 93MB sdist on a Studio worker.
# tilelang 0.1.8 ships wheels only for these Linux arches and macOS arm64; never fall back to its 93MB sdist.
_TILELANG_SUPPORTED_LINUX_MACHINES = frozenset(("x86_64", "amd64", "aarch64", "arm64"))
_TILELANG_INSTALL_TIMEOUT_S = 600
# apache-tvm-ffi 0.1.10/0.1.11 trigger "CUDA: misaligned address" on
# sm_100. If we detect a stale broken version, force a reinstall.
_TVM_FFI_BROKEN_VERSIONS = ("0.1.10", "0.1.11")
# Set to "1" to fall back to the substring-based gate for FLA / tilelang
# installs. Normal operation hooks transformers' availability functions
# so the install fires only when the loaded model actually checks them.
_FAST_PATH_HOOKS_SKIP_ENV = "UNSLOTH_STUDIO_SKIP_FAST_PATH_HOOKS"
@ -325,13 +305,7 @@ def _installed_torch_version_tuple() -> tuple[int, int] | None:
def _flash_linear_attention_importable() -> bool:
"""Best-effort import probe.
Catches arbitrary exceptions (not just ImportError) so a broken
optional package (OSError on missing native lib, RuntimeError from a
bad init) does not abort the worker; we fall back to reinstall or
the torch path.
"""
"""Catch any exception (not just ImportError) so a broken native lib doesn't abort the worker."""
try:
import fla.modules # noqa: F401
import fla.ops.gated_delta_rule # noqa: F401
@ -346,16 +320,7 @@ def _flash_linear_attention_importable() -> bool:
def _flash_linear_attention_current(already_importable: bool | None = None) -> bool:
"""True iff FLA is importable AND meets the PR's pinned versions.
A user with an older `flash-linear-attention` (e.g. 0.4.x) on the
venv would import fine but lack the gated_delta_rule kernels we
expect. Version-checking before short-circuiting forces a reinstall
to the pin.
`already_importable=True` lets the caller skip the import probe
when it has just performed it (call-count stability for tests).
"""
"""True iff FLA imports AND is at the pinned version (older FLA lacks gated_delta_rule kernels)."""
if already_importable is None:
already_importable = _flash_linear_attention_importable()
if not already_importable:
@ -378,25 +343,7 @@ def _flash_linear_attention_current(already_importable: bool | None = None) -> b
def _ensure_flash_linear_attention_unconditional(event_queue: Any) -> bool:
"""Install ``flash-linear-attention`` + ``fla-core`` unconditionally.
Returns True iff FLA is importable AT THE PINNED VERSION post-call;
False otherwise (skipped, install failed, deep import broken, etc).
Callers use the return value to decide whether to chain into
tilelang or short-circuit cleanly.
This is the body of the installer with the model-name substring gate
removed: the caller has already proven (via the runtime hook on
``is_flash_linear_attention_available``) that the loaded model
actually needs FLA, so we just need to make the import work.
Pinned ``flash-linear-attention``, ``fla-core`` and the runtime
deps we explicitly want (``einops``, ``packaging``, ``triton``)
are installed with ``--no-deps`` so pip never silently upgrades
torch from fla-core's ``torch>=2.7.0`` requirement.
Set ``UNSLOTH_STUDIO_SKIP_FLA_INSTALL=1`` to bypass entirely.
"""
"""Install pinned FLA + fla-core with --no-deps. Returns True iff importable post-call."""
if os.getenv(_FLA_SKIP_ENV) == "1":
return False
if sys.version_info < _FLA_MIN_PYTHON:
@ -419,8 +366,8 @@ def _ensure_flash_linear_attention_unconditional(event_queue: Any) -> bool:
)
return False
# Probe once; reuse the result for short-circuit AND
# --force-reinstall decision so call count stays stable.
# Probe once; reuse result so the --force-reinstall decision and the short-circuit
# share the same call count (stable for tests).
already_importable = _flash_linear_attention_importable()
if already_importable and _flash_linear_attention_current(already_importable = True):
logger.info("flash-linear-attention already importable at the pinned version")
@ -434,20 +381,15 @@ def _ensure_flash_linear_attention_unconditional(event_queue: Any) -> bool:
),
)
# Install fla-core's required non-torch runtime deps explicitly
# because `--no-deps` suppresses them. Without einops/packaging
# (and triton, on minimal torch builds), `import fla.modules`
# raises ModuleNotFoundError at runtime.
# `--no-deps` blocks the silent torch upgrade; we bring the non-torch runtime deps in by hand.
specs = [
*_FLA_RUNTIME_DEPS,
f"fla-core=={_FLA_CORE_PACKAGE_VERSION}",
f"flash-linear-attention=={_FLA_PACKAGE_VERSION}",
]
extra_args = ["--no-deps"]
# If an older FLA is importable we must force-reinstall to get the pinned
# version. Without --force-reinstall pip would see fla-core present and
# do nothing; --no-deps still applies so torch stays untouched.
if already_importable:
# Older FLA already imported; pip skips reinstall without this flag.
extra_args.append("--force-reinstall")
if shutil.which("uv"):
@ -496,9 +438,7 @@ def _ensure_flash_linear_attention_unconditional(event_queue: Any) -> bool:
)
return False
# Verify the install actually produced importable modules. Catches
# the case where pip exits 0 but a transitive runtime dep we did
# not list is missing.
# pip can exit 0 with a missing transitive runtime dep; verify the import.
if not _flash_linear_attention_importable():
_send_status(
event_queue,
@ -511,13 +451,7 @@ def _ensure_flash_linear_attention_unconditional(event_queue: Any) -> bool:
def _ensure_flash_linear_attention(event_queue: Any, model_name: str) -> None:
"""Legacy substring-gated installer.
Kept for the ``UNSLOTH_STUDIO_SKIP_FAST_PATH_HOOKS=1`` opt-out path,
where the runtime hook on ``is_flash_linear_attention_available`` is
disabled and we fall back to a model-name match. The hook is the
primary gate in normal operation.
"""
"""Legacy model-name-gated FLA install, used when UNSLOTH_STUDIO_SKIP_FAST_PATH_HOOKS=1."""
if not _model_wants_tilelang(model_name):
return
_ensure_flash_linear_attention_unconditional(event_queue)
@ -551,36 +485,49 @@ def _ensure_mamba_ssm(event_queue: Any, model_name: str) -> None:
)
# Linear-attention models that benefit from FLA's TileLang backend.
# FLA dispatches `chunk_bwd_dqkwg` / `parallel_attn_fwd` / `parallel_attn_bwd`
# to TileLang when both `tilelang` and `apache-tvm-ffi` are importable;
# this gives ~26% additional speedup on Qwen3.5-2B-Vision on B200 in our
# bench, on top of the FLA-Triton fast path.
#
# Restricted to GDN architectures (Qwen3.5 family). True SSM models
# (Nemotron-H, Falcon-H1, Granite-H, LFM2) take their own path and do not
# go through FLA's gated_delta_rule, so we do NOT install tilelang for them.
_TILELANG_MODEL_SUBSTRINGS = (
"qwen3.5",
"qwen3_5",
"qwen3.6",
"qwen3_6",
"qwen3-next",
"qwen3_next",
)
# Auto-derived from installed transformers: model_types whose modeling_*.py imports `from fla.*`.
# Cached per process. Empty when transformers can't be inspected -> we skip tilelang pre-install
# (the FLA Triton path still runs via the runtime hook).
_TRANSFORMERS_FLA_MODEL_TYPES_CACHE: frozenset[str] | None = None
_MODEL_NAME_SEP_CHARS = ("-", ".", "/", " ")
def _discover_fla_model_types() -> frozenset[str]:
"""Model_types in the installed transformers whose modeling file imports `from fla.*`."""
global _TRANSFORMERS_FLA_MODEL_TYPES_CACHE
if _TRANSFORMERS_FLA_MODEL_TYPES_CACHE is not None:
return _TRANSFORMERS_FLA_MODEL_TYPES_CACHE
found: set[str] = set()
try:
import transformers
models_root = Path(transformers.__file__).parent / "models"
for modeling in models_root.glob("*/modeling_*.py"):
try:
src = modeling.read_text(encoding = "utf-8", errors = "ignore")
except OSError:
continue
if "from fla." in src:
found.add(modeling.parent.name)
except Exception as exc:
logger.debug("FLA model-type discovery skipped: %s", exc)
_TRANSFORMERS_FLA_MODEL_TYPES_CACHE = frozenset(found)
return _TRANSFORMERS_FLA_MODEL_TYPES_CACHE
def _model_wants_tilelang(model_name: str) -> bool:
"""True iff model_name normalizes to contain a discovered FLA model_type."""
types = _discover_fla_model_types()
if not types:
return False
name = model_name.lower()
return any(sub in name for sub in _TILELANG_MODEL_SUBSTRINGS)
for sep in _MODEL_NAME_SEP_CHARS:
name = name.replace(sep, "_")
return any(t in name for t in types)
def _installed_tvm_ffi_version() -> str | None:
"""Return ``apache-tvm-ffi`` version if importable, else None.
Used to decide whether an in-place install needs to force a reinstall
because the existing version is on the broken list.
"""
"""Installed apache-tvm-ffi version, or None if missing/unimportable."""
try:
from importlib.metadata import version as _pkg_version
@ -590,7 +537,7 @@ def _installed_tvm_ffi_version() -> str | None:
def _tilelang_importable() -> bool:
"""Best-effort tilelang import probe; catches broader than ImportError."""
"""Catch any exception (not just ImportError) so a broken native lib doesn't abort the worker."""
try:
import tilelang # noqa: F401
import tvm_ffi # noqa: F401
@ -605,18 +552,7 @@ def _tilelang_importable() -> bool:
def _torch_has_hip() -> bool:
"""True iff the installed torch is a HIP / ROCm build.
We check `torch.version.hip` (non-None on ROCm wheels). This is the
reliable signal even on x86_64 Linux Strix Halo / MI300, where
`sys.platform` and `platform.machine()` look identical to a CUDA box.
Importing torch here is acceptable in the worker subprocess context:
the next step after kernel installers is the model load, which
imports torch anyway. We swallow import errors so a missing torch
(extremely unusual at this point) is treated as "not HIP" and the
rest of the gate stack handles it.
"""
"""True iff torch is a ROCm build; `torch.version.hip` is the only reliable signal on x86_64 ROCm."""
try:
import torch as _torch
@ -626,18 +562,9 @@ def _torch_has_hip() -> bool:
def _tilelang_platform_supported() -> bool:
"""True iff the current platform has a usable tilelang 0.1.8 backend.
"""True iff a tilelang 0.1.8 wheel will load: Linux x86_64/aarch64, non-HIP torch.
tilelang publishes manylinux x86_64/aarch64 and macOS arm64 wheels
plus a 93MB sdist; we never want the sdist on a Studio worker, so
we restrict to Linux x86_64/aarch64 explicitly.
Excludes HIP / ROCm torch builds: tilelang 0.1.8 has no HIP GEMM
instruction, so `_select_gemm_instruction` raises `Unsupported
target for gemm: hip` mid-compile during Qwen3.5 GDN backward.
Reported by h34v3nzc0dex on Strix Halo (gfx1151, ROCm 7.13). The
pip wheel installs fine and imports cleanly, but FLA's TileLang
dispatcher then crashes at first training step. See PR 5434.
HIP excluded because tilelang 0.1.8 has no HIP GEMM instruction and crashes mid-backward.
"""
import platform as _platform
@ -651,14 +578,14 @@ def _tilelang_platform_supported() -> bool:
def _pip_install_cmd(*args: str) -> list[str]:
"""Build a `uv pip install` or `python -m pip install` invocation."""
"""`uv pip install` if uv is on PATH, else `python -m pip install`."""
if shutil.which("uv"):
return ["uv", "pip", "install", "--python", sys.executable, *args]
return [sys.executable, "-m", "pip", "install", *args]
def _run_pip(cmd: list[str], event_queue: Any, label: str) -> bool:
"""Run a pip install command and report success/failure via status."""
"""Run a pip install and surface success/failure via status events."""
try:
result = _sp.run(
cmd,
@ -681,24 +608,11 @@ def _run_pip(cmd: list[str], event_queue: Any, label: str) -> bool:
def _ensure_tilelang_backend_unconditional(event_queue: Any) -> bool:
"""Install ``tilelang`` + pinned ``apache-tvm-ffi`` unconditionally.
"""Install pinned tilelang + apache-tvm-ffi; two-step repair if a broken tvm-ffi is present.
Returns True iff tilelang + tvm_ffi are importable post-call.
Called from the FLA hook because tilelang only matters once FLA is
active; the substring gate is gone here. Pre-existing platform,
Python, and skip-env guards remain.
Repair semantics for a broken `apache-tvm-ffi` (0.1.10/0.1.11):
step 1: ``--force-reinstall --no-deps apache-tvm-ffi==0.1.9``
(downgrades ONLY the broken package; does NOT touch
torch or the CUDA stack)
step 2: regular install for ``tilelang`` + ``apache-tvm-ffi``
resolves any missing transitive deps (z3-solver,
ml-dtypes) without --force-reinstall, so it never
replaces torch with a different CUDA build either.
Set ``UNSLOTH_STUDIO_SKIP_TILELANG_INSTALL=1`` to bypass.
Returns True iff both import post-call. Step 1 surgically downgrades a broken tvm-ffi
with --force-reinstall --no-deps so torch / CUDA stay untouched; step 2 is a regular
install for missing transitive deps. Bypass via UNSLOTH_STUDIO_SKIP_TILELANG_INSTALL=1.
"""
if os.getenv(_TILELANG_SKIP_ENV) == "1":
return False
@ -727,10 +641,7 @@ def _ensure_tilelang_backend_unconditional(event_queue: Any) -> bool:
logger.info("tilelang + apache-tvm-ffi already installed")
return True
# Step 1: if a broken tvm-ffi is present, surgically downgrade it
# without --no-deps' usual deps-only-once semantics. --no-deps here
# protects torch and the CUDA stack from being uninstalled by
# --force-reinstall pulling in apache-tvm-ffi's full dep graph.
# Step 1: --no-deps keeps --force-reinstall from touching torch/CUDA via the dep graph.
if needs_repair:
logger.info(
"Forcing apache-tvm-ffi downgrade: %s is on the broken list",
@ -752,10 +663,7 @@ def _ensure_tilelang_backend_unconditional(event_queue: Any) -> bool:
if not _run_pip(repair_cmd, event_queue, "TileLang backend repair"):
return False
# Step 2: regular dependency-resolving install so missing transitive
# deps (z3-solver, ml-dtypes, ...) get pulled in. Without
# --force-reinstall pip is a no-op for already-correct packages,
# so this never replaces torch.
# Step 2: regular install pulls in transitive deps (z3-solver, ml-dtypes) without touching torch.
_send_status(
event_queue,
(
@ -772,8 +680,7 @@ def _ensure_tilelang_backend_unconditional(event_queue: Any) -> bool:
if not _run_pip(install_cmd, event_queue, "TileLang backend"):
return False
# Verify imports succeed; pip can return 0 while a native library
# (libz3.so, ...) is missing for the runtime load.
# pip can exit 0 while a native lib (libz3.so) is missing; verify the import.
if not _tilelang_importable():
_send_status(
event_queue,
@ -792,54 +699,23 @@ def _ensure_tilelang_backend(event_queue: Any, model_name: str) -> None:
_ensure_tilelang_backend_unconditional(event_queue)
# ──────────────────────────────────────────────────────────────────────
# Runtime hook on transformers' fast-path availability gates.
#
# transformers' qwen3_5 / qwen3_5_moe / qwen3_next modeling files do
#
# if is_causal_conv1d_available():
# from causal_conv1d import causal_conv1d_fn, causal_conv1d_update
# if is_flash_linear_attention_available():
# from fla.modules import FusedRMSNormGated
# from fla.ops.gated_delta_rule import ...
#
# at MODULE IMPORT TIME. If the gate returns False then, the fast-path
# symbols are bound to None and the model falls back to a pure-Python
# torch loop forever in that process. We wrap the gates so the first
# call (always at modeling import time, because the worker has not
# loaded a model yet) drives the matching install synchronously and
# returns True post-install. That way:
#
# - Any model whose architecture actually queries the gates triggers
# the install, regardless of its name.
# - Models that never query the gates (Llama, Gemma, dense Qwen, …)
# never pay the install cost.
#
# This supersedes the substring-based `_model_wants_tilelang` check
# for these two kernels. Set `UNSLOTH_STUDIO_SKIP_FAST_PATH_HOOKS=1`
# to fall back to the legacy substring path.
# ──────────────────────────────────────────────────────────────────────
# ── Fast-path hooks ──
# Wrap transformers' is_{flash_linear_attention,causal_conv1d}_available so the first call
# (at modeling import time) drives the install. Any model that queries the gate gets the
# install; models that never query it (Llama, Gemma, dense Qwen) pay nothing.
# UNSLOTH_STUDIO_SKIP_FAST_PATH_HOOKS=1 falls back to the legacy substring path.
def _rebind_in_already_imported_modules(
*, attr_name: str, old_obj: Any, new_obj: Any
) -> int:
"""Replace `attr_name` in every loaded module that bound `old_obj`.
"""Rebind `attr_name -> new_obj` in every module that already imported `old_obj`.
Modeling files do `from transformers.utils.import_utils import
is_flash_linear_attention_available`, which creates a local binding
in the importing module. Reassigning the symbol on
`transformers.utils.import_utils` does NOT reach those bindings.
We use `module.__dict__.get(attr_name)` (NOT `getattr(mod, ...)`)
because transformers' lazy module aliases override `__getattr__` and
`getattr(mod, name)` will trigger an "Accessing X from .models..."
advisory warning AND can materialise lazy imports we have no
interest in. The dict lookup only sees real module-level bindings.
`from X import Y` creates a local binding that reassigning X.Y won't reach.
Uses `__dict__.get` (not `getattr`) to skip lazy `__getattr__` aliases.
"""
count = 0
missing = object()
# snapshot keys to avoid mutating during iteration
for mod_name, mod in list(sys.modules.items()):
if mod is None:
continue
@ -857,51 +733,19 @@ def _rebind_in_already_imported_modules(
def _install_fast_path_hooks(event_queue: Any, model_name: str) -> None:
"""Wrap `is_flash_linear_attention_available` and
`is_causal_conv1d_available` so the first call drives the matching
install if the underlying package is missing.
"""Hook transformers' is_*_available gates so the first call drives the install.
The wrapper:
1. Clears the original `@lru_cache` so the underlying check is
actually re-evaluated.
2. Calls the original. If it returns True, no work to do other
than the post-available action (e.g. tilelang repair).
3. If False, calls `install_fn(event_queue) -> bool`. The returned
bool is the authoritative post-install availability (NOT a
re-call of `original()`, which can lie when pip exited 0 but
deep imports are broken).
4. Calls `post_available_fn(event_queue)` if available, so
tilelang's broken-version repair runs even when FLA was
already True.
`model_name` is threaded through so the FLA install can gate
tilelang on `_model_wants_tilelang(model_name)`. tilelang is a
Qwen3.5-family optimisation; non-Qwen FLA-using architectures
(OLMo-Hybrid, future GDN models) only want FLA itself.
Idempotent: subsequent calls short-circuit on an `installed` flag.
Set `UNSLOTH_STUDIO_SKIP_FAST_PATH_HOOKS=1` to bypass.
Idempotent. UNSLOTH_STUDIO_SKIP_FAST_PATH_HOOKS=1 falls back to the substring gate.
"""
if os.getenv(_FAST_PATH_HOOKS_SKIP_ENV) == "1":
logger.info("Fast-path hooks disabled via env; using substring fallback")
return
# Defensive: on HIP/ROCm torch builds, FLA's TileLang backend (when
# tilelang is installed for any reason — e.g. a stale CUDA env that
# was reused for ROCm) crashes mid-backward with
# "Unsupported target for gemm: hip" inside
# `tilelang.tileop.gemm._select_gemm_instruction`. The install gate
# in `_ensure_tilelang_backend_unconditional` prevents NEW installs
# on HIP; this env-var setdefault disables FLA's TileLang dispatch
# for already-installed tilelang too. Users can override by setting
# FLA_TILELANG=1 explicitly. Reported by h34v3nzc0dex on Strix Halo.
# On HIP torch, even already-installed tilelang crashes FLA's TileLang dispatch.
# User can override with FLA_TILELANG=1.
if _torch_has_hip() and os.environ.get("FLA_TILELANG") is None:
os.environ["FLA_TILELANG"] = "0"
logger.info(
"HIP/ROCm torch detected; setting FLA_TILELANG=0 to keep "
"FLA on the safe Triton path (tilelang 0.1.8 has no HIP "
"GEMM backend)"
)
logger.info("HIP/ROCm torch detected; setting FLA_TILELANG=0 (no HIP GEMM in tilelang 0.1.8)")
try:
from transformers.utils import import_utils as _iu
@ -923,11 +767,8 @@ def _install_fast_path_hooks(event_queue: Any, model_name: str) -> None:
def wrapper() -> bool:
if state["installed"]:
return original()
# Clear the lru_cache so the underlying check re-evaluates
# after any pre-hook calls (defensive, the worker subprocess
# is freshly spawned so this should be a no-op).
try:
original.cache_clear()
original.cache_clear() # defensive; worker subprocess is fresh
except AttributeError:
pass
ok = original()
@ -935,86 +776,47 @@ def _install_fast_path_hooks(event_queue: Any, model_name: str) -> None:
if not ok:
ran_install = True
logger.info("Hook fired for %s; triggering install", gate_name)
_send_status(
event_queue,
f"Hook fired for {gate_name}; installing kernel...",
)
_send_status(event_queue, f"Hook fired for {gate_name}; installing kernel...")
try:
install_result = install_fn(event_queue)
ok = bool(install_result)
ok = bool(install_fn(event_queue))
except Exception as exc:
logger.warning(
"Install fired by %s hook raised: %s; continuing on torch fallback",
gate_name,
exc,
)
logger.warning("%s install raised: %s; falling back to torch", gate_name, exc)
ok = False
logger.info(
"Hook for %s completed; post-install availability=%s",
gate_name,
ok,
)
# post_available_fn handles edge cases that ONLY occur on
# the gate-was-already-True path (e.g. tilelang missing
# while FLA is already importable, or apache-tvm-ffi on
# the broken-versions list while FLA otherwise works).
# If install_fn ran, it already chained the matching
# follow-up install (`_fla_install` installs tilelang too),
# so running post_available_fn would double-install.
logger.info("%s hook done; available=%s", gate_name, ok)
# post_available_fn handles "gate already True but ancillary kernel broken" (e.g. tilelang
# missing while FLA imports fine); skip when install_fn already chained the follow-up.
if ok and not ran_install and post_available_fn is not None:
try:
post_available_fn(event_queue)
except Exception as exc:
logger.warning(
"%s post-available step raised: %s; continuing",
gate_name,
exc,
)
logger.warning("%s post-available step raised: %s; continuing", gate_name, exc)
state["installed"] = True
return ok
wrapper.__wrapped__ = original # type: ignore[attr-defined]
# Re-expose cache_clear so callers that introspect it still work.
wrapper.cache_clear = getattr(original, "cache_clear", lambda: None) # type: ignore[attr-defined]
return wrapper
def _fla_install(eq: Any) -> bool:
# FLA without tilelang gets ~2.35x speedup; tilelang adds ~26%.
# tilelang is a Qwen3.5-family optimisation only; non-Qwen FLA
# users (OLMo-Hybrid, ...) skip it. Order: install FLA first,
# gate tilelang on (FLA succeeded) AND (model wants tilelang).
fla_ok = _ensure_flash_linear_attention_unconditional(eq)
if not fla_ok:
logger.info(
"FLA install did not produce an importable runtime; "
"skipping TileLang backend"
)
# FLA alone ~2.35x; +tilelang adds ~26%. tilelang is GDN-only (Qwen3.5 family).
if not _ensure_flash_linear_attention_unconditional(eq):
logger.info("FLA install did not produce an importable runtime; skipping TileLang")
return False
if _model_wants_tilelang(model_name):
_ensure_tilelang_backend_unconditional(eq)
else:
logger.info(
"Model %r does not match the TileLang allowlist; "
"skipping TileLang backend (FLA Triton path is sufficient)",
model_name,
)
logger.info("Model %r outside TileLang allowlist; FLA Triton path is sufficient", model_name)
return True
def _fla_post_available(eq: Any) -> None:
# Runs when FLA was already importable (gate returned True
# without triggering install). If the model wants tilelang and
# tilelang is missing or `apache-tvm-ffi` is on the broken
# version list, the unconditional installer will repair it.
# FLA already imports; repair tilelang if missing or on the broken tvm-ffi list.
if not _model_wants_tilelang(model_name):
return
existing_tvm = _installed_tvm_ffi_version()
needs_repair = existing_tvm in _TVM_FFI_BROKEN_VERSIONS
if not needs_repair and _tilelang_importable():
if _installed_tvm_ffi_version() not in _TVM_FFI_BROKEN_VERSIONS and _tilelang_importable():
return
_ensure_tilelang_backend_unconditional(eq)
def _causal_conv1d_install(eq: Any) -> bool:
# Reuse the existing wheel-first installer.
ok = _install_package_wheel_first(
event_queue = eq,
import_name = "causal_conv1d",
@ -1029,39 +831,20 @@ def _install_fast_path_hooks(event_queue: Any, model_name: str) -> None:
)
return bool(ok)
rebound_total = 0
for gate_name, install_fn, post_fn in (
(
"is_flash_linear_attention_available",
_fla_install,
_fla_post_available,
),
("is_flash_linear_attention_available", _fla_install, _fla_post_available),
("is_causal_conv1d_available", _causal_conv1d_install, None),
):
original = getattr(_iu, gate_name, None)
if original is None:
logger.info(
"transformers.utils.import_utils.%s missing; skipping that hook",
gate_name,
)
logger.info("%s missing on transformers.utils.import_utils; skipping hook", gate_name)
continue
wrapped = _make_wrapper(original, install_fn, gate_name, post_fn)
setattr(_iu, gate_name, wrapped)
rebound = _rebind_in_already_imported_modules(
attr_name = gate_name, old_obj = original, new_obj = wrapped
)
rebound_total += rebound
logger.info(
"Installed fast-path hook on %s (rebound %d modules)",
gate_name,
rebound,
)
if rebound_total > 0:
logger.info(
"Rebound %d pre-existing module-level references to fast-path gates",
rebound_total,
)
logger.info("Installed fast-path hook on %s (rebound %d modules)", gate_name, rebound)
def _should_try_runtime_flash_attn_install(max_seq_length: int) -> bool:

View file

@ -264,6 +264,12 @@ def test_flash_linear_attention_matches_full_qwen3_family(monkeypatch):
monkeypatch.setattr(worker._sp, "run", run_mock)
_force_missing_fla_imports(monkeypatch)
monkeypatch.setattr(worker, "_send_status", lambda *a, **k: None)
# Hermetic discovery: pretend installed transformers ships all the Qwen GDN families.
monkeypatch.setattr(
worker,
"_discover_fla_model_types",
lambda: frozenset({"qwen3_5", "qwen3_5_moe", "qwen3_6", "qwen3_next"}),
)
for name in (
"unsloth/Qwen3.5-2B",
@ -903,22 +909,21 @@ def test_hook_skips_when_import_utils_unavailable(monkeypatch):
def test_substring_fallback_unchanged_when_hook_skipped(monkeypatch):
"""With the hook disabled, the orchestration falls back to the
substring path. Confirm _ensure_flash_linear_attention(model_name)
still gates on model name as before."""
"""Hook disabled -> legacy gate falls back to auto-discovered model types."""
install_mock = mock.Mock()
monkeypatch.setattr(
worker, "_ensure_flash_linear_attention_unconditional", install_mock
)
monkeypatch.setattr(
worker, "_discover_fla_model_types", lambda: frozenset({"qwen3_5"})
)
monkeypatch.setenv(worker._FAST_PATH_HOOKS_SKIP_ENV, "1")
# Qwen3.5 model triggers install.
worker._ensure_flash_linear_attention(
event_queue = [], model_name = "unsloth/Qwen3.5-2B"
)
assert install_mock.call_count == 1
# Llama doesn't.
worker._ensure_flash_linear_attention(
event_queue = [], model_name = "meta-llama/Llama-3.1-8B"
)
@ -1324,3 +1329,154 @@ def test_install_fast_path_hooks_does_not_set_fla_tilelang_on_cuda(monkeypatch):
)
assert _os.environ.get("FLA_TILELANG") is None
# ───────────────────────────────────────────────────────────────────
# Auto-discovery of FLA model_types from the installed transformers
# ───────────────────────────────────────────────────────────────────
def _make_fake_transformers_tree(tmp_path, fla_types: list[str], non_fla_types: list[str]):
"""Lay out a tmp dir as `transformers/models/{type}/modeling_{type}.py`."""
pkg = tmp_path / "transformers"
models = pkg / "models"
models.mkdir(parents = True)
(pkg / "__init__.py").write_text("")
for t in fla_types:
d = models / t
d.mkdir()
(d / f"modeling_{t}.py").write_text(
"from ...utils.import_utils import is_flash_linear_attention_available\n"
"if is_flash_linear_attention_available():\n"
" from fla.modules import FusedRMSNormGated\n"
" from fla.ops.gated_delta_rule import chunk_gated_delta_rule\n"
)
for t in non_fla_types:
d = models / t
d.mkdir()
(d / f"modeling_{t}.py").write_text("class Foo: pass\n")
return pkg
def _reset_fla_cache(monkeypatch):
monkeypatch.setattr(worker, "_TRANSFORMERS_FLA_MODEL_TYPES_CACHE", None)
def test_discover_fla_model_types_returns_only_fla_users(tmp_path, monkeypatch):
pkg = _make_fake_transformers_tree(
tmp_path,
fla_types = ["qwen3_5", "qwen3_5_moe", "qwen3_next"],
non_fla_types = ["llama", "gpt2", "mistral"],
)
fake = mock.MagicMock(__file__ = str(pkg / "__init__.py"))
monkeypatch.setitem(sys.modules, "transformers", fake)
_reset_fla_cache(monkeypatch)
result = worker._discover_fla_model_types()
assert result == frozenset({"qwen3_5", "qwen3_5_moe", "qwen3_next"})
assert "llama" not in result
assert "gpt2" not in result
def test_discover_fla_model_types_caches_across_calls(tmp_path, monkeypatch):
pkg = _make_fake_transformers_tree(
tmp_path, fla_types = ["qwen3_5"], non_fla_types = []
)
fake = mock.MagicMock(__file__ = str(pkg / "__init__.py"))
monkeypatch.setitem(sys.modules, "transformers", fake)
_reset_fla_cache(monkeypatch)
from pathlib import Path as _Path
read_calls = [0]
real_read = _Path.read_text
def counting_read(self, *a, **kw):
read_calls[0] += 1
return real_read(self, *a, **kw)
monkeypatch.setattr(_Path, "read_text", counting_read)
first = worker._discover_fla_model_types()
after_first = read_calls[0]
second = worker._discover_fla_model_types()
assert first == second
assert read_calls[0] == after_first # cache hit: no extra disk reads
def test_discover_fla_model_types_handles_missing_transformers(monkeypatch):
_reset_fla_cache(monkeypatch)
real_import = builtins.__import__
def fake_import(name, globals = None, locals = None, fromlist = (), level = 0):
if name == "transformers":
raise ImportError("transformers not installed")
return real_import(name, globals, locals, fromlist, level)
monkeypatch.setattr(builtins, "__import__", fake_import)
result = worker._discover_fla_model_types()
assert result == frozenset()
def test_discover_fla_model_types_handles_unreadable_file(tmp_path, monkeypatch):
pkg = _make_fake_transformers_tree(
tmp_path, fla_types = ["qwen3_5"], non_fla_types = []
)
fake = mock.MagicMock(__file__ = str(pkg / "__init__.py"))
monkeypatch.setitem(sys.modules, "transformers", fake)
_reset_fla_cache(monkeypatch)
from pathlib import Path as _Path
real_read = _Path.read_text
def boom_read(self, *a, **kw):
if "modeling_qwen3_5.py" in str(self):
raise OSError("permission denied")
return real_read(self, *a, **kw)
monkeypatch.setattr(_Path, "read_text", boom_read)
result = worker._discover_fla_model_types()
assert result == frozenset() # unreadable file simply doesn't contribute
def test_model_wants_tilelang_handles_real_repo_names(monkeypatch):
monkeypatch.setattr(
worker,
"_discover_fla_model_types",
lambda: frozenset({"qwen3_5", "qwen3_5_moe", "qwen3_next"}),
)
cases = [
("unsloth/Qwen3.5-2B", True),
("Qwen/Qwen3.5-MoE-A3B", True),
("mlx-community/qwen3-next-80b", True),
("unsloth/qwen3_5_moe_a3b_lora", True),
("meta-llama/Llama-3.1-8B", False),
("nvidia/Nemotron-H-4B", False),
("mistralai/Mistral-7B-v0.3", False),
("", False),
]
for name, expected in cases:
assert worker._model_wants_tilelang(name) is expected, name
def test_model_wants_tilelang_empty_when_transformers_has_no_fla(monkeypatch):
monkeypatch.setattr(worker, "_discover_fla_model_types", lambda: frozenset())
assert worker._model_wants_tilelang("unsloth/Qwen3.5-2B") is False
assert worker._model_wants_tilelang("meta-llama/Llama-3.1-8B") is False
def test_model_wants_tilelang_normalizes_separators(monkeypatch):
monkeypatch.setattr(
worker, "_discover_fla_model_types", lambda: frozenset({"qwen3_next"})
)
for variant in (
"qwen3-next",
"Qwen3.Next",
"Qwen/Qwen3 Next",
"anyone/qwen3_next",
"qwen3.next-80b",
):
assert worker._model_wants_tilelang(variant) is True, variant