Merge branch 'studio-fla-tilelang-qwen3.5' of https://github.com/unslothai/unsloth into studio-fla-tilelang-qwen3.5

This commit is contained in:
danielhanchen 2026-05-17 01:47:07 +00:00
commit 379cbb23eb

View file

@ -486,14 +486,14 @@ def test_tilelang_backend_reinstalls_when_tvm_ffi_is_broken(monkeypatch):
# Repair: --force-reinstall --no-deps, apache-tvm-ffi ONLY (no tilelang).
assert "--force-reinstall" in repair_args
assert "--no-deps" in repair_args, (
"Repair MUST use --no-deps to avoid replacing torch / CUDA"
)
assert (
"--no-deps" in repair_args
), "Repair MUST use --no-deps to avoid replacing torch / CUDA"
assert "--only-binary=:all:" in repair_args
assert f"apache-tvm-ffi=={worker._APACHE_TVM_FFI_PACKAGE_VERSION}" in repair_args
assert all("tilelang" not in a for a in repair_args), (
"Repair MUST only touch apache-tvm-ffi"
)
assert all(
"tilelang" not in a for a in repair_args
), "Repair MUST only touch apache-tvm-ffi"
# Install: regular dep-resolving install, NO --force-reinstall.
assert "--force-reinstall" not in install_args
@ -654,32 +654,33 @@ def _patch_iu_gates(monkeypatch, fla_gate, 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)
fla_gate = _make_fake_gate(initial_return = False)
conv_gate = _make_fake_gate(initial_return = False)
_patch_iu_gates(monkeypatch, fla_gate, conv_gate)
def _fla_install_side_effect(eq):
fla_gate.next_return = True
return True
fla_install = mock.Mock(side_effect=_fla_install_side_effect)
tile_install = mock.Mock(side_effect=lambda eq: None)
fla_install = mock.Mock(side_effect = _fla_install_side_effect)
tile_install = mock.Mock(side_effect = lambda eq: None)
def _conv_install_side_effect(**kw):
conv_gate.next_return = True
return True
conv_install = mock.Mock(side_effect=_conv_install_side_effect)
conv_install = mock.Mock(side_effect = _conv_install_side_effect)
monkeypatch.setattr(
worker, "_ensure_flash_linear_attention_unconditional", fla_install
)
monkeypatch.setattr(
worker, "_ensure_tilelang_backend_unconditional", tile_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)
monkeypatch.delenv(worker._FAST_PATH_HOOKS_SKIP_ENV, raising = False)
worker._install_fast_path_hooks(event_queue=_FakeQueue(), model_name="unsloth/Qwen3.5-2B")
worker._install_fast_path_hooks(
event_queue = _FakeQueue(), model_name = "unsloth/Qwen3.5-2B"
)
from transformers.utils import import_utils as _iu
@ -696,8 +697,8 @@ def test_hook_skips_install_when_gate_already_true(monkeypatch):
must do zero install work. (Tilelang repair on the already-True path
is covered by test_hook_runs_tilelang_repair_when_fla_already_true.)
"""
fla_gate = _make_fake_gate(initial_return=True)
conv_gate = _make_fake_gate(initial_return=True)
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()
@ -706,18 +707,18 @@ def test_hook_skips_install_when_gate_already_true(monkeypatch):
monkeypatch.setattr(
worker, "_ensure_flash_linear_attention_unconditional", fla_install
)
monkeypatch.setattr(
worker, "_ensure_tilelang_backend_unconditional", tile_install
)
monkeypatch.setattr(worker, "_ensure_tilelang_backend_unconditional", tile_install)
monkeypatch.setattr(worker, "_install_package_wheel_first", conv_install)
# Tilelang healthy so the post_available path is a no-op (otherwise
# it would call tile_install, which is correct behaviour but
# outside the scope of this test).
monkeypatch.setattr(worker, "_tilelang_importable", lambda: True)
monkeypatch.setattr(worker, "_installed_tvm_ffi_version", lambda: "0.1.9")
monkeypatch.delenv(worker._FAST_PATH_HOOKS_SKIP_ENV, raising=False)
monkeypatch.delenv(worker._FAST_PATH_HOOKS_SKIP_ENV, raising = False)
worker._install_fast_path_hooks(event_queue=_FakeQueue(), model_name="unsloth/Qwen3.5-2B")
worker._install_fast_path_hooks(
event_queue = _FakeQueue(), model_name = "unsloth/Qwen3.5-2B"
)
from transformers.utils import import_utils as _iu
@ -729,31 +730,32 @@ def test_hook_skips_install_when_gate_already_true(monkeypatch):
def test_hook_idempotent_on_repeat_call(monkeypatch):
fla_gate = _make_fake_gate(initial_return=False)
conv_gate = _make_fake_gate(initial_return=False)
fla_gate = _make_fake_gate(initial_return = False)
conv_gate = _make_fake_gate(initial_return = False)
_patch_iu_gates(monkeypatch, fla_gate, conv_gate)
def _fla_install_side_effect(eq):
fla_gate.next_return = True
return True
fla_install = mock.Mock(side_effect=_fla_install_side_effect)
fla_install = mock.Mock(side_effect = _fla_install_side_effect)
tile_install = mock.Mock()
def _conv_install_side_effect(**kw):
conv_gate.next_return = True
return True
conv_install = mock.Mock(side_effect=_conv_install_side_effect)
conv_install = mock.Mock(side_effect = _conv_install_side_effect)
monkeypatch.setattr(
worker, "_ensure_flash_linear_attention_unconditional", fla_install
)
monkeypatch.setattr(
worker, "_ensure_tilelang_backend_unconditional", tile_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)
monkeypatch.delenv(worker._FAST_PATH_HOOKS_SKIP_ENV, raising = False)
worker._install_fast_path_hooks(event_queue=_FakeQueue(), model_name="unsloth/Qwen3.5-2B")
worker._install_fast_path_hooks(
event_queue = _FakeQueue(), model_name = "unsloth/Qwen3.5-2B"
)
from transformers.utils import import_utils as _iu
@ -767,8 +769,8 @@ def test_hook_idempotent_on_repeat_call(monkeypatch):
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
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):
@ -781,9 +783,11 @@ def test_hook_handles_install_failure_gracefully(monkeypatch):
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)
monkeypatch.delenv(worker._FAST_PATH_HOOKS_SKIP_ENV, raising = False)
worker._install_fast_path_hooks(event_queue=_FakeQueue(), model_name="unsloth/Qwen3.5-2B")
worker._install_fast_path_hooks(
event_queue = _FakeQueue(), model_name = "unsloth/Qwen3.5-2B"
)
from transformers.utils import import_utils as _iu
@ -792,8 +796,8 @@ def test_hook_handles_install_failure_gracefully(monkeypatch):
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)
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()
@ -802,7 +806,9 @@ def test_hook_can_be_disabled_via_env(monkeypatch):
)
monkeypatch.setenv(worker._FAST_PATH_HOOKS_SKIP_ENV, "1")
worker._install_fast_path_hooks(event_queue=_FakeQueue(), model_name="unsloth/Qwen3.5-2B")
worker._install_fast_path_hooks(
event_queue = _FakeQueue(), model_name = "unsloth/Qwen3.5-2B"
)
from transformers.utils import import_utils as _iu
@ -813,8 +819,8 @@ def test_hook_can_be_disabled_via_env(monkeypatch):
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)
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(
@ -824,9 +830,11 @@ def test_hook_clears_lru_cache_before_first_check(monkeypatch):
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)
monkeypatch.delenv(worker._FAST_PATH_HOOKS_SKIP_ENV, raising = False)
worker._install_fast_path_hooks(event_queue=_FakeQueue(), model_name="unsloth/Qwen3.5-2B")
worker._install_fast_path_hooks(
event_queue = _FakeQueue(), model_name = "unsloth/Qwen3.5-2B"
)
from transformers.utils import import_utils as _iu
_iu.is_flash_linear_attention_available()
@ -840,8 +848,8 @@ def test_hook_rewrites_previously_imported_module_bindings(monkeypatch):
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)
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`.
@ -861,9 +869,11 @@ def test_hook_rewrites_previously_imported_module_bindings(monkeypatch):
worker, "_ensure_tilelang_backend_unconditional", lambda eq: True
)
monkeypatch.setattr(worker, "_install_package_wheel_first", lambda **kw: True)
monkeypatch.delenv(worker._FAST_PATH_HOOKS_SKIP_ENV, raising=False)
monkeypatch.delenv(worker._FAST_PATH_HOOKS_SKIP_ENV, raising = False)
worker._install_fast_path_hooks(event_queue=_FakeQueue(), model_name="unsloth/Qwen3.5-2B")
worker._install_fast_path_hooks(
event_queue = _FakeQueue(), model_name = "unsloth/Qwen3.5-2B"
)
# The fake module's local binding has been rewritten to the wrapper.
assert fake_mod.is_flash_linear_attention_available is not fla_gate
@ -884,10 +894,12 @@ def test_hook_skips_when_import_utils_unavailable(monkeypatch):
return real_import(name, *a, **kw)
monkeypatch.setattr(builtins, "__import__", fake_import)
monkeypatch.delenv(worker._FAST_PATH_HOOKS_SKIP_ENV, raising=False)
monkeypatch.delenv(worker._FAST_PATH_HOOKS_SKIP_ENV, raising = False)
# Should not raise.
worker._install_fast_path_hooks(event_queue=_FakeQueue(), model_name="unsloth/Qwen3.5-2B")
worker._install_fast_path_hooks(
event_queue = _FakeQueue(), model_name = "unsloth/Qwen3.5-2B"
)
def test_substring_fallback_unchanged_when_hook_skipped(monkeypatch):
@ -901,11 +913,15 @@ def test_substring_fallback_unchanged_when_hook_skipped(monkeypatch):
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")
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")
worker._ensure_flash_linear_attention(
event_queue = [], model_name = "meta-llama/Llama-3.1-8B"
)
assert install_mock.call_count == 1
@ -926,27 +942,27 @@ def test_hook_does_not_install_tilelang_for_non_qwen_fla_model(monkeypatch):
"""Finding #1: OLMo-Hybrid (and similar non-Qwen GDN models) call
`is_flash_linear_attention_available` but should NOT get tilelang,
which is a Qwen3.5-family optimisation. Was unconditional before."""
fla_gate = _make_fake_gate(initial_return=False)
conv_gate = _make_fake_gate(initial_return=True)
fla_gate = _make_fake_gate(initial_return = False)
conv_gate = _make_fake_gate(initial_return = True)
_patch_iu_gates(monkeypatch, fla_gate, conv_gate)
def _fla_install(eq):
fla_gate.next_return = True
return True
fla_install = mock.Mock(side_effect=_fla_install)
tile_install = mock.Mock(return_value=True)
fla_install = mock.Mock(side_effect = _fla_install)
tile_install = mock.Mock(return_value = True)
monkeypatch.setattr(
worker, "_ensure_flash_linear_attention_unconditional", fla_install
)
monkeypatch.setattr(worker, "_ensure_tilelang_backend_unconditional", tile_install)
monkeypatch.setattr(
worker, "_ensure_tilelang_backend_unconditional", tile_install
worker, "_install_package_wheel_first", mock.Mock(return_value = True)
)
monkeypatch.setattr(worker, "_install_package_wheel_first", mock.Mock(return_value=True))
monkeypatch.delenv(worker._FAST_PATH_HOOKS_SKIP_ENV, raising=False)
monkeypatch.delenv(worker._FAST_PATH_HOOKS_SKIP_ENV, raising = False)
worker._install_fast_path_hooks(
event_queue=_FakeQueue(), model_name="allenai/OLMo-Hybrid-1B"
event_queue = _FakeQueue(), model_name = "allenai/OLMo-Hybrid-1B"
)
from transformers.utils import import_utils as _iu
@ -958,27 +974,27 @@ def test_hook_does_not_install_tilelang_for_non_qwen_fla_model(monkeypatch):
def test_hook_does_install_tilelang_for_qwen35(monkeypatch):
"""Positive control for finding #1: Qwen3.5 still gets tilelang."""
fla_gate = _make_fake_gate(initial_return=False)
conv_gate = _make_fake_gate(initial_return=True)
fla_gate = _make_fake_gate(initial_return = False)
conv_gate = _make_fake_gate(initial_return = True)
_patch_iu_gates(monkeypatch, fla_gate, conv_gate)
def _fla_install(eq):
fla_gate.next_return = True
return True
fla_install = mock.Mock(side_effect=_fla_install)
tile_install = mock.Mock(return_value=True)
fla_install = mock.Mock(side_effect = _fla_install)
tile_install = mock.Mock(return_value = True)
monkeypatch.setattr(
worker, "_ensure_flash_linear_attention_unconditional", fla_install
)
monkeypatch.setattr(worker, "_ensure_tilelang_backend_unconditional", tile_install)
monkeypatch.setattr(
worker, "_ensure_tilelang_backend_unconditional", tile_install
worker, "_install_package_wheel_first", mock.Mock(return_value = True)
)
monkeypatch.setattr(worker, "_install_package_wheel_first", mock.Mock(return_value=True))
monkeypatch.delenv(worker._FAST_PATH_HOOKS_SKIP_ENV, raising=False)
monkeypatch.delenv(worker._FAST_PATH_HOOKS_SKIP_ENV, raising = False)
worker._install_fast_path_hooks(
event_queue=_FakeQueue(), model_name="unsloth/Qwen3.5-2B"
event_queue = _FakeQueue(), model_name = "unsloth/Qwen3.5-2B"
)
from transformers.utils import import_utils as _iu
@ -993,16 +1009,14 @@ def test_tilelang_repair_does_not_touch_torch_cuda_stack(monkeypatch):
forced step so --force-reinstall does not cascade through
apache-tvm-ffi's dep graph and pull a different torch wheel.
"""
monkeypatch.delenv(worker._TILELANG_SKIP_ENV, raising=False)
monkeypatch.delenv(worker._TILELANG_SKIP_ENV, raising = False)
monkeypatch.setattr(worker.shutil, "which", lambda name: "/usr/bin/uv")
monkeypatch.setattr(worker, "_installed_tvm_ffi_version", lambda: "0.1.10")
run_mock = mock.Mock(return_value=mock.Mock(returncode=0, stdout=""))
run_mock = mock.Mock(return_value = mock.Mock(returncode = 0, stdout = ""))
monkeypatch.setattr(worker._sp, "run", run_mock)
monkeypatch.setattr(worker, "_send_status", lambda *a, **k: None)
worker._ensure_tilelang_backend(
event_queue=[], model_name="unsloth/Qwen3.5-2B"
)
worker._ensure_tilelang_backend(event_queue = [], model_name = "unsloth/Qwen3.5-2B")
assert run_mock.call_count == 2
repair_args = run_mock.call_args_list[0][0][0]
@ -1029,28 +1043,30 @@ def test_hook_trusts_installer_bool_not_metadata(monkeypatch):
False so transformers takes the torch fallback.
"""
# Gate flips True after install (simulating "metadata sees fla").
fla_gate = _make_fake_gate(initial_return=False)
conv_gate = _make_fake_gate(initial_return=True)
fla_gate = _make_fake_gate(initial_return = False)
conv_gate = _make_fake_gate(initial_return = True)
_patch_iu_gates(monkeypatch, fla_gate, conv_gate)
# Installer "succeeds" at pip, AND flips the gate to True (metadata
# sees fla post-install), BUT returns False (deep import broken).
def _bad_install(eq):
fla_gate.next_return = True # metadata says yes after pip
return False # but deep import is broken
return False # but deep import is broken
fake_fla_install = mock.Mock(side_effect=_bad_install)
fake_fla_install = mock.Mock(side_effect = _bad_install)
monkeypatch.setattr(
worker, "_ensure_flash_linear_attention_unconditional", fake_fla_install
)
monkeypatch.setattr(
worker, "_ensure_tilelang_backend_unconditional", mock.Mock(return_value=True)
worker, "_ensure_tilelang_backend_unconditional", mock.Mock(return_value = True)
)
monkeypatch.setattr(worker, "_install_package_wheel_first", mock.Mock(return_value=True))
monkeypatch.delenv(worker._FAST_PATH_HOOKS_SKIP_ENV, raising=False)
monkeypatch.setattr(
worker, "_install_package_wheel_first", mock.Mock(return_value = True)
)
monkeypatch.delenv(worker._FAST_PATH_HOOKS_SKIP_ENV, raising = False)
worker._install_fast_path_hooks(
event_queue=_FakeQueue(), model_name="unsloth/Qwen3.5-2B"
event_queue = _FakeQueue(), model_name = "unsloth/Qwen3.5-2B"
)
from transformers.utils import import_utils as _iu
@ -1081,13 +1097,13 @@ def test_rebind_does_not_trigger_module_getattr(monkeypatch):
# No module-level binding to `is_flash_linear_attention_available`
# in __dict__, so the sweep must NOT trip the tripwire.
worker._rebind_in_already_imported_modules(
attr_name="is_flash_linear_attention_available",
old_obj=original,
new_obj=replacement,
)
assert not _GetattrTripwire.getattr_called, (
"Rebind sweep invoked __getattr__ — should use __dict__ probe"
attr_name = "is_flash_linear_attention_available",
old_obj = original,
new_obj = replacement,
)
assert (
not _GetattrTripwire.getattr_called
), "Rebind sweep invoked __getattr__ — should use __dict__ probe"
finally:
sys.modules.pop("_lazy_test_module", None)
@ -1097,20 +1113,20 @@ def test_hook_skips_tilelang_when_fla_install_is_skipped(monkeypatch):
_ensure_flash_linear_attention_unconditional; tilelang must NOT
install in that case.
"""
fla_gate = _make_fake_gate(initial_return=False)
conv_gate = _make_fake_gate(initial_return=True)
fla_gate = _make_fake_gate(initial_return = False)
conv_gate = _make_fake_gate(initial_return = True)
_patch_iu_gates(monkeypatch, fla_gate, conv_gate)
monkeypatch.setenv(worker._FLA_SKIP_ENV, "1")
tile_install = mock.Mock(return_value=True)
tile_install = mock.Mock(return_value = True)
monkeypatch.setattr(worker, "_ensure_tilelang_backend_unconditional", tile_install)
monkeypatch.setattr(
worker, "_ensure_tilelang_backend_unconditional", tile_install
worker, "_install_package_wheel_first", mock.Mock(return_value = True)
)
monkeypatch.setattr(worker, "_install_package_wheel_first", mock.Mock(return_value=True))
monkeypatch.delenv(worker._FAST_PATH_HOOKS_SKIP_ENV, raising=False)
monkeypatch.delenv(worker._FAST_PATH_HOOKS_SKIP_ENV, raising = False)
worker._install_fast_path_hooks(
event_queue=_FakeQueue(), model_name="unsloth/Qwen3.5-2B"
event_queue = _FakeQueue(), model_name = "unsloth/Qwen3.5-2B"
)
from transformers.utils import import_utils as _iu
@ -1125,26 +1141,26 @@ def test_hook_runs_tilelang_repair_when_fla_already_true(monkeypatch):
first probe) but tilelang is missing or apache-tvm-ffi is on the
broken list, the post-available action must still run tilelang.
"""
fla_gate = _make_fake_gate(initial_return=True)
conv_gate = _make_fake_gate(initial_return=True)
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(return_value=True)
tile_install = mock.Mock(return_value=True)
fla_install = mock.Mock(return_value = True)
tile_install = mock.Mock(return_value = True)
monkeypatch.setattr(
worker, "_ensure_flash_linear_attention_unconditional", fla_install
)
monkeypatch.setattr(worker, "_ensure_tilelang_backend_unconditional", tile_install)
monkeypatch.setattr(
worker, "_ensure_tilelang_backend_unconditional", tile_install
worker, "_install_package_wheel_first", mock.Mock(return_value = True)
)
monkeypatch.setattr(worker, "_install_package_wheel_first", mock.Mock(return_value=True))
# tilelang missing AND tvm-ffi is on broken list — both trigger repair.
monkeypatch.setattr(worker, "_tilelang_importable", lambda: False)
monkeypatch.setattr(worker, "_installed_tvm_ffi_version", lambda: "0.1.11")
monkeypatch.delenv(worker._FAST_PATH_HOOKS_SKIP_ENV, raising=False)
monkeypatch.delenv(worker._FAST_PATH_HOOKS_SKIP_ENV, raising = False)
worker._install_fast_path_hooks(
event_queue=_FakeQueue(), model_name="unsloth/Qwen3.5-2B"
event_queue = _FakeQueue(), model_name = "unsloth/Qwen3.5-2B"
)
from transformers.utils import import_utils as _iu
@ -1159,23 +1175,23 @@ def test_fla_installer_force_reinstalls_when_older_version_present(monkeypatch):
"""Finding #8: when an older `flash-linear-attention` is importable
but below the pin, the installer must force a reinstall (not no-op).
"""
monkeypatch.delenv(worker._FLA_SKIP_ENV, raising=False)
monkeypatch.delenv(worker._FLA_SKIP_ENV, raising = False)
monkeypatch.setattr(worker.shutil, "which", lambda name: "/usr/bin/uv")
monkeypatch.setattr(worker, "_installed_torch_version_tuple", lambda: (2, 9))
# Importable but stale (current() reports False even though importable() is True).
monkeypatch.setattr(worker, "_flash_linear_attention_importable", lambda: True)
monkeypatch.setattr(worker, "_flash_linear_attention_current", lambda **kw: False)
run_mock = mock.Mock(return_value=mock.Mock(returncode=0, stdout=""))
run_mock = mock.Mock(return_value = mock.Mock(returncode = 0, stdout = ""))
monkeypatch.setattr(worker._sp, "run", run_mock)
monkeypatch.setattr(worker, "_send_status", lambda *a, **k: None)
worker._ensure_flash_linear_attention_unconditional(event_queue=[])
worker._ensure_flash_linear_attention_unconditional(event_queue = [])
run_mock.assert_called_once()
args = run_mock.call_args[0][0]
assert "--force-reinstall" in args, (
"Stale FLA must trigger --force-reinstall, otherwise pip is a no-op"
)
assert (
"--force-reinstall" in args
), "Stale FLA must trigger --force-reinstall, otherwise pip is a no-op"
# --no-deps still applies so torch stays untouched.
assert "--no-deps" in args
@ -1191,6 +1207,7 @@ def test_run_training_process_eagerly_installs_causal_conv1d_in_normal_mode():
asserts the eager install is OUTSIDE the if/else hook branch.
"""
import inspect
src = inspect.getsource(worker.run_training_process)
# Find the orchestration block.
assert "_ensure_causal_conv1d_fast_path(event_queue, model_name)" in src