diff --git a/studio/backend/core/training/worker.py b/studio/backend/core/training/worker.py index 7111151fb3..7d489d8254 100644 --- a/studio/backend/core/training/worker.py +++ b/studio/backend/core/training/worker.py @@ -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: diff --git a/studio/backend/tests/test_training_worker_flash_attn.py b/studio/backend/tests/test_training_worker_flash_attn.py index 4df38abf0a..9927829c7a 100644 --- a/studio/backend/tests/test_training_worker_flash_attn.py +++ b/studio/backend/tests/test_training_worker_flash_attn.py @@ -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