Prefer installing mamba-ssm from wheel to speed up things
This commit is contained in:
parent
1d0328eeee
commit
d806da70e2
2 changed files with 152 additions and 88 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue