dora7's picture
Albedo SN97 workspace v16: RECORD, merged v16, adapters v13/v15/v16-dpo, packs, harness, evals
2abcc30 verified
Raw
History Blame Contribute Delete
2.64 kB
"""Point flashinfer/vLLM at the pip CUDA toolkit and skip JIT that needs nvcc."""
from __future__ import annotations
import os
import sys
from pathlib import Path
def find_cuda_home() -> Path | None:
existing = os.environ.get("CUDA_HOME") or os.environ.get("CUDA_PATH")
if existing:
home = Path(existing)
if (home / "bin" / "nvcc").is_file():
return home
# Do not resolve() sys.executable — venv/bin/python is a symlink to
# /usr/bin/python3.x, which walks out of the venv and misses nvidia/cu*.
prefixes = [
Path(sys.prefix),
Path(getattr(sys, "base_prefix", sys.prefix)),
Path(sys.executable).parent.parent,
]
seen: set[Path] = set()
for prefix in prefixes:
if prefix in seen:
continue
seen.add(prefix)
matches = sorted(prefix.glob("lib/python*/site-packages/nvidia/cu*/bin/nvcc"))
if matches:
return matches[-1].parent.parent
return None
def apply(*, disable_flashinfer_sampler: bool = True, gdn_backend: str = "triton") -> Path | None:
"""Set env for this process and every vLLM worker it later spawns."""
home = find_cuda_home()
if home is not None:
os.environ["CUDA_HOME"] = str(home)
os.environ["CUDA_PATH"] = str(home)
os.environ["PATH"] = f"{home / 'bin'}:{os.environ.get('PATH', '')}"
_link_system_cuda(home)
_ensure_lib64(home)
print(f"CUDA_HOME={home}", flush=True)
if disable_flashinfer_sampler:
os.environ.setdefault("VLLM_USE_FLASHINFER_SAMPLER", "0")
if gdn_backend:
os.environ.setdefault("ALBEDO_GDN_PREFILL_BACKEND", gdn_backend)
# Per-process Triton cache. Shared /root/.triton/cache races on overlayfs;
# /dev/shm is noexec and cannot mmap compiled .so files.
cache = Path("/workspace/data/triton-cache") / f"pid-{os.getpid()}"
cache.mkdir(parents=True, exist_ok=True)
os.environ["TRITON_CACHE_DIR"] = str(cache)
os.environ.setdefault("TRITON_HOME", "/workspace/data/triton-cache/home")
return home
def _link_system_cuda(home: Path) -> None:
target = Path("/usr/local/cuda")
if target.exists() or target.is_symlink():
return
try:
target.symlink_to(home)
print(f"linked {target} -> {home}", flush=True)
except OSError as exc:
print(f"could not link {target}: {exc}", flush=True)
def _ensure_lib64(home: Path) -> None:
lib64 = home / "lib64"
lib = home / "lib"
if lib64.exists() or not lib.is_dir():
return
try:
lib64.symlink_to("lib")
except OSError:
pass