From 04359be333529298f6fd99ba3ac2cf57c79e8cfe Mon Sep 17 00:00:00 2001 From: Datta Nimmaturi Date: Wed, 25 Mar 2026 14:52:26 +0530 Subject: [PATCH] [Studio] Try installing causal-conv1d from prebuilt wheels if avialable (#4547) * Try installing causal-conv1d from prebuilt wheels if avialable * Prefer installing mamba-ssm from wheel to speed up things * undo python stack install changes * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci * Revert "undo python stack install changes" This reverts commit d943551092ea080355acdb70438c3a5d0083ea7f. * add comments * Fix wheel installer: model detection, platform tags, torch pin, error handling - Add nemotron-h (hyphen) and granite-4.0-h / granitemoehybrid to model detection for both causal-conv1d and mamba-ssm. These hybrid Mamba models were silently skipped since nemotron_h (underscore) never matches real HF model IDs like nvidia/Nemotron-H-8B-Base, and granite was missing entirely despite being a supported model in model_config.py and loader.py. - Fix _causal_conv1d_platform_tag to detect linux_aarch64 via platform.machine() instead of hardcoding linux_x86_64. Both upstream releases publish aarch64 wheels. Drop win_amd64 since neither repo publishes Windows wheels (avoids a wasted HTTP probe on every run). - Pin torch to >=2.6.0,<2.11.0 instead of <=2.10.0 to add a version floor and document the wheel coverage range with upstream release links. - Strip non-numeric suffixes from torch minor version so nightly builds like 2.7a0 correctly resolve to wheel tag torch2.7 instead of torch2.7a0. - Use stderr=_sp.PIPE instead of stderr=_sp.STDOUT in the env probe so torch import warnings do not corrupt the JSON output. - Add timeout=30 to the env probe subprocess to prevent indefinite hangs. - Catch Exception (not just ImportError) on the existing-install check so ABI-broken installs with OSError/RuntimeError are retried rather than silently accepted. - Guard uv invocation with shutil.which("uv") to prevent FileNotFoundError crash when uv is not on PATH. Wrap the top-level ensure calls in try/except so failures do not kill the training worker. - Hoist _SSM_MODEL_SUBSTRINGS to module level. - Remove redundant --torch-backend=auto flag from direct wheel URL install. * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci * Add LFM2 to causal-conv1d detection; stop training on install failure - Add "lfm2" to _model_wants_causal_conv1d so Studio picks up the fast kernel path for Liquid Foundation Model 2. - Replace silent logger.warning on SSM dependency install failure with an error event that tells the user to choose another model and stops the training job immediately. * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci * Catch subprocess timeout in torch probe; narrow import guard to ImportError - _probe_causal_conv1d_env: wrap subprocess.run in try/except for TimeoutExpired so a slow torch import returns None (falls back to PyPI) instead of killing the training job. - _install_package_wheel_first: narrow except Exception to except ImportError on the __import__ check so unexpected errors from a broken module still propagate. * Remove unconditional torch pin from install_python_stack The torch>=2.6.0,<2.11.0 pin was added to ensure prebuilt causal-conv1d / mamba-ssm wheels exist, but it runs at install time for all users regardless of model choice. This can downgrade or unnecessarily upgrade torch. The worker already handles wheel compatibility at training time by probing the environment and falling back to PyPI, so the install-time pin is not needed. --------- Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com> Co-authored-by: Daniel Han --- studio/backend/core/training/worker.py | 336 ++++++++++++++++++++++--- 1 file changed, 297 insertions(+), 39 deletions(-) diff --git a/studio/backend/core/training/worker.py b/studio/backend/core/training/worker.py index d06dd6d358..891dfca8f7 100644 --- a/studio/backend/core/training/worker.py +++ b/studio/backend/core/training/worker.py @@ -16,15 +16,294 @@ from __future__ import annotations import structlog from loggers import get_logger import os +import platform +import shutil import sys import time import traceback +import json +import subprocess as _sp from pathlib import Path from typing import Any +import urllib.error +import urllib.request logger = get_logger(__name__) +_CAUSAL_CONV1D_RELEASE_TAG = "v1.6.1.post4" +_CAUSAL_CONV1D_PACKAGE_VERSION = "1.6.1" +_MAMBA_SSM_RELEASE_TAG = "v2.3.1" +_MAMBA_SSM_PACKAGE_VERSION = "2.3.1" + + +def _model_wants_causal_conv1d(model_name: str) -> bool: + name = model_name.lower() + return any( + key in name + for key in ( + "qwen3.5", + "qwen3_5", + "qwen3-next", + "qwen3_next", + "nemotron_h", + "nemotron-h", + "nemotron-3-nano", + "falcon_h1", + "falcon-h1", + "granite-4.0-h", + "granitemoehybrid", + "lfm2", + ) + ) + + +def _causal_conv1d_platform_tag() -> str | None: + machine = platform.machine().lower() + if sys.platform.startswith("linux"): + if machine in {"x86_64", "amd64"}: + return "linux_x86_64" + if machine in {"aarch64", "arm64"}: + return "linux_aarch64" + return None + # No prebuilt wheels published for macOS or Windows + return None + + +def _probe_causal_conv1d_env() -> dict[str, str] | None: + try: + probe = _sp.run( + [ + sys.executable, + "-c", + ( + "import json, sys, re, torch; " + "parts = torch.__version__.split('+', 1)[0].split('.')[:2]; " + "minor = re.sub(r'[^0-9].*', '', parts[1]) if len(parts) > 1 else '0'; " + "torch_mm = parts[0] + '.' + minor; " + "print(json.dumps({" + "'python_tag': f'cp{sys.version_info.major}{sys.version_info.minor}', " + "'torch_mm': torch_mm, " + "'cuda_major': str(int(str(torch.version.cuda).split('.', 1)[0])) if torch.version.cuda else '', " + "'cxx11abi': str(torch._C._GLIBCXX_USE_CXX11_ABI).upper()" + "}))" + ), + ], + stdout = _sp.PIPE, + stderr = _sp.PIPE, + text = True, + timeout = 30, + ) + except _sp.TimeoutExpired: + logger.warning("Torch environment probe timed out after 30s") + return None + if probe.returncode != 0: + logger.warning( + "Failed to probe torch environment for causal-conv1d wheel:\n%s", + probe.stdout, + ) + return None + + try: + return json.loads(probe.stdout.strip()) + except json.JSONDecodeError: + logger.warning( + "Failed to parse torch environment probe output: %s", probe.stdout + ) + return None + + +def _direct_wheel_url( + *, + filename_prefix: str, + package_version: str, + release_tag: str, + release_base_url: str, + env: dict[str, str] | None = None, +) -> str | None: + env = env or _probe_causal_conv1d_env() + platform_tag = _causal_conv1d_platform_tag() + if env is None or platform_tag is None or not env.get("cuda_major"): + return None + + filename = ( + f"{filename_prefix}-{package_version}" + f"+cu{env['cuda_major']}torch{env['torch_mm']}" + f"cxx11abi{env['cxx11abi']}-{env['python_tag']}-{env['python_tag']}-{platform_tag}.whl" + ) + return f"{release_base_url}/{release_tag}/{filename}" + + +def _url_exists(url: str) -> bool: + try: + request = urllib.request.Request(url, method = "HEAD") + with urllib.request.urlopen(request, timeout = 10): + return True + except urllib.error.HTTPError as exc: + if exc.code == 404: + return False + logger.warning("Unexpected HTTP error while probing %s: %s", url, exc) + return False + except Exception as exc: + logger.warning("Failed to probe %s: %s", url, exc) + return False + + +def _install_package_wheel_first( + *, + event_queue: Any, + import_name: str, + display_name: str, + pypi_name: str, + pypi_version: str, + filename_prefix: str, + release_tag: str, + release_base_url: str, +) -> None: + try: + __import__(import_name) + logger.info("%s already installed", display_name) + return + except ImportError: + pass + + env = _probe_causal_conv1d_env() + wheel_url = _direct_wheel_url( + filename_prefix = filename_prefix, + package_version = pypi_version, + release_tag = release_tag, + release_base_url = release_base_url, + env = env, + ) + + if wheel_url is None: + logger.info("No compatible %s wheel candidate", display_name) + else: + if _url_exists(wheel_url): + _send_status(event_queue, f"Installing prebuilt {display_name} wheel...") + installed = False + # Try uv first if available, then fall back to pip + if shutil.which("uv"): + uv_cmd = [ + "uv", + "pip", + "install", + "--python", + sys.executable, + "--no-deps", + wheel_url, + ] + result = _sp.run( + uv_cmd, + stdout = _sp.PIPE, + stderr = _sp.STDOUT, + text = True, + ) + if result.returncode == 0: + installed = True + else: + logger.warning( + "uv failed to install %s wheel:\n%s", + display_name, + result.stdout, + ) + if not installed: + pip_cmd = [ + sys.executable, + "-m", + "pip", + "install", + "--no-deps", + wheel_url, + ] + result = _sp.run( + pip_cmd, + stdout = _sp.PIPE, + stderr = _sp.STDOUT, + text = True, + ) + if result.returncode == 0: + installed = True + else: + logger.warning( + "pip failed to install %s wheel:\n%s", + display_name, + result.stdout, + ) + if installed: + logger.info("Installed prebuilt %s wheel successfully", display_name) + return + else: + logger.info("No published %s wheel found: %s", display_name, wheel_url) + + _send_status(event_queue, f"Installing {display_name} from PyPI...") + pypi_cmd = [ + sys.executable, + "-m", + "pip", + "install", + "--no-build-isolation", + "--no-deps", + "--no-cache-dir", + f"{pypi_name}=={pypi_version}", + ] + result = _sp.run( + pypi_cmd, + stdout = _sp.PIPE, + stderr = _sp.STDOUT, + text = True, + ) + if result.returncode != 0: + logger.error("Failed to install %s from PyPI:\n%s", display_name, result.stdout) + return + + logger.info("Installed %s from PyPI", display_name) + + +def _ensure_causal_conv1d_fast_path(event_queue: Any, model_name: str) -> None: + if not _model_wants_causal_conv1d(model_name): + return + + _install_package_wheel_first( + event_queue = event_queue, + import_name = "causal_conv1d", + display_name = "causal-conv1d", + pypi_name = "causal-conv1d", + pypi_version = _CAUSAL_CONV1D_PACKAGE_VERSION, + filename_prefix = "causal_conv1d", + release_tag = _CAUSAL_CONV1D_RELEASE_TAG, + release_base_url = "https://github.com/Dao-AILab/causal-conv1d/releases/download", + ) + + +_SSM_MODEL_SUBSTRINGS = ( + "nemotron_h", + "nemotron-h", + "nemotron-3-nano", + "falcon_h1", + "falcon-h1", + "granite-4.0-h", + "granitemoehybrid", +) + + +def _ensure_mamba_ssm(event_queue: Any, model_name: str) -> None: + if not any(sub in model_name.lower() for sub in _SSM_MODEL_SUBSTRINGS): + return + + logger.info("SSM model detected; setting up mamba-ssm after causal-conv1d") + _install_package_wheel_first( + event_queue = event_queue, + import_name = "mamba_ssm", + display_name = "mamba-ssm", + pypi_name = "mamba-ssm", + pypi_version = _MAMBA_SSM_PACKAGE_VERSION, + filename_prefix = "mamba_ssm", + release_tag = _MAMBA_SSM_RELEASE_TAG, + release_base_url = "https://github.com/state-spaces/mamba/releases/download", + ) + + def _activate_transformers_version(model_name: str) -> None: """Activate the correct transformers version BEFORE any ML imports. @@ -121,45 +400,24 @@ def run_training_process( model_name, ) - # ── 1b. Auto-install mamba-ssm for SSM/hybrid models (NemotronH, Falcon-H1) ── - _SSM_MODEL_SUBSTRINGS = ("nemotron_h", "nemotron-3-nano", "falcon_h1", "falcon-h1") - if any(sub in model_name.lower() for sub in _SSM_MODEL_SUBSTRINGS): - try: - import mamba_ssm # noqa: F401 - - logger.info("mamba-ssm already installed") - except ImportError: - logger.info( - "SSM model detected — installing mamba-ssm and causal-conv1d (this may take several minutes)..." - ) - _send_status( - event_queue, "Installing mamba-ssm (first time only, ~7 min)..." - ) - import subprocess as _sp - - # --no-build-isolation: compile against current torch (no version conflicts) - # --no-deps: don't pull in torch/transformers/triton (already installed) - for _pkg in ["causal_conv1d", "mamba_ssm"]: - _r = _sp.run( - [ - sys.executable, - "-m", - "pip", - "install", - "--no-build-isolation", - "--no-deps", - "--no-cache-dir", - _pkg, - ], - stdout = _sp.PIPE, - stderr = _sp.STDOUT, - text = True, - ) - if _r.returncode != 0: - logger.error("Failed to install %s:\n%s", _pkg, _r.stdout) - else: - logger.info("Installed %s successfully", _pkg) - logger.info("mamba-ssm installation complete") + # ── 1b. Set up causal-conv1d first, then install mamba-ssm if needed ── + try: + _ensure_causal_conv1d_fast_path(event_queue, model_name) + _ensure_mamba_ssm(event_queue, model_name) + except Exception as exc: + event_queue.put( + { + "type": "error", + "error": ( + f"Please choose another model to train, since " + f"causal-conv1d / mamba-ssm failed to install " + f"with error: {exc}" + ), + "stack": traceback.format_exc(limit = 20), + "ts": time.time(), + } + ) + return # ── 1c. Set fork start method so dataset.map() can multiprocess ── # The parent launched us via spawn (clean process), but the compiled