Prefer installing mamba-ssm from wheel to speed up things

This commit is contained in:
Datta Nimmaturi 2026-03-24 12:06:05 +00:00
commit d806da70e2
2 changed files with 152 additions and 88 deletions

View file

@ -20,6 +20,7 @@ import sys
import time
import traceback
import json
import subprocess as _sp
from pathlib import Path
from typing import Any
import urllib.error
@ -30,6 +31,8 @@ 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:
@ -60,8 +63,6 @@ def _causal_conv1d_platform_tag() -> str | None:
def _probe_causal_conv1d_env() -> dict[str, str] | None:
import subprocess as _sp
probe = _sp.run(
[
sys.executable,
@ -86,26 +87,30 @@ def _probe_causal_conv1d_env() -> dict[str, str] | None:
try:
return json.loads(probe.stdout.strip())
except Exception:
except json.JSONDecodeError:
logger.warning("Failed to parse torch environment probe output: %s", probe.stdout)
return None
def _causal_conv1d_direct_wheel_url() -> str | None:
env = _probe_causal_conv1d_env()
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"causal_conv1d-{_CAUSAL_CONV1D_PACKAGE_VERSION}"
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 (
"https://github.com/Dao-AILab/causal-conv1d/releases/download/"
f"{_CAUSAL_CONV1D_RELEASE_TAG}/{filename}"
)
return f"{release_base_url}/{release_tag}/{filename}"
def _url_exists(url: str) -> bool:
@ -123,69 +128,136 @@ def _url_exists(url: str) -> bool:
return False
def _ensure_causal_conv1d_fast_path(event_queue: Any, model_name: str) -> None:
if not _model_wants_causal_conv1d(model_name):
return
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 causal_conv1d # noqa: F401
logger.info("causal-conv1d already installed")
__import__(import_name)
logger.info("%s already installed", display_name)
return
except Exception:
except ImportError:
pass
wheel_url = _causal_conv1d_direct_wheel_url()
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 causal-conv1d wheel candidate for this environment")
return
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...")
uv_cmd = [
"uv",
"pip",
"install",
"--python",
sys.executable,
"--torch-backend=auto",
"--no-deps",
wheel_url,
]
result = _sp.run(
uv_cmd,
stdout = _sp.PIPE,
stderr = _sp.STDOUT,
text = True,
)
if result.returncode != 0:
logger.warning("uv failed to install %s wheel:\n%s", display_name, result.stdout)
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:
logger.info("Installed prebuilt %s wheel", display_name)
return
logger.warning("pip failed to install %s wheel:\n%s", display_name, result.stdout)
else:
logger.info("Installed prebuilt %s wheel", display_name)
return
else:
logger.info("No published %s wheel found: %s", display_name, wheel_url)
if not _url_exists(wheel_url):
logger.info("No published causal-conv1d wheel found for this environment: %s", wheel_url)
return
_send_status(event_queue, "Installing prebuilt causal-conv1d wheel...")
logger.info("Installing causal-conv1d from direct wheel URL: %s", wheel_url)
import subprocess as _sp
uv_cmd = [
"uv",
_send_status(event_queue, f"Installing {display_name} from PyPI...")
pypi_cmd = [
sys.executable,
"-m",
"pip",
"install",
"--python",
sys.executable,
"--torch-backend=auto",
"--no-build-isolation",
"--no-deps",
wheel_url,
"--no-cache-dir",
f"{pypi_name}=={pypi_version}",
]
result = _sp.run(
uv_cmd,
pypi_cmd,
stdout = _sp.PIPE,
stderr = _sp.STDOUT,
text = True,
)
if result.returncode != 0:
logger.warning("uv failed to install causal-conv1d wheel, falling back to pip:\n%s", result.stdout)
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:
logger.error("Failed to install causal-conv1d wheel:\n%s", result.stdout)
return
logger.error("Failed to install %s from PyPI:\n%s", display_name, result.stdout)
return
logger.info("Installed prebuilt causal-conv1d wheel successfully")
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",
)
def _ensure_mamba_ssm(event_queue: Any, model_name: str) -> None:
_SSM_MODEL_SUBSTRINGS = ("nemotron_h", "nemotron-3-nano", "falcon_h1", "falcon-h1")
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:
@ -284,28 +356,9 @@ def run_training_process(
model_name,
)
# ── 1b. Opportunistically install a matching prebuilt causal-conv1d wheel ──
# ── 1b. Set up causal-conv1d first, then install mamba-ssm if needed ──
_ensure_causal_conv1d_fast_path(event_queue, model_name)
# ── 1b. Do not block startup on PyPI source builds for optional SSM kernels ──
_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, but skipping blocking PyPI install for mamba-ssm. "
"causal-conv1d prebuilt wheel installation was already attempted above. "
"The model will rely on Transformers lazy kernel loading when available, "
"or fall back to the native slow path."
)
_send_status(
event_queue,
"Optional mamba-ssm kernels are not preinstalled; continuing without blocking install.",
)
logger.info("Continuing without eager mamba-ssm installation")
_ensure_mamba_ssm(event_queue, model_name)
# ── 1c. Set fork start method so dataset.map() can multiprocess ──
# The parent launched us via spawn (clean process), but the compiled

View file

@ -291,7 +291,7 @@ def patch_package_file(package_name: str, relative_path: str, url: str) -> None:
def install_python_stack() -> int:
global USE_UV, _STEP, _TOTAL
_STEP = 0
_TOTAL = 10 if IS_WINDOWS else 11
_TOTAL = 11 if IS_WINDOWS else 12
# 1. Upgrade pip (needed even with uv as fallback and for bootstrapping)
_progress("pip upgrade")
@ -300,7 +300,18 @@ def install_python_stack() -> int:
# Try to use uv for faster installs
USE_UV = _bootstrap_uv()
# 2. Core packages: unsloth-zoo + unsloth
# 2. Preinstall torch to a wheel-friendly line for causal-conv1d happy-path testing
_progress("torch pin")
pip_install(
"Installing pinned torch stack",
"--no-cache-dir",
"torch<2.10.0",
"torchvision",
"torchaudio",
constrain = False,
)
# 3. Core packages: unsloth-zoo + unsloth
_progress("base packages")
pip_install(
"Installing base packages",
@ -308,7 +319,7 @@ def install_python_stack() -> int:
req = REQ_ROOT / "base.txt",
)
# 3. Extra dependencies
# 4. Extra dependencies
_progress("unsloth extras")
pip_install(
"Installing additional unsloth dependencies",
@ -316,7 +327,7 @@ def install_python_stack() -> int:
req = REQ_ROOT / "extras.txt",
)
# 3b. Extra dependencies (no-deps) — audio model support etc.
# 4b. Extra dependencies (no-deps) — audio model support etc.
_progress("extra codecs")
pip_install(
"Installing extras (no-deps)",
@ -325,7 +336,7 @@ def install_python_stack() -> int:
req = REQ_ROOT / "extras-no-deps.txt",
)
# 4. Overrides (torchao, transformers) — force-reinstall
# 5. Overrides (torchao, transformers) — force-reinstall
_progress("dependency overrides")
pip_install(
"Installing dependency overrides",
@ -334,7 +345,7 @@ def install_python_stack() -> int:
req = REQ_ROOT / "overrides.txt",
)
# 5. Triton kernels (no-deps, from source)
# 6. Triton kernels (no-deps, from source)
if not IS_WINDOWS:
_progress("triton kernels")
pip_install(
@ -366,7 +377,7 @@ def install_python_stack() -> int:
# "https://raw.githubusercontent.com/unslothai/unsloth/refs/heads/main/unsloth/save.py",
# )
# 8. Studio dependencies
# 9. Studio dependencies
_progress("studio deps")
pip_install(
"Installing studio dependencies",
@ -374,7 +385,7 @@ def install_python_stack() -> int:
req = REQ_ROOT / "studio.txt",
)
# 9. Data-designer dependencies
# 10. Data-designer dependencies
_progress("data designer deps")
pip_install(
"Installing data-designer base dependencies",
@ -382,7 +393,7 @@ def install_python_stack() -> int:
req = SINGLE_ENV / "data-designer-deps.txt",
)
# 10. Data-designer packages (no-deps to avoid conflicts)
# 11. Data-designer packages (no-deps to avoid conflicts)
_progress("data designer")
pip_install(
"Installing data-designer",
@ -391,7 +402,7 @@ def install_python_stack() -> int:
req = SINGLE_ENV / "data-designer.txt",
)
# 11. Local Data Designer seed plugin
# 12. Local Data Designer seed plugin
if not LOCAL_DD_UNSTRUCTURED_PLUGIN.is_dir():
print(
_red(
@ -408,7 +419,7 @@ def install_python_stack() -> int:
constrain = False,
)
# 12. Patch metadata for single-env compatibility
# 13. Patch metadata for single-env compatibility
_progress("finalizing")
run(
"Patching single-env metadata",