studio: hook transformers' fast-path gates for just-in-time FLA + causal-conv1d install
The substring-based detection in this PR (`_model_wants_tilelang` /
`_model_wants_causal_conv1d`) is brittle: it depends on what the user
typed for the model name, not on what the architecture actually needs.
Users typing custom model paths, future Qwen3.7 / non-Qwen GDN
architectures, and any model whose author renamed it would silently
fall back to the torch loop.
The correct signal is the one transformers itself uses to gate the
fast path. `transformers/models/qwen3_5_moe/modeling_qwen3_5_moe.py`
does at module import time:
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 (
chunk_gated_delta_rule, fused_recurrent_gated_delta_rule,
)
Wrap both gates so the first call (always at modeling import, before
any forward pass) installs the matching kernel synchronously and
delegates to the original function. Any model whose architecture
queries those gates auto-triggers the install; models that never
query them (Llama, Gemma, dense Qwen, ...) never pay the cost.
Mechanics:
- Split `_ensure_flash_linear_attention` and `_ensure_tilelang_backend`
into `_unconditional` variants (no substring gate, retains python
/ torch / platform / skip-env guards) plus thin substring wrappers
used by the legacy fallback path.
- New `_install_fast_path_hooks(event_queue)` patches both gates on
`transformers.utils.import_utils` AND sweeps `sys.modules` so any
modeling file that already did `from ... import is_X` sees the
wrapper (the local binding survives a module-level reassignment).
- Wrappers clear the original's `lru_cache` before delegating, install
on False, re-check, and short-circuit on subsequent calls.
- Set `UNSLOTH_STUDIO_SKIP_FAST_PATH_HOOKS=1` to fall back to the
substring path.
Verified end-to-end against `transformers.models.qwen3_5_moe`:
PRE_STATE fla=False tilelang=False causal_conv1d=False
HOOK_INSTALLED
Hook fired for is_causal_conv1d_available; installing kernel...
Installing prebuilt causal-conv1d wheel...
Hook fired for is_flash_linear_attention_available; installing kernel...
Installing flash-linear-attention==0.5.0 (with fla-core==0.5.0) for the fast path...
Installed flash-linear-attention for the FLA fast path
Installing TileLang backend (apache-tvm-ffi==0.1.9, tilelang==0.1.8)...
Installed TileLang backend for FLA fast path
MODELING_IMPORT_OK
FAST_PATH_SYMBOLS {"chunk_gated_delta_rule": true,
"fused_recurrent_gated_delta_rule": true,
"FusedRMSNormGated": true,
"causal_conv1d_fn": true,
"causal_conv1d_update": true}
POST_STATE fla=True tilelang=True causal_conv1d=True
Adds 9 new tests covering: install-on-False, skip-on-True, idempotency,
install-failure handling, env-disable, lru_cache clear, sys.modules
rebind, missing-transformers fallback, substring fallback. Total
test count is now 36 (was 27).
This commit is contained in:
parent
66dface7d7
commit
6ce495a42d
2 changed files with 530 additions and 29 deletions
|
|
@ -85,6 +85,10 @@ _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"
|
||||
|
||||
|
||||
def _model_wants_causal_conv1d(model_name: str) -> bool:
|
||||
|
|
@ -341,16 +345,13 @@ def _flash_linear_attention_importable() -> bool:
|
|||
return False
|
||||
|
||||
|
||||
def _ensure_flash_linear_attention(event_queue: Any, model_name: str) -> None:
|
||||
"""Install ``flash-linear-attention`` + ``fla-core`` for Qwen3.5 family.
|
||||
def _ensure_flash_linear_attention_unconditional(event_queue: Any) -> None:
|
||||
"""Install ``flash-linear-attention`` + ``fla-core`` unconditionally.
|
||||
|
||||
Qwen3.5 / Qwen3.6 / Qwen3-Next gate their transformers fast path on
|
||||
FLA's ``chunk_gated_delta_rule`` / ``fused_recurrent_gated_delta_rule``
|
||||
being importable. Without FLA the path falls back to a pure-Python
|
||||
torch loop (~2.35x slower in our Qwen3.5-2B-Vision bench).
|
||||
|
||||
True SSM families (Nemotron-H, Falcon-H1, Granite-H, LFM2) take the
|
||||
mamba_ssm path and never call FLA's GDN kernels, so we skip them.
|
||||
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``)
|
||||
|
|
@ -361,8 +362,6 @@ def _ensure_flash_linear_attention(event_queue: Any, model_name: str) -> None:
|
|||
"""
|
||||
if os.getenv(_FLA_SKIP_ENV) == "1":
|
||||
return
|
||||
if not _model_wants_tilelang(model_name):
|
||||
return
|
||||
if sys.version_info < _FLA_MIN_PYTHON:
|
||||
logger.info(
|
||||
"Skipping flash-linear-attention install: requires Python >= %d.%d, have %s",
|
||||
|
|
@ -463,6 +462,19 @@ def _ensure_flash_linear_attention(event_queue: Any, model_name: str) -> None:
|
|||
logger.info("Installed flash-linear-attention for the FLA fast path")
|
||||
|
||||
|
||||
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.
|
||||
"""
|
||||
if not _model_wants_tilelang(model_name):
|
||||
return
|
||||
_ensure_flash_linear_attention_unconditional(event_queue)
|
||||
|
||||
|
||||
_SSM_MODEL_SUBSTRINGS = (
|
||||
"nemotron_h",
|
||||
"nemotron-h",
|
||||
|
|
@ -558,8 +570,12 @@ def _tilelang_platform_supported() -> bool:
|
|||
return _platform.machine().lower() in _TILELANG_SUPPORTED_LINUX_MACHINES
|
||||
|
||||
|
||||
def _ensure_tilelang_backend(event_queue: Any, model_name: str) -> None:
|
||||
"""Install ``tilelang`` + pinned ``apache-tvm-ffi`` for FLA's TileLang backend.
|
||||
def _ensure_tilelang_backend_unconditional(event_queue: Any) -> None:
|
||||
"""Install ``tilelang`` + pinned ``apache-tvm-ffi`` unconditionally.
|
||||
|
||||
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.
|
||||
|
||||
The combined pin is important: `tilelang` declares
|
||||
``apache-tvm-ffi>=0.1.2,~=0.1.0`` which lets pip pull the latest 0.1.10/
|
||||
|
|
@ -567,15 +583,10 @@ def _ensure_tilelang_backend(event_queue: Any, model_name: str) -> None:
|
|||
Triton kernels on sm_100 (Blackwell). Pinning to 0.1.9 (the upper bound
|
||||
that ``mamba_ssm 2.3.2`` itself uses) avoids the regression.
|
||||
|
||||
Both packages are pure-Python wheels on PyPI; no wheel-matching dance
|
||||
is needed.
|
||||
|
||||
Set ``UNSLOTH_STUDIO_SKIP_TILELANG_INSTALL=1`` to bypass.
|
||||
"""
|
||||
if os.getenv(_TILELANG_SKIP_ENV) == "1":
|
||||
return
|
||||
if not _model_wants_tilelang(model_name):
|
||||
return
|
||||
if sys.version_info < _FLA_MIN_PYTHON:
|
||||
logger.info(
|
||||
"Skipping tilelang install: requires Python >= %d.%d, have %s",
|
||||
|
|
@ -685,6 +696,209 @@ def _ensure_tilelang_backend(event_queue: Any, model_name: str) -> None:
|
|||
logger.info("Installed TileLang backend for FLA fast path")
|
||||
|
||||
|
||||
def _ensure_tilelang_backend(event_queue: Any, model_name: str) -> None:
|
||||
"""Legacy substring-gated tilelang installer (opt-out path)."""
|
||||
if not _model_wants_tilelang(model_name):
|
||||
return
|
||||
_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.
|
||||
# ──────────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
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`.
|
||||
|
||||
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 sweep `sys.modules` for any module whose module-level dict
|
||||
contains `attr_name` bound to `old_obj` and rebind it to `new_obj`.
|
||||
Returns the number of bindings rewritten.
|
||||
"""
|
||||
count = 0
|
||||
# snapshot keys to avoid mutating during iteration
|
||||
for mod_name, mod in list(sys.modules.items()):
|
||||
if mod is None:
|
||||
continue
|
||||
try:
|
||||
existing = getattr(mod, attr_name, None)
|
||||
except Exception:
|
||||
continue
|
||||
if existing is old_obj:
|
||||
try:
|
||||
setattr(mod, attr_name, new_obj)
|
||||
count += 1
|
||||
except Exception as exc:
|
||||
logger.debug(
|
||||
"Could not rebind %s in %s: %s", attr_name, mod_name, exc
|
||||
)
|
||||
return count
|
||||
|
||||
|
||||
def _install_fast_path_hooks(event_queue: Any) -> 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.
|
||||
|
||||
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.
|
||||
3. If False, triggers `_ensure_*_unconditional(event_queue)` (FLA
|
||||
pulls tilelang too), clears the cache again, and re-checks.
|
||||
4. Returns the final boolean. If the install failed, returns
|
||||
False — same observable behaviour as before this PR, the
|
||||
model just falls back to the torch loop.
|
||||
|
||||
Idempotent: subsequent calls short-circuit on an `installed` flag.
|
||||
"""
|
||||
if os.getenv(_FAST_PATH_HOOKS_SKIP_ENV) == "1":
|
||||
logger.info("Fast-path hooks disabled via env; using substring fallback")
|
||||
return
|
||||
|
||||
try:
|
||||
from transformers.utils import import_utils as _iu
|
||||
except Exception as exc:
|
||||
logger.warning(
|
||||
"transformers.utils.import_utils not importable; skipping fast-path hooks: %s",
|
||||
exc,
|
||||
)
|
||||
return
|
||||
|
||||
def _make_wrapper(
|
||||
original: Callable[[], bool],
|
||||
install_fn: Callable[[Any], None],
|
||||
gate_name: str,
|
||||
) -> Callable[[], bool]:
|
||||
state = {"installed": False}
|
||||
|
||||
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()
|
||||
except AttributeError:
|
||||
pass
|
||||
ok = original()
|
||||
if not ok:
|
||||
logger.info("Hook fired for %s; triggering install", gate_name)
|
||||
_send_status(
|
||||
event_queue,
|
||||
f"Hook fired for {gate_name}; installing kernel...",
|
||||
)
|
||||
try:
|
||||
install_fn(event_queue)
|
||||
except Exception as exc:
|
||||
logger.warning(
|
||||
"Install fired by %s hook raised: %s; continuing on torch fallback",
|
||||
gate_name,
|
||||
exc,
|
||||
)
|
||||
# Re-check post-install.
|
||||
try:
|
||||
original.cache_clear()
|
||||
except AttributeError:
|
||||
pass
|
||||
ok = original()
|
||||
logger.info(
|
||||
"Hook for %s completed; post-install availability=%s",
|
||||
gate_name,
|
||||
ok,
|
||||
)
|
||||
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) -> None:
|
||||
# FLA without tilelang gets ~2.35x speedup; tilelang adds ~26%.
|
||||
# They pair up, so install both on the same trigger.
|
||||
_ensure_flash_linear_attention_unconditional(eq)
|
||||
_ensure_tilelang_backend_unconditional(eq)
|
||||
|
||||
def _causal_conv1d_install(eq: Any) -> None:
|
||||
# Reuse the existing wheel-first installer. It does its own
|
||||
# idempotency check via `__import__("causal_conv1d")`.
|
||||
_install_package_wheel_first(
|
||||
event_queue=eq,
|
||||
import_name="causal_conv1d",
|
||||
display_name="causal-conv1d",
|
||||
pypi_name="causal-conv1d",
|
||||
pypi_version=_CAUSAL_CONV1D_PACKAGE_VERSION,
|
||||
filename_prefix="causal_conv1d",
|
||||
release_tag=_CAUSAL_CONV1D_RELEASE_TAG,
|
||||
release_base_url=(
|
||||
"https://github.com/Dao-AILab/causal-conv1d/releases/download"
|
||||
),
|
||||
)
|
||||
|
||||
rebound_total = 0
|
||||
for gate_name, install_fn in (
|
||||
("is_flash_linear_attention_available", _fla_install),
|
||||
("is_causal_conv1d_available", _causal_conv1d_install),
|
||||
):
|
||||
original = getattr(_iu, gate_name, None)
|
||||
if original is None:
|
||||
logger.info(
|
||||
"transformers.utils.import_utils.%s missing; skipping that hook",
|
||||
gate_name,
|
||||
)
|
||||
continue
|
||||
wrapped = _make_wrapper(original, install_fn, gate_name)
|
||||
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,
|
||||
)
|
||||
|
||||
|
||||
def _should_try_runtime_flash_attn_install(max_seq_length: int) -> bool:
|
||||
if os.getenv(_FLASH_ATTN_SKIP_ENV) == "1":
|
||||
return False
|
||||
|
|
@ -1496,19 +1710,29 @@ def run_training_process(
|
|||
)
|
||||
|
||||
# ── 1b. Install fast-path kernel libraries for the chosen model.
|
||||
# Order:
|
||||
# 1) causal-conv1d (gates transformers' qwen3_5 / qwen3_next fast path)
|
||||
# 2) flash-linear-attention (the other half of that gate; without it
|
||||
# the conv kernel alone gives ~no measurable speedup)
|
||||
# 3) mamba-ssm (true SSM families only: Nemotron-H, Falcon-H1, etc.)
|
||||
# 4) tilelang + apache-tvm-ffi (FLA's TileLang backend, optional but
|
||||
# adds ~26% on Qwen3.5 GDN layers on Hopper+)
|
||||
# 5) flash-attn (only for max_seq_length >= 32k, separate concern)
|
||||
#
|
||||
# Primary gate: hook transformers' fast-path availability functions.
|
||||
# When the loaded model's modeling file does
|
||||
# if is_flash_linear_attention_available(): from fla.modules import ...
|
||||
# the hook fires and installs FLA + tilelang on the spot. Models that
|
||||
# never call those gates never trigger the install. This supersedes
|
||||
# substring-based detection for FLA and tilelang.
|
||||
#
|
||||
# Legacy substring path remains for:
|
||||
# - causal-conv1d (also covered by its own hook, but the substring
|
||||
# installer prefers the wheel-first install and is reused here
|
||||
# from the hook closure too)
|
||||
# - mamba-ssm (true SSM families)
|
||||
# - flash-attn (long-context only)
|
||||
# plus an opt-out via UNSLOTH_STUDIO_SKIP_FAST_PATH_HOOKS=1.
|
||||
try:
|
||||
_ensure_causal_conv1d_fast_path(event_queue, model_name)
|
||||
_ensure_flash_linear_attention(event_queue, model_name)
|
||||
if os.getenv(_FAST_PATH_HOOKS_SKIP_ENV) == "1":
|
||||
_ensure_causal_conv1d_fast_path(event_queue, model_name)
|
||||
_ensure_flash_linear_attention(event_queue, model_name)
|
||||
_ensure_tilelang_backend(event_queue, model_name)
|
||||
else:
|
||||
_install_fast_path_hooks(event_queue)
|
||||
_ensure_mamba_ssm(event_queue, model_name)
|
||||
_ensure_tilelang_backend(event_queue, model_name)
|
||||
_ensure_flash_attn_for_long_context(
|
||||
event_queue,
|
||||
int(config.get("max_seq_length", 2048)),
|
||||
|
|
|
|||
|
|
@ -580,3 +580,280 @@ def test_tilelang_backend_swallows_install_failure(monkeypatch):
|
|||
|
||||
run_mock.assert_called_once()
|
||||
assert any("failed" in s.lower() for s in statuses)
|
||||
|
||||
|
||||
# ───────────────────────────────────────────────────────────────────
|
||||
# Runtime hook on `is_flash_linear_attention_available` /
|
||||
# `is_causal_conv1d_available`. These are the primary gate in
|
||||
# normal operation; the substring tests above cover the
|
||||
# UNSLOTH_STUDIO_SKIP_FAST_PATH_HOOKS=1 fallback.
|
||||
# ───────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class _FakeQueue(list):
|
||||
"""List with `.put` so worker._send_status can send into it during tests."""
|
||||
|
||||
def put(self, item):
|
||||
self.append(item)
|
||||
|
||||
|
||||
def _make_fake_gate(initial_return: bool):
|
||||
"""Build a callable that mimics transformers' lru_cache-decorated gates.
|
||||
|
||||
Tracks call count and exposes a `cache_clear` attribute. The return
|
||||
value can be flipped to mimic install-then-True behaviour by setting
|
||||
`.next_return`.
|
||||
"""
|
||||
|
||||
class Gate:
|
||||
def __init__(self, initial: bool) -> None:
|
||||
self.next_return = initial
|
||||
self.call_count = 0
|
||||
self.cache_clear_count = 0
|
||||
|
||||
def __call__(self) -> bool:
|
||||
self.call_count += 1
|
||||
return self.next_return
|
||||
|
||||
def cache_clear(self) -> None:
|
||||
self.cache_clear_count += 1
|
||||
|
||||
return Gate(initial_return)
|
||||
|
||||
|
||||
def _patch_iu_gates(monkeypatch, fla_gate, conv_gate):
|
||||
"""Drop fake gates onto transformers.utils.import_utils for the test."""
|
||||
from transformers.utils import import_utils as _iu
|
||||
|
||||
monkeypatch.setattr(_iu, "is_flash_linear_attention_available", fla_gate)
|
||||
monkeypatch.setattr(_iu, "is_causal_conv1d_available", conv_gate)
|
||||
|
||||
|
||||
def test_hook_installs_when_gate_returns_false(monkeypatch):
|
||||
fla_gate = _make_fake_gate(initial_return=False)
|
||||
conv_gate = _make_fake_gate(initial_return=False)
|
||||
_patch_iu_gates(monkeypatch, fla_gate, conv_gate)
|
||||
|
||||
fla_install = mock.Mock(side_effect=lambda eq: setattr(fla_gate, "next_return", True))
|
||||
tile_install = mock.Mock(side_effect=lambda eq: None)
|
||||
conv_install = mock.Mock(side_effect=lambda **kw: setattr(conv_gate, "next_return", True))
|
||||
|
||||
monkeypatch.setattr(
|
||||
worker, "_ensure_flash_linear_attention_unconditional", fla_install
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
worker, "_ensure_tilelang_backend_unconditional", tile_install
|
||||
)
|
||||
monkeypatch.setattr(worker, "_install_package_wheel_first", conv_install)
|
||||
monkeypatch.delenv(worker._FAST_PATH_HOOKS_SKIP_ENV, raising=False)
|
||||
|
||||
worker._install_fast_path_hooks(event_queue=_FakeQueue())
|
||||
|
||||
from transformers.utils import import_utils as _iu
|
||||
|
||||
# Both gates are now wrapped. Call them — the hook should drive the install.
|
||||
assert _iu.is_flash_linear_attention_available() is True
|
||||
fla_install.assert_called_once()
|
||||
tile_install.assert_called_once()
|
||||
assert _iu.is_causal_conv1d_available() is True
|
||||
conv_install.assert_called_once()
|
||||
|
||||
|
||||
def test_hook_skips_install_when_gate_already_true(monkeypatch):
|
||||
fla_gate = _make_fake_gate(initial_return=True)
|
||||
conv_gate = _make_fake_gate(initial_return=True)
|
||||
_patch_iu_gates(monkeypatch, fla_gate, conv_gate)
|
||||
|
||||
fla_install = mock.Mock()
|
||||
tile_install = mock.Mock()
|
||||
conv_install = mock.Mock()
|
||||
monkeypatch.setattr(
|
||||
worker, "_ensure_flash_linear_attention_unconditional", fla_install
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
worker, "_ensure_tilelang_backend_unconditional", tile_install
|
||||
)
|
||||
monkeypatch.setattr(worker, "_install_package_wheel_first", conv_install)
|
||||
monkeypatch.delenv(worker._FAST_PATH_HOOKS_SKIP_ENV, raising=False)
|
||||
|
||||
worker._install_fast_path_hooks(event_queue=_FakeQueue())
|
||||
|
||||
from transformers.utils import import_utils as _iu
|
||||
|
||||
assert _iu.is_flash_linear_attention_available() is True
|
||||
assert _iu.is_causal_conv1d_available() is True
|
||||
fla_install.assert_not_called()
|
||||
tile_install.assert_not_called()
|
||||
conv_install.assert_not_called()
|
||||
|
||||
|
||||
def test_hook_idempotent_on_repeat_call(monkeypatch):
|
||||
fla_gate = _make_fake_gate(initial_return=False)
|
||||
conv_gate = _make_fake_gate(initial_return=False)
|
||||
_patch_iu_gates(monkeypatch, fla_gate, conv_gate)
|
||||
|
||||
fla_install = mock.Mock(side_effect=lambda eq: setattr(fla_gate, "next_return", True))
|
||||
tile_install = mock.Mock()
|
||||
conv_install = mock.Mock(side_effect=lambda **kw: setattr(conv_gate, "next_return", True))
|
||||
monkeypatch.setattr(
|
||||
worker, "_ensure_flash_linear_attention_unconditional", fla_install
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
worker, "_ensure_tilelang_backend_unconditional", tile_install
|
||||
)
|
||||
monkeypatch.setattr(worker, "_install_package_wheel_first", conv_install)
|
||||
monkeypatch.delenv(worker._FAST_PATH_HOOKS_SKIP_ENV, raising=False)
|
||||
|
||||
worker._install_fast_path_hooks(event_queue=_FakeQueue())
|
||||
|
||||
from transformers.utils import import_utils as _iu
|
||||
|
||||
# First call: hook fires.
|
||||
_iu.is_flash_linear_attention_available()
|
||||
# Subsequent calls: must not re-trigger the installer.
|
||||
_iu.is_flash_linear_attention_available()
|
||||
_iu.is_flash_linear_attention_available()
|
||||
assert fla_install.call_count == 1
|
||||
assert tile_install.call_count == 1
|
||||
|
||||
|
||||
def test_hook_handles_install_failure_gracefully(monkeypatch):
|
||||
fla_gate = _make_fake_gate(initial_return=False)
|
||||
conv_gate = _make_fake_gate(initial_return=True) # bypass to focus on FLA
|
||||
_patch_iu_gates(monkeypatch, fla_gate, conv_gate)
|
||||
|
||||
def raising_install(eq):
|
||||
raise RuntimeError("pip failed to fetch wheel")
|
||||
|
||||
monkeypatch.setattr(
|
||||
worker, "_ensure_flash_linear_attention_unconditional", raising_install
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
worker, "_ensure_tilelang_backend_unconditional", lambda eq: None
|
||||
)
|
||||
monkeypatch.setattr(worker, "_install_package_wheel_first", lambda **kw: None)
|
||||
monkeypatch.delenv(worker._FAST_PATH_HOOKS_SKIP_ENV, raising=False)
|
||||
|
||||
worker._install_fast_path_hooks(event_queue=_FakeQueue())
|
||||
|
||||
from transformers.utils import import_utils as _iu
|
||||
|
||||
# Must not raise; returns False so transformers falls back to torch loop.
|
||||
assert _iu.is_flash_linear_attention_available() is False
|
||||
|
||||
|
||||
def test_hook_can_be_disabled_via_env(monkeypatch):
|
||||
fla_gate = _make_fake_gate(initial_return=False)
|
||||
conv_gate = _make_fake_gate(initial_return=False)
|
||||
_patch_iu_gates(monkeypatch, fla_gate, conv_gate)
|
||||
|
||||
fla_install = mock.Mock()
|
||||
monkeypatch.setattr(
|
||||
worker, "_ensure_flash_linear_attention_unconditional", fla_install
|
||||
)
|
||||
monkeypatch.setenv(worker._FAST_PATH_HOOKS_SKIP_ENV, "1")
|
||||
|
||||
worker._install_fast_path_hooks(event_queue=_FakeQueue())
|
||||
|
||||
from transformers.utils import import_utils as _iu
|
||||
|
||||
# Hook should NOT have been installed; gates remain the fakes.
|
||||
assert _iu.is_flash_linear_attention_available is fla_gate
|
||||
assert _iu.is_causal_conv1d_available is conv_gate
|
||||
fla_install.assert_not_called()
|
||||
|
||||
|
||||
def test_hook_clears_lru_cache_before_first_check(monkeypatch):
|
||||
fla_gate = _make_fake_gate(initial_return=True)
|
||||
conv_gate = _make_fake_gate(initial_return=True)
|
||||
_patch_iu_gates(monkeypatch, fla_gate, conv_gate)
|
||||
|
||||
monkeypatch.setattr(
|
||||
worker, "_ensure_flash_linear_attention_unconditional", lambda eq: None
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
worker, "_ensure_tilelang_backend_unconditional", lambda eq: None
|
||||
)
|
||||
monkeypatch.setattr(worker, "_install_package_wheel_first", lambda **kw: None)
|
||||
monkeypatch.delenv(worker._FAST_PATH_HOOKS_SKIP_ENV, raising=False)
|
||||
|
||||
worker._install_fast_path_hooks(event_queue=_FakeQueue())
|
||||
from transformers.utils import import_utils as _iu
|
||||
|
||||
_iu.is_flash_linear_attention_available()
|
||||
# The wrapper called cache_clear at least once before delegating.
|
||||
assert fla_gate.cache_clear_count >= 1
|
||||
|
||||
|
||||
def test_hook_rewrites_previously_imported_module_bindings(monkeypatch):
|
||||
"""Modeling files bind `is_flash_linear_attention_available` locally
|
||||
via `from ... import is_X`. Reassigning the attribute on
|
||||
transformers.utils.import_utils alone does NOT reach those local
|
||||
bindings. The hook installer sweeps sys.modules and rebinds them.
|
||||
"""
|
||||
fla_gate = _make_fake_gate(initial_return=False)
|
||||
conv_gate = _make_fake_gate(initial_return=True)
|
||||
_patch_iu_gates(monkeypatch, fla_gate, conv_gate)
|
||||
|
||||
# Create a fake modeling module that did `from ... import is_flash_linear_attention_available`.
|
||||
fake_mod = sys.modules.setdefault(
|
||||
"_test_fake_modeling_qwen35", type(sys)("_test_fake_modeling_qwen35")
|
||||
)
|
||||
fake_mod.is_flash_linear_attention_available = fla_gate
|
||||
|
||||
def fake_install(eq):
|
||||
fla_gate.next_return = True
|
||||
|
||||
monkeypatch.setattr(
|
||||
worker, "_ensure_flash_linear_attention_unconditional", fake_install
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
worker, "_ensure_tilelang_backend_unconditional", lambda eq: None
|
||||
)
|
||||
monkeypatch.setattr(worker, "_install_package_wheel_first", lambda **kw: None)
|
||||
monkeypatch.delenv(worker._FAST_PATH_HOOKS_SKIP_ENV, raising=False)
|
||||
|
||||
worker._install_fast_path_hooks(event_queue=_FakeQueue())
|
||||
|
||||
# The fake module's local binding has been rewritten to the wrapper.
|
||||
assert fake_mod.is_flash_linear_attention_available is not fla_gate
|
||||
# Calling through the fake module's reference triggers the install.
|
||||
assert fake_mod.is_flash_linear_attention_available() is True
|
||||
|
||||
del sys.modules["_test_fake_modeling_qwen35"]
|
||||
|
||||
|
||||
def test_hook_skips_when_import_utils_unavailable(monkeypatch):
|
||||
"""If transformers.utils.import_utils can't be imported, the hook
|
||||
installer must log and return cleanly rather than crash the worker."""
|
||||
real_import = builtins.__import__
|
||||
|
||||
def fake_import(name, *a, **kw):
|
||||
if name == "transformers.utils" or name == "transformers.utils.import_utils":
|
||||
raise ImportError("transformers missing in worker venv")
|
||||
return real_import(name, *a, **kw)
|
||||
|
||||
monkeypatch.setattr(builtins, "__import__", fake_import)
|
||||
monkeypatch.delenv(worker._FAST_PATH_HOOKS_SKIP_ENV, raising=False)
|
||||
|
||||
# Should not raise.
|
||||
worker._install_fast_path_hooks(event_queue=_FakeQueue())
|
||||
|
||||
|
||||
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."""
|
||||
install_mock = mock.Mock()
|
||||
monkeypatch.setattr(
|
||||
worker, "_ensure_flash_linear_attention_unconditional", install_mock
|
||||
)
|
||||
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")
|
||||
assert install_mock.call_count == 1
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue