Merge branch 'studio-fla-tilelang-qwen3.5' of https://github.com/unslothai/unsloth into studio-fla-tilelang-qwen3.5
This commit is contained in:
commit
73f7e32bfc
2 changed files with 33 additions and 17 deletions
|
|
@ -619,6 +619,7 @@ def _torch_has_hip() -> bool:
|
|||
"""
|
||||
try:
|
||||
import torch as _torch
|
||||
|
||||
return getattr(_torch.version, "hip", None) is not None
|
||||
except Exception:
|
||||
return False
|
||||
|
|
|
|||
|
|
@ -1243,13 +1243,13 @@ def test_tilelang_platform_unsupported_on_hip_torch(monkeypatch):
|
|||
|
||||
def test_tilelang_install_skipped_on_hip_torch(monkeypatch):
|
||||
"""End-to-end: the unconditional installer must not call pip on HIP torch."""
|
||||
monkeypatch.delenv(worker._TILELANG_SKIP_ENV, raising=False)
|
||||
monkeypatch.delenv(worker._TILELANG_SKIP_ENV, raising = False)
|
||||
monkeypatch.setattr(worker, "_torch_has_hip", lambda: True)
|
||||
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)
|
||||
|
||||
result = worker._ensure_tilelang_backend_unconditional(event_queue=[])
|
||||
result = worker._ensure_tilelang_backend_unconditional(event_queue = [])
|
||||
|
||||
assert result is False
|
||||
run_mock.assert_not_called()
|
||||
|
|
@ -1261,15 +1261,20 @@ def test_install_fast_path_hooks_sets_fla_tilelang_zero_on_hip(monkeypatch):
|
|||
PRE-EXISTING tilelang install isn't used by FLA's dispatcher.
|
||||
"""
|
||||
import os as _os
|
||||
monkeypatch.delenv("FLA_TILELANG", raising=False)
|
||||
monkeypatch.delenv(worker._FAST_PATH_HOOKS_SKIP_ENV, raising=False)
|
||||
|
||||
monkeypatch.delenv("FLA_TILELANG", raising = False)
|
||||
monkeypatch.delenv(worker._FAST_PATH_HOOKS_SKIP_ENV, raising = False)
|
||||
monkeypatch.setattr(worker, "_torch_has_hip", lambda: True)
|
||||
monkeypatch.setattr(worker, "_ensure_flash_linear_attention_unconditional", lambda eq: True)
|
||||
monkeypatch.setattr(worker, "_ensure_tilelang_backend_unconditional", lambda eq: True)
|
||||
monkeypatch.setattr(
|
||||
worker, "_ensure_flash_linear_attention_unconditional", lambda eq: True
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
worker, "_ensure_tilelang_backend_unconditional", lambda eq: True
|
||||
)
|
||||
monkeypatch.setattr(worker, "_install_package_wheel_first", lambda **kw: True)
|
||||
|
||||
worker._install_fast_path_hooks(
|
||||
event_queue=_FakeQueue(), model_name="unsloth/Qwen3.5-2B"
|
||||
event_queue = _FakeQueue(), model_name = "unsloth/Qwen3.5-2B"
|
||||
)
|
||||
|
||||
assert _os.environ.get("FLA_TILELANG") == "0"
|
||||
|
|
@ -1280,15 +1285,20 @@ def test_install_fast_path_hooks_respects_user_fla_tilelang_override(monkeypatch
|
|||
overwrite — they may know they have a HIP-aware tilelang fork.
|
||||
"""
|
||||
import os as _os
|
||||
|
||||
monkeypatch.setenv("FLA_TILELANG", "1")
|
||||
monkeypatch.delenv(worker._FAST_PATH_HOOKS_SKIP_ENV, raising=False)
|
||||
monkeypatch.delenv(worker._FAST_PATH_HOOKS_SKIP_ENV, raising = False)
|
||||
monkeypatch.setattr(worker, "_torch_has_hip", lambda: True)
|
||||
monkeypatch.setattr(worker, "_ensure_flash_linear_attention_unconditional", lambda eq: True)
|
||||
monkeypatch.setattr(worker, "_ensure_tilelang_backend_unconditional", lambda eq: True)
|
||||
monkeypatch.setattr(
|
||||
worker, "_ensure_flash_linear_attention_unconditional", lambda eq: True
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
worker, "_ensure_tilelang_backend_unconditional", lambda eq: True
|
||||
)
|
||||
monkeypatch.setattr(worker, "_install_package_wheel_first", lambda **kw: True)
|
||||
|
||||
worker._install_fast_path_hooks(
|
||||
event_queue=_FakeQueue(), model_name="unsloth/Qwen3.5-2B"
|
||||
event_queue = _FakeQueue(), model_name = "unsloth/Qwen3.5-2B"
|
||||
)
|
||||
|
||||
assert _os.environ["FLA_TILELANG"] == "1"
|
||||
|
|
@ -1297,15 +1307,20 @@ def test_install_fast_path_hooks_respects_user_fla_tilelang_override(monkeypatch
|
|||
def test_install_fast_path_hooks_does_not_set_fla_tilelang_on_cuda(monkeypatch):
|
||||
"""CUDA path must NOT set FLA_TILELANG (tilelang is wanted there)."""
|
||||
import os as _os
|
||||
monkeypatch.delenv("FLA_TILELANG", raising=False)
|
||||
monkeypatch.delenv(worker._FAST_PATH_HOOKS_SKIP_ENV, raising=False)
|
||||
|
||||
monkeypatch.delenv("FLA_TILELANG", raising = False)
|
||||
monkeypatch.delenv(worker._FAST_PATH_HOOKS_SKIP_ENV, raising = False)
|
||||
monkeypatch.setattr(worker, "_torch_has_hip", lambda: False)
|
||||
monkeypatch.setattr(worker, "_ensure_flash_linear_attention_unconditional", lambda eq: True)
|
||||
monkeypatch.setattr(worker, "_ensure_tilelang_backend_unconditional", lambda eq: True)
|
||||
monkeypatch.setattr(
|
||||
worker, "_ensure_flash_linear_attention_unconditional", lambda eq: True
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
worker, "_ensure_tilelang_backend_unconditional", lambda eq: True
|
||||
)
|
||||
monkeypatch.setattr(worker, "_install_package_wheel_first", lambda **kw: True)
|
||||
|
||||
worker._install_fast_path_hooks(
|
||||
event_queue=_FakeQueue(), model_name="unsloth/Qwen3.5-2B"
|
||||
event_queue = _FakeQueue(), model_name = "unsloth/Qwen3.5-2B"
|
||||
)
|
||||
|
||||
assert _os.environ.get("FLA_TILELANG") is None
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue