Skip to content
Merged
Show file tree
Hide file tree
Changes from 2 commits
Commits
Show all changes
61 commits
Select commit Hold shift + click to select a range
c07cdda
studio: install flash-linear-attention and tilelang for Qwen3.5 family
danielhanchen May 15, 2026
0bb03e0
[pre-commit.ci] auto fixes from pre-commit.com hooks
pre-commit-ci[bot] May 15, 2026
57afa62
Merge branch 'main' into studio-fla-tilelang-qwen3.5
danielhanchen May 15, 2026
92f9d4b
tests/studio: accept new grad_norm arg in MLX smoke _on_step callback
danielhanchen May 15, 2026
bbd715e
[pre-commit.ci] auto fixes from pre-commit.com hooks
pre-commit-ci[bot] May 15, 2026
1ff38ae
Merge branch 'main' into studio-fla-tilelang-qwen3.5
danielhanchen May 15, 2026
49a0db9
ci: retrigger after zoo drift + IPython fixes landed in main
danielhanchen May 15, 2026
b92afb7
tests/studio: pin max_grad_value=0 in MLX smoke so max_grad_norm=1.0 …
danielhanchen May 15, 2026
d079859
tests/studio: clarify why MLX smoke pins max_grad_value=0
danielhanchen May 15, 2026
a9982a0
Merge branch 'main' into studio-fla-tilelang-qwen3.5
danielhanchen May 15, 2026
6a1a215
Merge branch 'main' into studio-fla-tilelang-qwen3.5
danielhanchen May 15, 2026
9e9c3ac
Merge branch 'main' into studio-fla-tilelang-qwen3.5
danielhanchen May 15, 2026
b399247
tests/studio: replace fragile substring gate with loss + round-trip g…
danielhanchen May 15, 2026
d7f3a3e
[pre-commit.ci] auto fixes from pre-commit.com hooks
pre-commit-ci[bot] May 15, 2026
453c31a
Merge branch 'main' into studio-fla-tilelang-qwen3.5
danielhanchen May 15, 2026
39559cc
tests/studio: gate MLX reload on training-row loss, not greedy text
danielhanchen May 15, 2026
994688d
Merge branch 'main' into studio-fla-tilelang-qwen3.5
danielhanchen May 15, 2026
4454608
ci: retrigger Backend CI after transient pwsh-startup timeout
danielhanchen May 15, 2026
e7aeb32
Merge branch 'main' into studio-fla-tilelang-qwen3.5
danielhanchen May 15, 2026
1f32279
Merge branch 'main' into studio-fla-tilelang-qwen3.5
danielhanchen May 15, 2026
d56313e
Merge branch 'main' into studio-fla-tilelang-qwen3.5
danielhanchen May 15, 2026
f9b3d26
Merge branch 'main' into studio-fla-tilelang-qwen3.5
danielhanchen May 15, 2026
a8a15f7
ci: retrigger MLX dispatch after pytorch CDN DNS flake
danielhanchen May 15, 2026
2f9a6b0
Merge branch 'main' into studio-fla-tilelang-qwen3.5
danielhanchen May 15, 2026
6760682
Merge branch 'main' into studio-fla-tilelang-qwen3.5
danielhanchen May 15, 2026
17d4213
Merge branch 'main' into studio-fla-tilelang-qwen3.5
danielhanchen May 15, 2026
ec9e643
Merge branch 'main' into studio-fla-tilelang-qwen3.5
danielhanchen May 15, 2026
6681421
Merge branch 'main' into studio-fla-tilelang-qwen3.5
danielhanchen May 15, 2026
a415f6b
Merge branch 'main' into studio-fla-tilelang-qwen3.5
danielhanchen May 16, 2026
1a4df61
Merge branch 'main' into studio-fla-tilelang-qwen3.5
danielhanchen May 16, 2026
5ed13a9
Merge branch 'main' into studio-fla-tilelang-qwen3.5
danielhanchen May 16, 2026
fc2ee99
Merge branch 'main' into studio-fla-tilelang-qwen3.5
danielhanchen May 16, 2026
3fde343
studio: harden FLA + tilelang installers per reviewer feedback
danielhanchen May 16, 2026
d137a67
[pre-commit.ci] auto fixes from pre-commit.com hooks
pre-commit-ci[bot] May 16, 2026
0f246a6
studio: address reviewer.py P1/P2 findings on FLA + tilelang installers
danielhanchen May 16, 2026
27dc546
[pre-commit.ci] auto fixes from pre-commit.com hooks
pre-commit-ci[bot] May 16, 2026
66dface
studio: pin packaging + triton with FLA --no-deps install
danielhanchen May 16, 2026
6ce495a
studio: hook transformers' fast-path gates for just-in-time FLA + cau…
danielhanchen May 16, 2026
d2d758d
[pre-commit.ci] auto fixes from pre-commit.com hooks
pre-commit-ci[bot] May 16, 2026
8e7859d
Merge remote-tracking branch 'origin/main' into studio-fla-tilelang-q…
danielhanchen May 17, 2026
29e9f31
studio: address reviewer.py n=12 findings on the FLA hook path
danielhanchen May 17, 2026
6981149
Merge branch 'studio-fla-tilelang-qwen3.5' of https://github.com/unsl…
danielhanchen May 17, 2026
78b07a2
[pre-commit.ci] auto fixes from pre-commit.com hooks
pre-commit-ci[bot] May 17, 2026
85cdafc
studio: fix double-install of tilelang on the FLA hook install path
danielhanchen May 17, 2026
3913a66
Merge branch 'studio-fla-tilelang-qwen3.5' of https://github.com/unsl…
danielhanchen May 17, 2026
800fc98
[pre-commit.ci] auto fixes from pre-commit.com hooks
pre-commit-ci[bot] May 17, 2026
d73935c
ci: retrigger Mac Studio GGUF after transient HF DNS resolve flake
danielhanchen May 17, 2026
379cbb2
Merge branch 'studio-fla-tilelang-qwen3.5' of https://github.com/unsl…
danielhanchen May 17, 2026
10b50c8
studio: skip tilelang on HIP / ROCm torch (Strix Halo crash report)
danielhanchen May 17, 2026
038906c
[pre-commit.ci] auto fixes from pre-commit.com hooks
pre-commit-ci[bot] May 17, 2026
a4e63ec
ci: retrigger Windows Studio UI after transient Playwright tab-lookup…
danielhanchen May 17, 2026
73f7e32
Merge branch 'studio-fla-tilelang-qwen3.5' of https://github.com/unsl…
danielhanchen May 17, 2026
c358b05
studio: auto-discover FLA-using model types from installed transformers
danielhanchen May 17, 2026
5c2511d
[pre-commit.ci] auto fixes from pre-commit.com hooks
pre-commit-ci[bot] May 17, 2026
bb0e0b2
test: hermetize the non-allowlist hook test against transformers 5.4.0+
danielhanchen May 17, 2026
a68acdc
Merge branch 'studio-fla-tilelang-qwen3.5' of https://github.com/unsl…
danielhanchen May 17, 2026
3e60db3
ci: retrigger Windows Studio API after llama.cpp prebuilt staging Win…
danielhanchen May 17, 2026
a68c078
tests: move MLX smoke gate changes to dedicated PR #5537
danielhanchen May 18, 2026
f505f73
studio: friendlier install banners (drop hook / gate-name jargon)
danielhanchen May 18, 2026
1844402
Merge remote-tracking branch 'origin/main' into studio-fla-tilelang-q…
danielhanchen May 18, 2026
dcfb47c
[pre-commit.ci] auto fixes from pre-commit.com hooks
pre-commit-ci[bot] May 18, 2026
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
157 changes: 155 additions & 2 deletions studio/backend/core/training/worker.py
Original file line number Diff line number Diff line change
Expand Up @@ -52,6 +52,14 @@ def _output_dir_from_resume_checkpoint(
_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.
_TILELANG_PACKAGE_VERSION = "0.1.8"

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P2 Badge Install a TileLang version FLA declares compatible

For a fresh Qwen3.5-family Studio job this pins tilelang to 0.1.8, but the unpinned flash-linear-attention being installed just above currently declares its TileLang extra as tilelang>=0.1.9 in upstream pyproject.toml. That leaves the new TileLang backend on a version FLA does not claim to support, so the intended chunk_bwd_dqkwg / parallel_attn_* TileLang dispatch can fail or silently fall back despite the helper reporting the backend installed; pin a compatible TileLang release while separately constraining apache-tvm-ffi to the safe version.

Useful? React with 👍 / 👎.

_APACHE_TVM_FFI_PACKAGE_VERSION = "0.1.9"
_TILELANG_SKIP_ENV = "UNSLOTH_STUDIO_SKIP_TILELANG_INSTALL"


def _model_wants_causal_conv1d(model_name: str) -> bool:
Expand Down Expand Up @@ -275,6 +283,34 @@ def _ensure_causal_conv1d_fast_path(event_queue: Any, model_name: str) -> None:
)


def _ensure_flash_linear_attention(event_queue: Any, model_name: str) -> None:
"""Install ``flash-linear-attention`` from PyPI for models that need it.

Qwen3.5 / Qwen3.6 / Qwen3-Next (and the SSM hybrids covered by
``_model_wants_causal_conv1d``) gate their fast path on FLA's
``chunk_gated_delta_rule`` / ``fused_recurrent_gated_delta_rule``
being importable. Without FLA, transformers falls back to a pure
Python torch loop (~2.35x slower in our Qwen3.5-2B-Vision bench).

FLA ships as a universal py3-none-any wheel on PyPI (Triton kernels
JIT-compile at runtime), so no wheel-matching dance is needed.
"""
if not _model_wants_causal_conv1d(model_name):
return

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P2 Badge Avoid installing FLA for SSM-only models

When model_name is Falcon-H1/Nemotron-H/Granite-H/LFM2, this predicate is true because _model_wants_causal_conv1d includes those substrings, so _ensure_flash_linear_attention now attempts pip install flash-linear-attention before the mamba setup. The new TileLang comments below explicitly say those true SSM models do not go through FLA's gated_delta_rule, so on machines where the existing causal-conv1d/mamba dependencies are already present this adds an unnecessary network/dependency mutation path, with the helper's unbounded pip wait, for no fast-path benefit; restrict the FLA predicate to the Qwen GDN families.

Useful? React with 👍 / 👎.


_install_package_wheel_first(
event_queue = event_queue,
import_name = "fla",
display_name = "flash-linear-attention",
pypi_name = "flash-linear-attention",
wheel_url_builder = lambda env: None,
pypi_spec = "flash-linear-attention",

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P1 Badge Prevent FLA from upgrading the torch stack

When a Studio install is on one of the supported torch 2.4/2.5/2.6 environments (for example install.sh defaults to torch>=2.4,<2.11.0), this bare pip install flash-linear-attention lets pip resolve the current fla-core dependency on torch>=2.7.0 by replacing torch/triton/torchvision inside the Studio venv before training starts. That can silently move users off the CUDA/ROCm wheel set selected by the installer and break unrelated training jobs; install FLA without dependencies or gate it on an already-compatible torch stack.

Useful? React with 👍 / 👎.

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P2 Badge Bypass the ROCm source-build guard for FLA

When a Qwen3.5-family job runs on a ROCm runtime image where probe_torch_wheel_env() reports hip_version but hipcc is not installed, this pure-PyPI FLA install still goes through _install_package_wheel_first; because wheel_url_builder returns None, the helper reaches its generic HIP guard and returns before running pip. The helper comment says flash-linear-attention is a universal py3-none-any wheel, so these ROCm jobs silently miss the FLA fast path even though a normal pip install could proceed; use a direct install path or an option that skips the hipcc source-build check for pure-Python packages.

Useful? React with 👍 / 👎.

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P2 Badge Guard FLA install on supported Python

For Studio runs under Python 3.9, which this project still declares as supported in pyproject.toml, this unpinned flash-linear-attention install cannot succeed because current PyPI releases require Python >=3.10. A Qwen3.5/3.6 job in that environment will try and fail this pip install on every launch, then continue without the intended FLA fast path while still proceeding to the TileLang install; add a Python-version guard or a compatible pinned package path before attempting the install.

Useful? React with 👍 / 👎.

pypi_status_message = (
"Installing flash-linear-attention from PyPI for the fast path..."
),
)


_SSM_MODEL_SUBSTRINGS = (
"nemotron_h",
"nemotron-h",
Expand Down Expand Up @@ -303,6 +339,111 @@ 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",
)


def _model_wants_tilelang(model_name: str) -> bool:
name = model_name.lower()
return any(sub in name for sub in _TILELANG_MODEL_SUBSTRINGS)


def _ensure_tilelang_backend(event_queue: Any, model_name: str) -> None:
"""Install ``tilelang`` + pinned ``apache-tvm-ffi`` for FLA's TileLang backend.

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/
0.1.11, but those versions hit a "CUDA: misaligned address" crash in
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

try:
import tilelang # noqa: F401
import tvm_ffi # noqa: F401

logger.info("tilelang + apache-tvm-ffi already installed")
return

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P1 Badge Reinstall pinned TileLang deps when unsafe versions are present

For Qwen3.5-family jobs on a machine that already has tilelang plus apache-tvm-ffi 0.1.10/0.1.11 installed, this import-only check returns without applying the pinned apache-tvm-ffi==0.1.9 pair. The new helper’s own comment says those newer apache-tvm-ffi versions crash with CUDA misaligned-address errors on Blackwell, so an existing Studio/container environment with the bad version still takes the broken TileLang path instead of being downgraded; check installed package versions before returning here.

Useful? React with 👍 / 👎.

except ImportError:
pass

_send_status(
event_queue,
(
f"Installing TileLang backend ("
f"apache-tvm-ffi=={_APACHE_TVM_FFI_PACKAGE_VERSION}, "
f"tilelang=={_TILELANG_PACKAGE_VERSION}) for FLA fast path..."
),
)

# Install both in one pip resolve so the apache-tvm-ffi pin wins over
# tilelang's `>=0.1.2,~=0.1.0` constraint.
specs = [
f"apache-tvm-ffi=={_APACHE_TVM_FFI_PACKAGE_VERSION}",
f"tilelang=={_TILELANG_PACKAGE_VERSION}",
]
if shutil.which("uv"):
pypi_cmd = [
"uv",
"pip",
"install",
"--python",
sys.executable,
*specs,
]
else:
pypi_cmd = [
sys.executable,
"-m",
"pip",
"install",
*specs,

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P2 Badge Guard TileLang installs to wheel-supported platforms

For Qwen3.5-family training on Windows or any platform without a tilelang==0.1.8 wheel, this unqualified install lets pip/uv fall back to the 93MB TileLang sdist even though the backend is optional and failures are swallowed below. PyPI metadata for this pinned version only publishes Linux x86_64/aarch64 and macOS arm64 wheels, while Studio also has a Windows setup path, so those jobs can spend a long time downloading/building before continuing without TileLang; add a platform/wheel guard or force binary-only/fail-fast installation for this optional backend.

Useful? React with 👍 / 👎.

]

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

medium

The subprocess.run call here lacks a timeout. While consistent with other non-HIP installs in this file, a network hang during pip install could block the training worker indefinitely. Consider adding a generous timeout (e.g., 300s or 600s) to ensure the process can recover or fail gracefully if the network is unresponsive.


result = _sp.run(
pypi_cmd,
stdout = _sp.PIPE,
stderr = _sp.STDOUT,
text = True,
)
if result.returncode != 0:
logger.warning(
"TileLang backend install failed (continuing without it):\n%s",
result.stdout,
)
_send_status(
event_queue,
"TileLang backend install failed; continuing on the FLA Triton path",
)
return

logger.info("Installed TileLang backend for FLA fast path")


def _should_try_runtime_flash_attn_install(max_seq_length: int) -> bool:
if os.getenv(_FLASH_ATTN_SKIP_ENV) == "1":
return False
Expand Down Expand Up @@ -1111,10 +1252,20 @@ def run_training_process(
model_name,
)

# ── 1b. Set up causal-conv1d first, then install mamba-ssm if needed ──
# ── 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)
try:
_ensure_causal_conv1d_fast_path(event_queue, model_name)
_ensure_flash_linear_attention(event_queue, model_name)

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P2 Badge Install TileLang before probing FLA

On machines where flash-linear-attention is already installed but TileLang is not, this call probes availability by importing fla before _ensure_tilelang_backend runs. FLA's top-level import initializes its layers/models and backend dispatch machinery based on the backends importable at that time, so installing TileLang later in the same training subprocess can leave the Qwen3.5 job on the Triton backend until a fresh process starts. Install the TileLang pair before the FLA import probe, or explicitly refresh/reload the FLA backend registration after the TileLang install.

Useful? React with 👍 / 👎.

_ensure_mamba_ssm(event_queue, model_name)
Comment on lines +1779 to 1781

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P2 Badge Restore causal-conv1d setup for hook-enabled SSM path

When UNSLOTH_STUDIO_SKIP_FAST_PATH_HOOKS is not set, this branch installs hooks but no longer calls _ensure_causal_conv1d_fast_path before _ensure_mamba_ssm. I checked the upstream transformers SSM model files (modeling_falcon_h1.py and modeling_nemotron_h.py), and they load causal-conv1d via lazy_load_kernel("causal-conv1d") instead of calling is_causal_conv1d_available, so the new hook never triggers for those families; on a fresh worker they can now run without causal-conv1d installed and fall back to the slower path that this installer previously avoided.

Useful? React with 👍 / 👎.

_ensure_tilelang_backend(event_queue, model_name)

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P2 Badge Skip TileLang install when FLA is not eligible

Gatekeeping for flash-linear-attention now skips installation on unsupported stacks (for example torch <2.7 or UNSLOTH_STUDIO_SKIP_FLA_INSTALL=1), but _ensure_tilelang_backend is still called unconditionally afterward. In those environments, TileLang cannot be used (it is only a backend for FLA), so this still triggers a network/dependency mutation path (apache-tvm-ffi/tilelang install or reinstall) with no runtime benefit and potential environment churn; add a shared eligibility guard so TileLang is skipped whenever FLA is unavailable.

Useful? React with 👍 / 👎.

_ensure_flash_attn_for_long_context(
event_queue,
int(config.get("max_seq_length", 2048)),
Expand All @@ -1125,7 +1276,9 @@ def run_training_process(
"type": "error",
"error": (
f"Please choose another model to train, since "
f"causal-conv1d / mamba-ssm failed to install "
f"a fast-path kernel library "
f"(causal-conv1d / flash-linear-attention / "
f"mamba-ssm / tilelang) failed to install "
f"with error: {exc}"
),
"stack": traceback.format_exc(limit = 20),
Expand Down
138 changes: 138 additions & 0 deletions studio/backend/tests/test_training_worker_flash_attn.py
Original file line number Diff line number Diff line change
Expand Up @@ -193,3 +193,141 @@ def test_mamba_ssm_path_preserves_wheel_first_install_args(monkeypatch):
release_tag = worker._MAMBA_SSM_RELEASE_TAG,
release_base_url = "https://github.com/state-spaces/mamba/releases/download",
)


def test_flash_linear_attention_uses_pypi_for_qwen3_5(monkeypatch):
install_mock = mock.Mock(return_value = True)
monkeypatch.setattr(worker, "_install_package_wheel_first", install_mock)

worker._ensure_flash_linear_attention(
event_queue = [],
model_name = "unsloth/Qwen3.5-2B",
)

install_mock.assert_called_once()
_, kwargs = install_mock.call_args
assert kwargs["import_name"] == "fla"
assert kwargs["display_name"] == "flash-linear-attention"
assert kwargs["pypi_name"] == "flash-linear-attention"
assert kwargs["pypi_spec"] == "flash-linear-attention"
# Pure-Python wheel from PyPI: no version pin, no github wheel lookup.
assert "pypi_version" not in kwargs or kwargs["pypi_version"] is None
assert callable(kwargs["wheel_url_builder"])
assert kwargs["wheel_url_builder"](None) is None


def test_flash_linear_attention_skips_for_unrelated_models(monkeypatch):
install_mock = mock.Mock(return_value = True)
monkeypatch.setattr(worker, "_install_package_wheel_first", install_mock)

worker._ensure_flash_linear_attention(
event_queue = [],
model_name = "meta-llama/Llama-3.2-1B-Instruct",
)

install_mock.assert_not_called()


def test_flash_linear_attention_matches_full_qwen3_family(monkeypatch):
install_mock = mock.Mock(return_value = True)
monkeypatch.setattr(worker, "_install_package_wheel_first", install_mock)

for name in (
"unsloth/Qwen3.5-2B",
"unsloth/Qwen3_5-MoE-A22B",
"unsloth/Qwen3.6-4B",
"unsloth/Qwen3_6-4B",
"unsloth/Qwen3-Next-80B-A3B",
"unsloth/Qwen3_Next-80B-A3B",
):
worker._ensure_flash_linear_attention(event_queue = [], model_name = name)

assert install_mock.call_count == 6


def test_tilelang_backend_installs_pinned_pair_for_qwen3_5(monkeypatch):
monkeypatch.delenv(worker._TILELANG_SKIP_ENV, raising = False)
monkeypatch.setattr(worker.shutil, "which", lambda name: "/usr/bin/uv")
run_mock = mock.Mock(return_value = mock.Mock(returncode = 0, stdout = ""))
monkeypatch.setattr(worker._sp, "run", run_mock)

# Force the "not installed" branch by making the imports fail.
real_import = builtins.__import__

def fake_import(name, *a, **kw):
if name in ("tilelang", "tvm_ffi"):
raise ImportError
return real_import(name, *a, **kw)

monkeypatch.setattr(builtins, "__import__", fake_import)
statuses: list[str] = []
monkeypatch.setattr(worker, "_send_status", lambda queue, msg: statuses.append(msg))

worker._ensure_tilelang_backend(
event_queue = [],
model_name = "unsloth/Qwen3.5-2B",
)

run_mock.assert_called_once()
args = run_mock.call_args[0][0]
assert f"apache-tvm-ffi=={worker._APACHE_TVM_FFI_PACKAGE_VERSION}" in args
assert f"tilelang=={worker._TILELANG_PACKAGE_VERSION}" in args
assert any("TileLang backend" in s for s in statuses)


def test_tilelang_backend_skipped_for_ssm_models(monkeypatch):
monkeypatch.delenv(worker._TILELANG_SKIP_ENV, raising = False)
run_mock = mock.Mock(return_value = mock.Mock(returncode = 0, stdout = ""))
monkeypatch.setattr(worker._sp, "run", run_mock)

# Nemotron-H / Falcon-H1 / Granite-H take the mamba_ssm path, not FLA's
# gated_delta_rule -> tilelang has no effect on them.
for name in (
"tiiuae/Falcon-H1-0.5B-Instruct",
"nvidia/Nemotron-H-8B-Base",
"ibm-granite/granite-4.0-h-tiny",
"meta-llama/Llama-3.2-1B-Instruct",
):
worker._ensure_tilelang_backend(event_queue = [], model_name = name)

run_mock.assert_not_called()


def test_tilelang_backend_skipped_via_env(monkeypatch):
monkeypatch.setenv(worker._TILELANG_SKIP_ENV, "1")
run_mock = mock.Mock(return_value = mock.Mock(returncode = 0, stdout = ""))
monkeypatch.setattr(worker._sp, "run", run_mock)

worker._ensure_tilelang_backend(
event_queue = [],
model_name = "unsloth/Qwen3.5-2B",
)

run_mock.assert_not_called()


def test_tilelang_backend_swallows_install_failure(monkeypatch):
monkeypatch.delenv(worker._TILELANG_SKIP_ENV, raising = False)
monkeypatch.setattr(worker.shutil, "which", lambda name: None)
run_mock = mock.Mock(return_value = mock.Mock(returncode = 1, stdout = "boom"))
monkeypatch.setattr(worker._sp, "run", run_mock)

real_import = builtins.__import__

def fake_import(name, *a, **kw):
if name in ("tilelang", "tvm_ffi"):
raise ImportError
return real_import(name, *a, **kw)

monkeypatch.setattr(builtins, "__import__", fake_import)
statuses: list[str] = []
monkeypatch.setattr(worker, "_send_status", lambda queue, msg: statuses.append(msg))

# Should not raise even when pip exits non-zero.
worker._ensure_tilelang_backend(
event_queue = [],
model_name = "unsloth/Qwen3.5-2B",
)

run_mock.assert_called_once()
assert any("failed" in s.lower() for s in statuses)
Loading