studio: install flash-linear-attention and tilelang for Qwen3.5 family

Studio currently only installs causal-conv1d for qwen3.5 / qwen3.6 /
qwen3-next models. Without flash-linear-attention installed alongside
it, transformers' Qwen3.5 fast-path gate stays False and the model
falls back to a pure-PyTorch loop for the GatedDeltaNet layers. In a
60-step run on unsloth/Qwen3.5-2B on B200, this fallback costs ~2.35x
vs the full fast path.

On top of that, FLA dispatches its hottest GDN kernels through a
TileLang backend when tilelang is importable. Adding tilelang plus a
pinned apache-tvm-ffi gives another ~26% on the same workload (4.73
s/step to 3.50 s/step) and is what users have been getting indirectly
when they install mamba-ssm (mamba-ssm transitively pulls tilelang and
pins apache-tvm-ffi<=0.1.9, which is the last working version on
sm_100; 0.1.10 and 0.1.11 crash Triton with misaligned address).

Changes:
  * _ensure_flash_linear_attention: pure-Python PyPI install gated on
    the same model match set as _ensure_causal_conv1d_fast_path.
  * _ensure_tilelang_backend: installs apache-tvm-ffi==0.1.9 and
    tilelang==0.1.8 in one pip resolve so the tvm-ffi pin wins over
    tilelang's >=0.1.2 constraint. Gated on the Qwen3.5 family only;
    SSM models (Nemotron-H, Falcon-H1, Granite-H, LFM2) do not use
    FLA's GDN dispatch.
  * UNSLOTH_STUDIO_SKIP_TILELANG_INSTALL=1 escape hatch matching the
    flash-attn pattern.
  * Orchestration block reordered: causal-conv1d -> fla -> mamba-ssm
    -> tilelang -> flash-attn (long context).
  * 7 new tests covering the new helpers, including SSM-model skip,
    skip-env, full Qwen3 family name variants, and graceful pip
    install failure.

Combined Qwen3.5-2B-Vision step time on B200 in our bench goes from
5.0 s/step (current Studio: causal-conv1d only) to 3.5 s/step
(causal-conv1d + fla + tilelang), a 1.43x speedup with no notebook
or user code changes required.
This commit is contained in:
danielhanchen 2026-05-15 05:26:03 +00:00
commit c07cddae35
2 changed files with 286 additions and 2 deletions

View file

@ -52,6 +52,14 @@ _MAMBA_SSM_RELEASE_TAG = "v2.3.1"
_MAMBA_SSM_PACKAGE_VERSION = "2.3.1"
_FLASH_ATTN_RUNTIME_MIN_SEQ_LEN = 32768
_FLASH_ATTN_SKIP_ENV = "UNSLOTH_STUDIO_SKIP_FLASHATTN_INSTALL"
# tilelang 0.1.9+ pairs with apache-tvm-ffi >=0.1.10 by default, but
# apache-tvm-ffi 0.1.10/0.1.11 has an alignment regression that crashes
# subsequent Triton kernels with "CUDA: misaligned address" on sm_100
# (Blackwell). 0.1.9 is the last known-good. mamba_ssm 2.3.2 also pins
# apache-tvm-ffi<=0.1.9, which is the original source of this pin.
_TILELANG_PACKAGE_VERSION = "0.1.8"
_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:
@ -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
_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..."
),
)
_SSM_MODEL_SUBSTRINGS = (
"nemotron_h",
"nemotron-h",
@ -303,6 +339,104 @@ 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
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,
]
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
@ -1111,10 +1245,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)
_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)),
@ -1125,7 +1269,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),

View file

@ -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)