import argparse import datetime import hashlib import json import os import shlex import sys import warnings from pathlib import Path from _provenance import ( atomic_write, find_package_repo, git_diff, pkg_version, ) from _provenance import ( git_sha as _git_sha, ) # ``hs_connectors`` is a workspace package, while the dedicated vLLM virtual # environment only installs the main ``speculators`` package. Add it explicitly # for both this launcher and the child ``vllm serve`` process so out-of-tree KV # connectors are importable without modifying the vLLM environment. _HS_CONNECTORS_SRC = Path(__file__).resolve().parents[1] / "hs_connectors" / "src" if str(_HS_CONNECTORS_SRC) not in sys.path: sys.path.insert(0, str(_HS_CONNECTORS_SRC)) _existing_pythonpath = os.environ.get("PYTHONPATH") os.environ["PYTHONPATH"] = ( f"{_HS_CONNECTORS_SRC}{os.pathsep}{_existing_pythonpath}" if _existing_pythonpath else str(_HS_CONNECTORS_SRC) ) try: from hs_connectors import HiddenStatesBackend _backend_registry: dict[str, type[HiddenStatesBackend]] = dict( HiddenStatesBackend.registry # type: ignore[misc] ) except ImportError: _backend_registry = {} # type: ignore[assignment] if "file" not in _backend_registry: # Vendored File backend in case hs_connectors is not available class _InlineFileBackend: @staticmethod def add_launch_args(parser: argparse.ArgumentParser) -> None: parser.add_argument( "--hidden-states-path", type=str, default="/tmp/hidden_states", # noqa: S108 help=( "The directory to save hidden states to. " "Default '/tmp/hidden_states'" ), ) @staticmethod def build_kv_transfer_config(args: argparse.Namespace) -> dict: return { "kv_connector": "ExampleHiddenStatesConnector", "kv_role": "kv_producer", "kv_connector_extra_config": { "shared_storage_path": args.hidden_states_path, }, } _backend_registry["file"] = _InlineFileBackend # type: ignore[assignment] # Keep the preprocessing workers and vLLM front end within one CPU budget. # Four preprocessing workers are paired with one API server. The former is # estimated at 3 CPUs and the latter at 4 CPUs, so each preprocessing worker # represents 4 CPUs of combined capacity. Leave 25% for native runtime # threads and other application work. CPU_BUDGET_FRACTION = 0.75 DEFAULT_RENDERER_NUM_WORKERS = 2 MAX_API_SERVER_COUNT = 32 WORKERS_PER_API_SERVER = 4 CPUS_PER_API_SERVER = 4 MAX_PREPROCESSING_WORKERS = 128 EFFECTIVE_CPUS_PER_PREPROCESSING_WORKER = 4 # The vLLM API-server processes each create their own native thread pools. # Keep those pools bounded by default; explicit environment settings still # take precedence for users who have measured a different configuration. DEFAULT_RENDER_THREAD_ENV = { "OMP_NUM_THREADS": "1", "OPENBLAS_NUM_THREADS": "1", "MKL_NUM_THREADS": "1", "RAYON_NUM_THREADS": "2", } def _usable_cpu_count() -> int: """Return the CPUs available to this process, respecting affinity.""" if hasattr(os, "process_cpu_count"): # Python 3.13+ return os.process_cpu_count() or 1 if hasattr(os, "sched_getaffinity"): # Linux return len(os.sched_getaffinity(0)) return os.cpu_count() or 1 def _preprocessing_workers(cpus: int) -> int: """Mirror prepare_data.py's shared render CPU budget.""" return max( 1, min( MAX_PREPROCESSING_WORKERS, int(cpus * CPU_BUDGET_FRACTION) // EFFECTIVE_CPUS_PER_PREPROCESSING_WORKER, ), ) def render_throughput_defaults(cpus: int | None = None) -> tuple[int, int]: """Return affinity-aware API-server and renderer-worker defaults.""" if cpus is None: cpus = _usable_cpu_count() api_servers = max( 1, min( MAX_API_SERVER_COUNT, _preprocessing_workers(cpus) // WORKERS_PER_API_SERVER, cpus // CPUS_PER_API_SERVER, ), ) return api_servers, DEFAULT_RENDERER_NUM_WORKERS def _with_render_defaults(vllm_args: list[str]) -> list[str]: """Prepend high-throughput render defaults, unless no API server is wanted.""" if "--headless" in vllm_args: return vllm_args api_servers, renderer_workers = render_throughput_defaults() return [ "--api-server-count", str(api_servers), "--renderer-num-workers", str(renderer_workers), *vllm_args, ] def _set_render_thread_defaults() -> None: """Bound native pools inherited by vLLM's API-server processes.""" for name, value in DEFAULT_RENDER_THREAD_ENV.items(): os.environ.setdefault(name, value) def _enable_scale_out_endpoints() -> None: """Enable vLLM's render route while preserving an explicit user setting.""" os.environ.setdefault("VLLM_ENABLE_SCALE_OUT_ENDPOINTS", "1") def _add_shared_args(parser: argparse.ArgumentParser) -> None: parser.add_argument( "--provenance-dir", type=str, default=None, help=( "Directory to write vllm_command.txt, vllm.patch, and " "checkpoint_sha256.txt (plus drafter_checkpoint_sha256.txt in " "eval mode). Provenance is only captured when this is set; omit " "it to skip logging (e.g. for ad-hoc test/debug runs)." ), ) parser.add_argument( "--no-hash-checkpoints", action="store_true", default=False, help=( "Skip SHA256 hashing of .safetensors files. Useful " "for large checkpoints where hashing adds significant " "launch latency. When set, checkpoint_sha256.txt " "records file sizes and modification times instead." ), ) parser.add_argument( "--dry-run", action="store_true", help="Print the command without running it", ) def parse_args(): parser = argparse.ArgumentParser( description="Launch vLLM for training or evaluation", ) sub = parser.add_subparsers(dest="subcommand") # --- train subcommand (default when no subcommand given) --- train_parser = sub.add_parser( "train", help="Hidden-states extraction for training data generation", ) train_parser.add_argument( "model", type=str, help="Model name or path to extract hidden states from", ) train_parser.add_argument( "--hidden-states-backend", choices=list(_backend_registry.keys()), default="file", help=( "Hidden states transfer backend. Each backend may " "add its own CLI arguments (see below). " "Default: 'file'." ), ) for backend_cls in _backend_registry.values(): backend_cls.add_launch_args(train_parser) train_parser.add_argument( "--target-layer-ids", type=int, nargs="+", help=( "(Optional) Space-separated list of integer layer " "ids. Defaults to " "[2, num_hidden_layers // 2, num_hidden_layers - 3]." " Note: if set, you must also pass the same value " "into the training process" ), ) train_parser.add_argument( "--include-last-layer", action=argparse.BooleanOptionalAction, default=True, help=( "Append the last layer (num_hidden_layers) to " "target_layer_ids for verifier hidden states " "extraction. Default: True" ), ) train_parser.add_argument( "--trust-remote-code", action="store_true", help=( "Allow custom model configuration code while resolving " "hidden-state layer ids. Pass the same flag after '--' for " "vLLM itself." ), ) _add_shared_args(train_parser) # --- eval subcommand --- eval_parser = sub.add_parser( "eval", help="Speculative decoding serving for evaluation", ) eval_parser.add_argument( "model", type=str, help="Target model name or path", ) eval_parser.add_argument( "--spec-model", type=str, required=True, help="Drafter model name or path", ) eval_parser.add_argument( "--spec-tokens", type=int, default=None, help="Number of speculative tokens", ) eval_parser.add_argument( "--spec-method", type=str, default=None, help="Speculative decoding method (optional)", ) _add_shared_args(eval_parser) subcommands = set(sub.choices) argv = sys.argv[1:] if not argv or (argv[0] not in subcommands and not argv[0].startswith("-")): argv = ["train", *argv] args, vllm_args = parser.parse_known_args(argv) if args.subcommand is None: args, vllm_args = parser.parse_known_args(["train"] + sys.argv[1:]) return args, vllm_args def _warn(msg: str) -> None: print(f"Warning: {msg}", file=sys.stderr) def _find_vllm_repo() -> str | None: """Find the vllm git checkout by walking up from the installed package.""" repo = find_package_repo("vllm") if repo and (repo / "vllm" / "__init__.py").is_file(): return str(repo) return None def _sha256_file(path: str) -> str: h = hashlib.sha256() with open(path, "rb") as f: while chunk := f.read(1 << 20): h.update(chunk) return h.hexdigest() # --------------------------------------------------------------------------- # Individual provenance writers — each is self-contained and best-effort. # --------------------------------------------------------------------------- def _save_vllm_command( prov_dir: Path, cmd: list[str], sha: str, diff: str, vllm_ver: str, ) -> None: sha_label = f"{sha} (dirty)" if diff else sha ts = datetime.datetime.now(tz=datetime.timezone.utc).isoformat() header = "\n".join( [ f"# Timestamp: {ts}", f"# Python: {sys.executable}", f"# Git SHA: {sha_label}", f"# vllm: {vllm_ver}", ] ) atomic_write( prov_dir / "vllm_command.txt", f"{header}\n{shlex.join(cmd)}\n", ) def _save_vllm_patch( prov_dir: Path, vllm_repo: str | None, sha: str, diff: str, vllm_ver: str, ) -> None: if vllm_repo: content = f"# repo: {vllm_repo} ({sha})\n{diff}\n" else: content = f"# vllm {vllm_ver} (wheel install, no git repo found)\n" atomic_write(prov_dir / "vllm.patch", content) def _save_checkpoint_sha256( prov_dir: Path, model: str, *, skip_hash: bool = False, dest_name: str = "checkpoint_sha256.txt", ) -> None: dest = prov_dir / dest_name model_path = os.path.expanduser(model) if not os.path.isdir(model_path): atomic_write(dest, f"# model: {model} (not a local path)\n") return safetensors = sorted( f for f in os.listdir(model_path) if f.endswith(".safetensors") ) if not safetensors: atomic_write(dest, f"# no .safetensors files in {model_path}\n") return if skip_hash: lines = [] for name in safetensors: fp = os.path.join(model_path, name) st = os.stat(fp) lines.append(f"size={st.st_size} mtime={st.st_mtime} {name}") header = "# hashing skipped (--no-hash-checkpoints)\n" atomic_write(dest, header + "\n".join(lines) + "\n") else: lines = [ f"{_sha256_file(os.path.join(model_path, name))} {name}" for name in safetensors ] atomic_write(dest, "\n".join(lines) + "\n") # --------------------------------------------------------------------------- # Top-level entry point # --------------------------------------------------------------------------- def _save_vllm_provenance( cmd: list[str], provenance_dir: str, model: str, *, skip_hash: bool = False, spec_model: str | None = None, ) -> None: """Write vllm_command.txt, vllm.patch, and checkpoint_sha256.txt. In eval mode (``spec_model`` given) also writes drafter_checkpoint_sha256.txt. Best-effort — failures warn but never block the vLLM launch. """ prov_dir = Path(provenance_dir) try: prov_dir.mkdir(parents=True, exist_ok=True) except OSError as exc: _warn(f"could not create provenance dir: {exc}") return vllm_repo = _find_vllm_repo() vllm_root = Path(vllm_repo) if vllm_repo else None sha = _git_sha(vllm_root) diff = git_diff(vllm_root) vllm_ver = pkg_version("vllm") writers = [ ( "vllm_command.txt", lambda: _save_vllm_command(prov_dir, cmd, sha, diff, vllm_ver), ), ( "vllm.patch", lambda: _save_vllm_patch(prov_dir, vllm_repo, sha, diff, vllm_ver), ), ( "checkpoint_sha256.txt", lambda: _save_checkpoint_sha256(prov_dir, model, skip_hash=skip_hash), ), ] if spec_model is not None: drafter_dest = "drafter_checkpoint_sha256.txt" writers.append( ( drafter_dest, lambda: _save_checkpoint_sha256( prov_dir, spec_model, skip_hash=skip_hash, dest_name=drafter_dest, ), ) ) for artifact, write in writers: try: write() except Exception as exc: # noqa: BLE001 _warn(f"could not save {artifact}: {exc}") def _build_train_cmd(args, vllm_args): from transformers import AutoConfig # noqa: PLC0415 config = AutoConfig.from_pretrained( args.model, trust_remote_code=args.trust_remote_code, ) if hasattr(config, "text_config"): config = config.text_config num_hidden_layers = config.num_hidden_layers if args.target_layer_ids: target_layer_ids = args.target_layer_ids if args.include_last_layer and num_hidden_layers not in target_layer_ids: target_layer_ids.append(num_hidden_layers) warnings.warn( f"Using custom target layer ids {target_layer_ids}. These " "must also be explicitly passed into the training script.", stacklevel=2, ) else: target_layer_ids = [ 2, num_hidden_layers // 2, num_hidden_layers - 3, num_hidden_layers, ] # Layer id ``num_hidden_layers`` (the final hidden state) is valid: the # default above and --include-last-layer both emit it. if ( min(target_layer_ids) < 0 or max(target_layer_ids) > num_hidden_layers or len(set(target_layer_ids)) != len(target_layer_ids) ): raise ValueError( f"Invalid target layer ids {target_layer_ids}; ids must be " f"distinct and within [0, {num_hidden_layers}]." ) speculative_config = { "method": "extract_hidden_states", "num_speculative_tokens": 1, "draft_model_config": { "hf_config": {"eagle_aux_hidden_state_layer_ids": target_layer_ids} }, } backend_cls = _backend_registry[args.hidden_states_backend] kv_transfer_config = backend_cls.build_kv_transfer_config(args) return [ sys.executable, "-m", "vllm.entrypoints.cli.main", "serve", args.model, "--speculative_config", json.dumps(speculative_config), "--kv_transfer_config", json.dumps(kv_transfer_config), *_with_render_defaults(vllm_args), ] def _build_eval_cmd(args, vllm_args): cmd = [ sys.executable, "-m", "vllm.entrypoints.cli.main", "serve", args.model, "--spec-model", args.spec_model, ] if args.spec_tokens is not None: cmd.extend(["--spec-tokens", str(args.spec_tokens)]) if args.spec_method is not None: cmd.extend(["--spec-method", args.spec_method]) cmd.extend(vllm_args) return cmd def main(): args, vllm_args = parse_args() if "--" in vllm_args: vllm_args.remove("--") if args.subcommand == "train": cmd = _build_train_cmd(args, vllm_args) elif args.subcommand == "eval": cmd = _build_eval_cmd(args, vllm_args) else: raise ValueError(f"Unknown subcommand: {args.subcommand}") print("Running command:") print(" ".join(cmd)) if args.provenance_dir: _save_vllm_provenance( cmd, args.provenance_dir, args.model, skip_hash=args.no_hash_checkpoints, spec_model=getattr(args, "spec_model", None), ) if not args.dry_run: # Render tuning applies to the train pipeline only; eval serving skips it. if args.subcommand == "train" and "--headless" not in vllm_args: _set_render_thread_defaults() _enable_scale_out_endpoints() os.execvp(cmd[0], cmd) # noqa: S606 if __name__ == "__main__": main()