* studio: skip flash-attn install on Blackwell GPUs (sm_100+) Dao-AILab does not publish prebuilt flash-attn wheels for sm_100, sm_120, or sm_121, and the older-arch wheels fail to load on Blackwell. Add a shared has_blackwell_gpu() helper and gate both the install-time (install_python_stack._ensure_flash_attn) and runtime (worker._ensure_flash_attn_for_long_context) paths on it. Detection uses nvidia-smi --query-gpu=compute_cap, which works on Linux and Windows. * test: stub has_blackwell_gpu in pre-existing runtime flash-attn tests prefers_prebuilt_wheel and falls_back_to_pypi exercise the install paths that the Blackwell guard now short-circuits. Make them explicit about non-Blackwell so they pass on real Blackwell hosts. * studio: cache has_blackwell_gpu, skip Blackwell warning under NO_TORCH - Wrap has_blackwell_gpu in functools.lru_cache so repeated calls in a single process avoid redundant nvidia-smi spawns. Tests clear the cache via setup_method/teardown_method. - In _ensure_flash_attn, run the NO_TORCH short-circuit before the Blackwell check so GGUF-only users (who never install torch anyway) do not see a Blackwell warning. Blackwell check still runs above the IS_WINDOWS / IS_MACOS gates so Blackwell-on-Windows users still see the explicit reason rather than a silent OS skip. * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci * test: add has_blackwell_gpu to mlx worker test wheel_utils stub test_mlx_training_worker_config loads worker.py against a hand-rolled utils.wheel_utils stub. Adding has_blackwell_gpu to the stub symbol list so worker's import line resolves. --------- Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
84 lines
2.7 KiB
Python
84 lines
2.7 KiB
Python
# SPDX-License-Identifier: AGPL-3.0-only
|
|
|
|
import importlib.util
|
|
import sys
|
|
import types
|
|
from pathlib import Path
|
|
|
|
import pytest
|
|
|
|
|
|
def _load_worker_module():
|
|
stub_names = (
|
|
"structlog",
|
|
"loggers",
|
|
"utils",
|
|
"utils.hardware",
|
|
"utils.wheel_utils",
|
|
)
|
|
previous_modules = {name: sys.modules.get(name) for name in stub_names}
|
|
|
|
try:
|
|
sys.modules["structlog"] = types.ModuleType("structlog")
|
|
|
|
loggers = types.ModuleType("loggers")
|
|
loggers.get_logger = lambda *_args, **_kwargs: None
|
|
sys.modules["loggers"] = loggers
|
|
|
|
utils = types.ModuleType("utils")
|
|
utils.__path__ = []
|
|
sys.modules["utils"] = utils
|
|
|
|
hardware = types.ModuleType("utils.hardware")
|
|
hardware.apply_gpu_ids = lambda *_args, **_kwargs: None
|
|
sys.modules["utils.hardware"] = hardware
|
|
|
|
wheel_utils = types.ModuleType("utils.wheel_utils")
|
|
for name in (
|
|
"direct_wheel_url",
|
|
"flash_attn_wheel_url",
|
|
"has_blackwell_gpu",
|
|
"install_wheel",
|
|
"probe_torch_wheel_env",
|
|
"url_exists",
|
|
):
|
|
setattr(wheel_utils, name, lambda *_args, **_kwargs: None)
|
|
sys.modules["utils.wheel_utils"] = wheel_utils
|
|
|
|
worker_path = (
|
|
Path(__file__).resolve().parents[1] / "core" / "training" / "worker.py"
|
|
)
|
|
spec = importlib.util.spec_from_file_location(
|
|
"mlx_training_worker_under_test", worker_path
|
|
)
|
|
module = importlib.util.module_from_spec(spec)
|
|
assert spec.loader is not None
|
|
spec.loader.exec_module(module)
|
|
return module
|
|
finally:
|
|
for name, module in previous_modules.items():
|
|
if module is None:
|
|
sys.modules.pop(name, None)
|
|
else:
|
|
sys.modules[name] = module
|
|
|
|
|
|
_worker = _load_worker_module()
|
|
_normalize_mlx_studio_optimizer = _worker._normalize_mlx_studio_optimizer
|
|
_normalize_mlx_studio_scheduler = _worker._normalize_mlx_studio_scheduler
|
|
|
|
|
|
def test_mlx_studio_optimizer_aliases_are_explicit():
|
|
assert _normalize_mlx_studio_optimizer("adamw_8bit") == "adamw"
|
|
assert _normalize_mlx_studio_optimizer("paged_adamw_8bit") == "adamw"
|
|
assert _normalize_mlx_studio_optimizer("adafactor") == "adafactor"
|
|
|
|
|
|
def test_mlx_studio_rejects_unknown_optimizer():
|
|
with pytest.raises(ValueError, match = "Unsupported optimizer for MLX training"):
|
|
_normalize_mlx_studio_optimizer("adamw_typo")
|
|
|
|
|
|
def test_mlx_studio_rejects_unknown_scheduler():
|
|
with pytest.raises(ValueError, match = "Unsupported LR scheduler for MLX training"):
|
|
_normalize_mlx_studio_scheduler("linear_typo")
|