studio: harden FLA + tilelang installers per reviewer feedback

Addresses bot review on #5434:

  * Narrow `_ensure_flash_linear_attention` from `_model_wants_causal_conv1d`
    (which also matches Nemotron-H / Falcon-H1 / Granite-H / LFM2) to
    `_model_wants_tilelang` (Qwen3.5 / Qwen3.6 / Qwen3-Next only). True
    SSM families take the mamba_ssm path and never call FLA's GDN
    kernels, so installing FLA there is wasted bandwidth.

  * Pin both `flash-linear-attention==0.5.0` and `fla-core==0.5.0` and
    install with `--no-deps`. Otherwise pip resolves fla-core's
    declared `torch>=2.7.0` requirement and may silently upgrade the
    Studio venv's torch on environments running torch 2.4/2.5/2.6.

  * Skip both installs on Python <3.10 (FLA, fla-core, and tilelang
    all declare `Requires-Python: >=3.10`). On older interpreters the
    pip install would fail every launch and leave the worker on the
    slow torch fallback while still claiming to have set up the fast
    path.

  * Skip tilelang install on non-Linux platforms. `tilelang==0.1.8`
    only publishes Linux x86_64 / aarch64 and macOS arm64 wheels.
    Falling back to its 93MB sdist on a Studio worker is undesirable.

  * Detect an existing `apache-tvm-ffi` 0.1.10 / 0.1.11 install and
    force a reinstall to 0.1.9 with `--force-reinstall --no-deps`.
    Previously the import-only probe returned early and left the
    broken version in place, which crashes Triton on sm_100.

  * Add a 600s timeout to the tilelang and FLA subprocess.run calls,
    matching the existing flash-attn install pattern, so a network
    hang cannot block the training subprocess indefinitely.

  * 13 new / updated tests covering all six guards plus the
    pinned-spec, timeout, and force-reinstall code paths.

Total: 21 passing tests (8 original + 13 new / updated).
This commit is contained in:
danielhanchen 2026-05-16 07:03:26 +00:00
commit 3fde3439e8
2 changed files with 296 additions and 72 deletions

View file

@ -60,6 +60,20 @@ _FLASH_ATTN_SKIP_ENV = "UNSLOTH_STUDIO_SKIP_FLASHATTN_INSTALL"
_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.
_FLA_PACKAGE_VERSION = "0.5.0"
_FLA_CORE_PACKAGE_VERSION = "0.5.0"
# flash-linear-attention and tilelang both require Python >=3.10.
_FLA_MIN_PYTHON = (3, 10)
# tilelang wheels exist for Linux x86_64/aarch64 and macOS arm64. We
# never want to fall back to its 93MB sdist on a Studio worker, so
# skip on platforms outside that set.
_TILELANG_SUPPORTED_PLATFORMS = ("linux",)
_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")
def _model_wants_causal_conv1d(model_name: str) -> bool:
@ -284,32 +298,89 @@ 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.
"""Install ``flash-linear-attention`` + ``fla-core`` for Qwen3.5 family.
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).
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).
FLA ships as a universal py3-none-any wheel on PyPI (Triton kernels
JIT-compile at runtime), so no wheel-matching dance is needed.
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.
Both packages are pure-Python wheels on PyPI. We install with
``--no-deps`` to prevent fla-core's ``torch>=2.7.0`` requirement
from silently upgrading the Studio venv's torch.
"""
if not _model_wants_causal_conv1d(model_name):
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",
_FLA_MIN_PYTHON[0], _FLA_MIN_PYTHON[1], sys.version.split()[0],
)
return
_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",
pypi_status_message = (
"Installing flash-linear-attention from PyPI for the fast path..."
try:
import fla.modules # noqa: F401
import fla.ops.gated_delta_rule # noqa: F401
logger.info("flash-linear-attention already importable")
return
except ImportError:
pass
_send_status(
event_queue,
(
f"Installing flash-linear-attention=={_FLA_PACKAGE_VERSION} "
f"(with fla-core=={_FLA_CORE_PACKAGE_VERSION}) for the fast path..."
),
)
specs = [
f"fla-core=={_FLA_CORE_PACKAGE_VERSION}",
f"flash-linear-attention=={_FLA_PACKAGE_VERSION}",
]
if shutil.which("uv"):
pypi_cmd = [
"uv", "pip", "install",
"--python", sys.executable,
"--no-deps",
*specs,
]
else:
pypi_cmd = [
sys.executable, "-m", "pip", "install",
"--no-deps",
*specs,
]
try:
result = _sp.run(
pypi_cmd,
stdout = _sp.PIPE,
stderr = _sp.STDOUT,
text = True,
timeout = _TILELANG_INSTALL_TIMEOUT_S,
)
except _sp.TimeoutExpired:
logger.warning("flash-linear-attention install timed out; continuing")
_send_status(event_queue, "flash-linear-attention install timed out; continuing")
return
if result.returncode != 0:
logger.warning(
"flash-linear-attention install failed (continuing on torch fallback):\n%s",
result.stdout,
)
_send_status(
event_queue,
"flash-linear-attention install failed; continuing on torch fallback",
)
return
logger.info("Installed flash-linear-attention for the FLA fast path")
_SSM_MODEL_SUBSTRINGS = (
"nemotron_h",
@ -363,6 +434,19 @@ def _model_wants_tilelang(model_name: str) -> bool:
return any(sub in name for sub in _TILELANG_MODEL_SUBSTRINGS)
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.
"""
try:
from importlib.metadata import version as _pkg_version
return _pkg_version("apache-tvm-ffi")
except Exception:
return None
def _ensure_tilelang_backend(event_queue: Any, model_name: str) -> None:
"""Install ``tilelang`` + pinned ``apache-tvm-ffi`` for FLA's TileLang backend.
@ -381,15 +465,35 @@ def _ensure_tilelang_backend(event_queue: Any, model_name: str) -> None:
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")
if sys.version_info < _FLA_MIN_PYTHON:
logger.info(
"Skipping tilelang install: requires Python >= %d.%d, have %s",
_FLA_MIN_PYTHON[0], _FLA_MIN_PYTHON[1], sys.version.split()[0],
)
return
except ImportError:
pass
if not any(sys.platform.startswith(p) for p in _TILELANG_SUPPORTED_PLATFORMS):
logger.info(
"Skipping tilelang install: no prebuilt wheel for platform %s",
sys.platform,
)
return
existing_tvm_ffi = _installed_tvm_ffi_version()
needs_reinstall = existing_tvm_ffi in _TVM_FFI_BROKEN_VERSIONS
if not needs_reinstall:
try:
import tilelang # noqa: F401
import tvm_ffi # noqa: F401
logger.info("tilelang + apache-tvm-ffi already installed")
return
except ImportError:
pass
else:
logger.info(
"Forcing tilelang reinstall: apache-tvm-ffi %s is on the broken list",
existing_tvm_ffi,
)
_send_status(
event_queue,
@ -406,30 +510,34 @@ def _ensure_tilelang_backend(event_queue: Any, model_name: str) -> None:
f"apache-tvm-ffi=={_APACHE_TVM_FFI_PACKAGE_VERSION}",
f"tilelang=={_TILELANG_PACKAGE_VERSION}",
]
extra_args = ["--force-reinstall", "--no-deps"] if needs_reinstall else []
if shutil.which("uv"):
pypi_cmd = [
"uv",
"pip",
"install",
"--python",
sys.executable,
"uv", "pip", "install",
"--python", sys.executable,
*extra_args,
*specs,
]
else:
pypi_cmd = [
sys.executable,
"-m",
"pip",
"install",
sys.executable, "-m", "pip", "install",
*extra_args,
*specs,
]
result = _sp.run(
pypi_cmd,
stdout = _sp.PIPE,
stderr = _sp.STDOUT,
text = True,
)
try:
result = _sp.run(
pypi_cmd,
stdout = _sp.PIPE,
stderr = _sp.STDOUT,
text = True,
timeout = _TILELANG_INSTALL_TIMEOUT_S,
)
except _sp.TimeoutExpired:
logger.warning("TileLang backend install timed out; continuing")
_send_status(event_queue, "TileLang backend install timed out; continuing")
return
if result.returncode != 0:
logger.warning(
"TileLang backend install failed (continuing without it):\n%s",

View file

@ -195,42 +195,75 @@ def test_mamba_ssm_path_preserves_wheel_first_install_args(monkeypatch):
)
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)
def _force_missing_fla_imports(monkeypatch):
"""Make fla.modules / fla.ops.gated_delta_rule imports raise ImportError."""
real_import = builtins.__import__
def fake_import(name, *a, **kw):
if name.startswith("fla.modules") or name.startswith("fla.ops"):
raise ImportError
return real_import(name, *a, **kw)
monkeypatch.setattr(builtins, "__import__", fake_import)
def test_flash_linear_attention_installs_pinned_pair_for_qwen3_5(monkeypatch):
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_missing_fla_imports(monkeypatch)
statuses: list[str] = []
monkeypatch.setattr(worker, "_send_status", lambda queue, msg: statuses.append(msg))
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
run_mock.assert_called_once()
args = run_mock.call_args[0][0]
assert f"flash-linear-attention=={worker._FLA_PACKAGE_VERSION}" in args
assert f"fla-core=={worker._FLA_CORE_PACKAGE_VERSION}" in args
assert "--no-deps" in args
assert run_mock.call_args.kwargs["timeout"] == worker._TILELANG_INSTALL_TIMEOUT_S
assert any("flash-linear-attention" in s for s in statuses)
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)
run_mock = mock.Mock(return_value = mock.Mock(returncode = 0, stdout = ""))
monkeypatch.setattr(worker._sp, "run", run_mock)
worker._ensure_flash_linear_attention(
event_queue = [],
model_name = "meta-llama/Llama-3.2-1B-Instruct",
)
install_mock.assert_not_called()
run_mock.assert_not_called()
def test_flash_linear_attention_skips_for_ssm_only_models(monkeypatch):
# Nemotron-H / Falcon-H1 / Granite-H / LFM2 take the mamba_ssm path
# and never call FLA's gated_delta_rule kernels.
run_mock = mock.Mock(return_value = mock.Mock(returncode = 0, stdout = ""))
monkeypatch.setattr(worker._sp, "run", run_mock)
for name in (
"tiiuae/Falcon-H1-0.5B-Instruct",
"nvidia/Nemotron-H-8B-Base",
"ibm-granite/granite-4.0-h-tiny",
"LiquidAI/LFM2-1.2B-Instruct",
):
worker._ensure_flash_linear_attention(event_queue = [], model_name = name)
run_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)
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_missing_fla_imports(monkeypatch)
monkeypatch.setattr(worker, "_send_status", lambda *a, **k: None)
for name in (
"unsloth/Qwen3.5-2B",
@ -242,16 +275,25 @@ def test_flash_linear_attention_matches_full_qwen3_family(monkeypatch):
):
worker._ensure_flash_linear_attention(event_queue = [], model_name = name)
assert install_mock.call_count == 6
assert run_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")
def test_flash_linear_attention_skipped_below_python_3_10(monkeypatch):
# sys.version_info is a structseq, not constructible; substitute a
# plain tuple so the `< _FLA_MIN_PYTHON` comparison still works.
monkeypatch.setattr(worker.sys, "version_info", (3, 9, 0, "final", 0))
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.
worker._ensure_flash_linear_attention(
event_queue = [],
model_name = "unsloth/Qwen3.5-2B",
)
run_mock.assert_not_called()
def _force_missing_tilelang_imports(monkeypatch):
real_import = builtins.__import__
def fake_import(name, *a, **kw):
@ -260,6 +302,15 @@ def test_tilelang_backend_installs_pinned_pair_for_qwen3_5(monkeypatch):
return real_import(name, *a, **kw)
monkeypatch.setattr(builtins, "__import__", fake_import)
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")
monkeypatch.setattr(worker, "_installed_tvm_ffi_version", lambda: None)
run_mock = mock.Mock(return_value = mock.Mock(returncode = 0, stdout = ""))
monkeypatch.setattr(worker._sp, "run", run_mock)
_force_missing_tilelang_imports(monkeypatch)
statuses: list[str] = []
monkeypatch.setattr(worker, "_send_status", lambda queue, msg: statuses.append(msg))
@ -272,9 +323,81 @@ def test_tilelang_backend_installs_pinned_pair_for_qwen3_5(monkeypatch):
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 run_mock.call_args.kwargs["timeout"] == worker._TILELANG_INSTALL_TIMEOUT_S
assert any("TileLang backend" in s for s in statuses)
def test_tilelang_backend_reinstalls_when_tvm_ffi_is_broken(monkeypatch):
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.11")
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",
)
run_mock.assert_called_once()
args = run_mock.call_args[0][0]
assert "--force-reinstall" in args
assert "--no-deps" in args
def test_tilelang_backend_skipped_below_python_3_10(monkeypatch):
monkeypatch.delenv(worker._TILELANG_SKIP_ENV, raising = False)
# sys.version_info is a structseq, not constructible; substitute a
# plain tuple so the `< _FLA_MIN_PYTHON` comparison still works.
monkeypatch.setattr(worker.sys, "version_info", (3, 9, 0, "final", 0))
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_skipped_on_windows(monkeypatch):
monkeypatch.delenv(worker._TILELANG_SKIP_ENV, raising = False)
monkeypatch.setattr(worker.sys, "platform", "win32")
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_timeout(monkeypatch):
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: None)
_force_missing_tilelang_imports(monkeypatch)
def raise_timeout(*a, **kw):
raise subprocess.TimeoutExpired(cmd = "pip", timeout = 1)
monkeypatch.setattr(worker._sp, "run", raise_timeout)
statuses: list[str] = []
monkeypatch.setattr(worker, "_send_status", lambda queue, msg: statuses.append(msg))
# Should not raise.
worker._ensure_tilelang_backend(
event_queue = [],
model_name = "unsloth/Qwen3.5-2B",
)
assert any("timed out" in s.lower() 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 = ""))
@ -309,17 +432,10 @@ def test_tilelang_backend_skipped_via_env(monkeypatch):
def test_tilelang_backend_swallows_install_failure(monkeypatch):
monkeypatch.delenv(worker._TILELANG_SKIP_ENV, raising = False)
monkeypatch.setattr(worker.shutil, "which", lambda name: None)
monkeypatch.setattr(worker, "_installed_tvm_ffi_version", lambda: 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)
_force_missing_tilelang_imports(monkeypatch)
statuses: list[str] = []
monkeypatch.setattr(worker, "_send_status", lambda queue, msg: statuses.append(msg))