studio: skip flash-attn install on Blackwell GPUs (sm_100+) (#5420)
* 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>
This commit is contained in:
parent
000ca89301
commit
79adfd9c71
6 changed files with 284 additions and 2 deletions
|
|
@ -28,6 +28,7 @@ if str(_BACKEND_DIR) not in sys.path:
|
|||
from backend.utils.wheel_utils import (
|
||||
flash_attn_package_version,
|
||||
flash_attn_wheel_url,
|
||||
has_blackwell_gpu,
|
||||
install_wheel,
|
||||
probe_torch_wheel_env,
|
||||
url_exists,
|
||||
|
|
@ -628,10 +629,19 @@ def _flash_attn_install_disabled() -> bool:
|
|||
|
||||
|
||||
def _ensure_flash_attn() -> None:
|
||||
if NO_TORCH or IS_WINDOWS or IS_MACOS:
|
||||
return
|
||||
if _flash_attn_install_disabled():
|
||||
return
|
||||
if NO_TORCH:
|
||||
return
|
||||
if has_blackwell_gpu():
|
||||
_step(
|
||||
"warning",
|
||||
"Skipping flash-attn: Blackwell GPU detected (sm_100+); no compatible prebuilt wheel",
|
||||
_cyan,
|
||||
)
|
||||
return
|
||||
if IS_WINDOWS or IS_MACOS:
|
||||
return
|
||||
if (
|
||||
subprocess.run(
|
||||
[sys.executable, "-c", "import flash_attn"],
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue