Download source/tests/e2e/utils.py from khazic/spec-b300: direct link, hf CLI and curl.
- Browser
- Download file 19.6 kB
-
https://huggingface.co/khazic/spec-b300/resolve/main/source/tests/e2e/utils.py
- Command line
-
hf download hf://khazic/spec-b300/source/tests/e2e/utils.py
-
curl -L -o utils.py https://huggingface.co/khazic/spec-b300/resolve/main/source/tests/e2e/utils.py
19.6 kB
| import json | |
| import os | |
| import shutil | |
| import signal | |
| import subprocess | |
| import sys | |
| import time | |
| import urllib.error | |
| import urllib.request | |
| from collections.abc import Callable, Iterable | |
| from contextlib import contextmanager, suppress | |
| from functools import wraps | |
| from pathlib import Path | |
| from textwrap import indent | |
| import pytest | |
| from loguru import logger | |
| from PIL import Image | |
| from speculators.data_generation.preprocessing import load_raw_dataset | |
| __all__ = [ | |
| "SCRIPTS_DIR", | |
| "VLLM_PYTHON", | |
| "launch_mooncake_master", | |
| "launch_mooncake_master_context", | |
| "launch_vllm_server", | |
| "launch_vllm_server_context", | |
| "purge_newfiles", | |
| "run_data_generation_offline", | |
| "run_prepare_data", | |
| "run_stitch_mtp", | |
| "run_training", | |
| "run_vllm_engine", | |
| "stop_mooncake_master", | |
| "stop_vllm_server", | |
| "wait_for_server", | |
| ] | |
| def purge_newfiles(fn: Callable[..., Path]): | |
| """Decorator that turns a Path-returning function into a context manager. | |
| On exit, deletes top-level files in the resolved directory whose mtime is | |
| newer than when the wrapped function returned. Does not recurse into | |
| subdirectories. This prevents generated artifacts (e.g. size-keyed | |
| ``d2t-*.npy`` and ``t2d-*.npy`` files potentially cached by ``train.py``) | |
| from persisting in | |
| shared directories (such as the HF snapshot cache) between test runs. | |
| """ | |
| def wrapper(*args, **kwargs): | |
| path = fn(*args, **kwargs) | |
| cutoff = time.time() | |
| try: | |
| yield path | |
| finally: | |
| if path.is_dir(): | |
| for f in path.iterdir(): | |
| if f.is_file() and f.stat().st_mtime > cutoff: | |
| f.unlink() | |
| logger.info("Purged generated artifact: {}", f.name) | |
| return wrapper | |
| VLLM_PYTHON = os.environ.get("VLLM_PYTHON", sys.executable) | |
| SCRIPTS_DIR = Path(__file__).resolve().parent.parent.parent / "scripts" | |
| PROCESS_GROUP_CLEANUP_TIMEOUT = 5.0 | |
| PROCESS_GROUP_POLL_INTERVAL = 0.1 | |
| def _signal_process_group(process_group_id: int, sig: int) -> None: | |
| """Signal a vLLM process group, ignoring an already-gone group.""" | |
| if process_group_id == os.getpgrp(): | |
| raise RuntimeError("Refusing to signal the test runner's process group") | |
| with suppress(ProcessLookupError): | |
| os.killpg(process_group_id, sig) | |
| def _wait_for_process_group_exit( | |
| process_group_id: int, | |
| timeout: float = PROCESS_GROUP_CLEANUP_TIMEOUT, | |
| ) -> bool: | |
| """Wait until no process remains in a vLLM process group.""" | |
| deadline = time.monotonic() + timeout | |
| while time.monotonic() < deadline: | |
| try: | |
| os.killpg(process_group_id, 0) | |
| except ProcessLookupError: | |
| return True | |
| time.sleep(PROCESS_GROUP_POLL_INTERVAL) | |
| return False | |
| def wait_for_server( | |
| port: int, | |
| timeout: float = 600.0, | |
| poll_interval: float = 2.0, | |
| process: subprocess.Popen | None = None, | |
| readiness_stability: float = 5.0, | |
| ): | |
| """Poll vLLM server health endpoint until stably ready or timeout. | |
| If *process* is provided, checks whether it has exited between polls | |
| so that startup failures are reported immediately instead of waiting | |
| for the full timeout. A continuous healthy window is required because | |
| multi-process vLLM can answer one health request before another API | |
| server process finishes starting or fails. | |
| """ | |
| logger.info("Waiting for server") | |
| url = f"http://localhost:{port}/health" | |
| deadline = time.monotonic() + timeout | |
| healthy_since: float | None = None | |
| while time.monotonic() < deadline: | |
| if process is not None and process.poll() is not None: | |
| raise RuntimeError( | |
| f"vLLM server process exited with code {process.returncode} " | |
| "before becoming ready" | |
| ) | |
| try: | |
| with urllib.request.urlopen(url, timeout=5) as resp: # noqa: S310 | |
| healthy = resp.status == 200 | |
| except (urllib.error.URLError, ConnectionError, OSError): | |
| healthy = False | |
| now = time.monotonic() | |
| if healthy: | |
| if healthy_since is None: | |
| healthy_since = now | |
| if now - healthy_since >= readiness_stability: | |
| return | |
| else: | |
| healthy_since = None | |
| time.sleep(poll_interval) | |
| raise TimeoutError(f"vLLM server on port {port} not ready after {timeout}s") | |
| def launch_vllm_server( | |
| model: str, | |
| port: int, | |
| hidden_states_path: str, | |
| *, | |
| max_model_len: int = 513, | |
| gpu_memory_utilization: float = 0.5, | |
| target_layer_ids: list[int] | None = None, | |
| enforce_eager: bool = False, | |
| allowed_local_media_path: str | None = None, | |
| hidden_states_backend: str = "file", | |
| mooncake_master: str | None = None, | |
| mooncake_metadata_server: str | None = None, | |
| mooncake_protocol: str | None = None, | |
| ) -> subprocess.Popen: | |
| """Launch a vLLM server configured for hidden-state extraction. | |
| Returns the server subprocess. Caller is responsible for stopping it | |
| via stop_vllm_server(). | |
| """ | |
| cmd = [ | |
| VLLM_PYTHON, | |
| str(SCRIPTS_DIR / "launch_vllm.py"), | |
| model, | |
| "--hidden-states-path", | |
| str(hidden_states_path), | |
| "--hidden-states-backend", | |
| hidden_states_backend, | |
| ] | |
| if mooncake_master is not None: | |
| cmd += ["--mooncake-master", mooncake_master] | |
| if mooncake_metadata_server is not None: | |
| cmd += ["--mooncake-metadata-server", mooncake_metadata_server] | |
| if mooncake_protocol is not None: | |
| cmd += ["--mooncake-protocol", mooncake_protocol] | |
| if target_layer_ids is not None: | |
| cmd += ["--target-layer-ids"] + [str(lid) for lid in target_layer_ids] | |
| if enforce_eager: | |
| cmd += ["--enforce-eager"] | |
| if allowed_local_media_path is not None: | |
| cmd += ["--allowed-local-media-path", allowed_local_media_path] | |
| cmd += [ | |
| "--", | |
| "--port", | |
| str(port), | |
| "--max-model-len", | |
| str(max_model_len), | |
| "--gpu-memory-utilization", | |
| str(gpu_memory_utilization), | |
| "--disable-uvicorn-access-log", | |
| ] | |
| logger.info("Starting vLLM server: {}", " ".join(cmd)) | |
| # vLLM creates an engine process and multiple API-server descendants. | |
| # Isolate the whole tree so teardown cannot leave workers attached to the | |
| # pytest process group or interfere with the next test's launch. | |
| process = subprocess.Popen(cmd, start_new_session=True) # noqa: S603 | |
| try: | |
| wait_for_server(port, process=process) | |
| logger.info("vLLM server ready on port {}", port) | |
| except Exception: | |
| stop_vllm_server(process) | |
| raise | |
| return process | |
| def stop_vllm_server(process: subprocess.Popen): | |
| """Gracefully stop a vLLM server and all of its descendants.""" | |
| process_group_id = process.pid | |
| if process.poll() is None: | |
| # Give vLLM's process manager the first opportunity to shut down its | |
| # children cleanly and reap them. | |
| process.terminate() | |
| try: | |
| process.wait(timeout=30) | |
| except subprocess.TimeoutExpired: | |
| _signal_process_group(process_group_id, signal.SIGTERM) | |
| try: | |
| process.wait(timeout=5) | |
| except subprocess.TimeoutExpired: | |
| _signal_process_group(process_group_id, signal.SIGKILL) | |
| process.wait(timeout=10) | |
| # The manager may have exited before all descendants were reaped. The | |
| # dedicated session lets us clean up those stragglers without touching | |
| # pytest or unrelated processes, then wait for the reset to complete. | |
| _signal_process_group(process_group_id, signal.SIGTERM) | |
| if not _wait_for_process_group_exit(process_group_id): | |
| _signal_process_group(process_group_id, signal.SIGKILL) | |
| if not _wait_for_process_group_exit(process_group_id): | |
| logger.error( | |
| "vLLM process group {} did not exit after forced cleanup", | |
| process_group_id, | |
| ) | |
| if process.returncode not in (0, -15): # -15 = SIGTERM (expected) | |
| logger.error("vLLM server exited with code {}", process.returncode) | |
| logger.info("vLLM server stopped (exit code {})", process.returncode) | |
| def launch_vllm_server_context(*args, **kwargs): | |
| process = launch_vllm_server(*args, **kwargs) | |
| try: | |
| yield | |
| finally: | |
| stop_vllm_server(process) | |
| def launch_mooncake_master(port: int) -> subprocess.Popen: | |
| """Launch a mooncake_master process. | |
| Returns the subprocess. Caller is responsible for stopping it | |
| via stop_mooncake_master(). | |
| The process is started in its own session (start_new_session=True) | |
| because the installed ``mooncake_master`` entry-point is a Python | |
| wrapper that spawns the real binary via subprocess.call(). Killing | |
| only the wrapper leaves the child binary running as an orphan. | |
| Using a dedicated session lets stop_mooncake_master() kill the | |
| entire process group at once. | |
| """ | |
| exe = shutil.which("mooncake_master") | |
| if exe is None: | |
| pytest.skip("mooncake_master not found on PATH") | |
| cmd = [exe, "--port", str(port)] | |
| logger.info("Starting mooncake_master: {}", " ".join(cmd)) | |
| proc = subprocess.Popen(cmd, start_new_session=True) # noqa: S603 | |
| time.sleep(2) | |
| if proc.poll() is not None: | |
| raise RuntimeError( | |
| f"mooncake_master exited immediately with code {proc.returncode}" | |
| ) | |
| return proc | |
| def stop_mooncake_master(process: subprocess.Popen): | |
| """Stop the mooncake_master process group.""" | |
| if process.poll() is not None: | |
| logger.info("mooncake_master already exited (code {})", process.returncode) | |
| return | |
| pgid = os.getpgid(process.pid) | |
| os.killpg(pgid, signal.SIGTERM) | |
| try: | |
| process.wait(timeout=10) | |
| except subprocess.TimeoutExpired: | |
| os.killpg(pgid, signal.SIGKILL) | |
| process.wait(timeout=10) | |
| if process.returncode not in (0, -15, -signal.SIGKILL): | |
| logger.error("mooncake_master exited with code {}", process.returncode) | |
| logger.info("mooncake_master stopped (exit code {})", process.returncode) | |
| def launch_mooncake_master_context(port: int): | |
| process = launch_mooncake_master(port) | |
| try: | |
| yield | |
| finally: | |
| stop_mooncake_master(process) | |
| def setup_dummy_sharegpt4v_coco(coco_dir: Path): | |
| """Enable ShareGPT4V to be used without downloading the actual COCO dataset.""" | |
| coco_dir.mkdir(parents=True, exist_ok=True) | |
| dummy_image = Image.new("RGB", (256, 256)) | |
| dummy_image_path = coco_dir / "dummy.png" | |
| dummy_image.save(dummy_image_path) | |
| raw_dataset, normalize_fn = load_raw_dataset("sharegpt4v_coco") | |
| # Use symlinks to avoid copying the image | |
| for raw_path in raw_dataset["image"]: | |
| image_path = coco_dir / raw_path.removeprefix("coco/") | |
| if not image_path.exists(): | |
| image_path.parent.mkdir(parents=True, exist_ok=True) | |
| image_path.symlink_to(dummy_image_path) | |
| def run_prepare_data( | |
| model: str, | |
| data: str, | |
| data_path: Path, | |
| max_samples: int = 50, | |
| seq_length: int = 512, | |
| timeout: float | None = None, | |
| render_endpoint: str | None = None, | |
| ): | |
| """Tokenize data using prepare_data.py.""" | |
| cmd = [ | |
| sys.executable, | |
| str(SCRIPTS_DIR / "prepare_data.py"), | |
| "--model", | |
| model, | |
| "--data", | |
| data, | |
| "--output", | |
| str(data_path), | |
| "--max-samples", | |
| str(max_samples), | |
| "--seq-length", | |
| str(seq_length), | |
| ] | |
| if render_endpoint is not None: | |
| cmd += ["--render-endpoint", render_endpoint] | |
| logger.info("Preparing data: {}", " ".join(cmd)) | |
| result = subprocess.run( # noqa: S603 | |
| cmd, check=False, timeout=timeout | |
| ) | |
| assert result.returncode == 0, "prepare_data.py failed" | |
| def run_data_generation_offline( | |
| data_path: Path, | |
| hidden_states_path: Path | None = None, | |
| port: int = 8321, | |
| max_samples: int = 50, | |
| concurrency: int = 4, | |
| validate_outputs: bool = True, | |
| timeout: float | None = None, | |
| fail_on_error: bool = True, | |
| ): | |
| datagen_cmd = [ | |
| sys.executable, | |
| str(SCRIPTS_DIR / "data_generation_offline.py"), | |
| "--preprocessed-data", | |
| str(data_path), | |
| "--endpoint", | |
| f"http://localhost:{port}/v1", | |
| "--max-samples", | |
| str(max_samples), | |
| "--concurrency", | |
| str(concurrency), | |
| ] | |
| if validate_outputs: | |
| datagen_cmd.append("--validate-outputs") | |
| if fail_on_error: | |
| datagen_cmd.append("--fail-on-error") | |
| if hidden_states_path is not None: | |
| datagen_cmd += ["--output", str(hidden_states_path)] | |
| logger.info("Generating hidden states offline: {}", " ".join(datagen_cmd)) | |
| result = subprocess.run( # noqa: S603 | |
| datagen_cmd, stderr=subprocess.PIPE, text=True, check=False, timeout=timeout | |
| ) | |
| assert result.returncode == 0, ( | |
| f"data_generation_offline.py failed:\n{result.stderr}" | |
| ) | |
| def run_training( | |
| model: str, | |
| data_path: Path, | |
| save_path: Path, | |
| seq_length: int = 512, | |
| port: int = 8321, | |
| draft_vocab_size: int | None = 8192, | |
| epochs: int = 1, | |
| lr: float = 3e-4, | |
| online: bool = True, | |
| hidden_states_path: Path | None = None, | |
| timeout: float | None = None, | |
| speculator_type: str = "eagle3", | |
| extra_train_args: list[str] | None = None, | |
| target_layer_ids: list[int] | None = None, | |
| num_layers: int | None = None, | |
| log_freq: int = 1, | |
| hidden_states_backend: str = "file", | |
| mooncake_master: str | None = None, | |
| mooncake_metadata_server: str | None = None, | |
| mooncake_protocol: str | None = None, | |
| ): | |
| train_cmd = [ | |
| sys.executable, | |
| str(SCRIPTS_DIR / "train.py"), | |
| "--verifier-name-or-path", | |
| model, | |
| "--data-path", | |
| str(data_path), | |
| "--vllm-endpoint", | |
| f"http://localhost:{port}/v1", | |
| "--save-path", | |
| str(save_path), | |
| "--epochs", | |
| str(epochs), | |
| "--lr", | |
| str(lr), | |
| "--total-seq-len", | |
| str(seq_length), | |
| "--speculator-type", | |
| speculator_type, | |
| "--log-freq", | |
| str(log_freq), | |
| "--hidden-states-backend", | |
| hidden_states_backend, | |
| ] | |
| if draft_vocab_size is not None: | |
| train_cmd += ["--draft-vocab-size", str(draft_vocab_size)] | |
| if mooncake_master is not None: | |
| train_cmd += ["--mooncake-master", mooncake_master] | |
| if mooncake_metadata_server is not None: | |
| train_cmd += ["--mooncake-metadata-server", mooncake_metadata_server] | |
| if mooncake_protocol is not None: | |
| train_cmd += ["--mooncake-protocol", mooncake_protocol] | |
| if online: | |
| train_cmd += [ | |
| "--on-missing", | |
| "generate", | |
| "--on-generate", | |
| "delete", | |
| ] | |
| else: | |
| train_cmd += [ | |
| "--on-missing", | |
| "raise", | |
| ] | |
| if hidden_states_path is not None: | |
| train_cmd += ["--hidden-states-path", str(hidden_states_path)] | |
| if target_layer_ids is not None: | |
| train_cmd += ["--target-layer-ids"] + [str(lid) for lid in target_layer_ids] | |
| if num_layers is not None: | |
| train_cmd += ["--num-layers", str(num_layers)] | |
| if extra_train_args: | |
| train_cmd += extra_train_args | |
| logger.info("Running training: {}", " ".join(train_cmd)) | |
| result = subprocess.run( # noqa: S603 | |
| train_cmd, stderr=subprocess.PIPE, text=True, check=False, timeout=timeout | |
| ) | |
| assert result.returncode == 0, f"train.py failed:\n{result.stderr}" | |
| def run_stitch_mtp( | |
| finetuned_checkpoint: Path, | |
| verifier_path: str, | |
| output_path: Path, | |
| timeout: float | None = None, | |
| ): | |
| cmd = [ | |
| sys.executable, | |
| str(SCRIPTS_DIR / "stitch_mtp.py"), | |
| str(finetuned_checkpoint), | |
| verifier_path, | |
| "--output-path", | |
| str(output_path), | |
| ] | |
| logger.info("Stitching MTP weights: {}", " ".join(cmd)) | |
| result = subprocess.run( # noqa: S603 | |
| cmd, capture_output=True, text=True, check=False, timeout=timeout | |
| ) | |
| assert result.returncode == 0, f"stitch_mtp.py failed:\n{result.stderr}" | |
| def run_vllm_engine( | |
| model_path: str, | |
| tmp_path: Path, | |
| prompts: list[list[dict[str, str]]], | |
| max_model_len: int = 1024, | |
| gpu_memory_utilization: float = 0.8, | |
| enforce_eager: bool = False, | |
| allowed_local_media_path: str | None = None, | |
| speculative_config: dict | None = None, | |
| disable_compile_cache: bool = False, | |
| max_tokens: int = 50, | |
| ignore_eos: bool = True, | |
| acceptance_thresholds: Iterable[float] | None = None, | |
| timeout: float | None = None, | |
| ): | |
| VLLM_PYTHON = os.environ.get("VLLM_PYTHON", sys.executable) | |
| logger.info("vLLM Python executable: {}", VLLM_PYTHON) | |
| run_vllm_file = str(Path(__file__).with_name("run_vllm.py")) | |
| results_file = str(tmp_path / "results.json") | |
| sampling_params_dict = { | |
| "temperature": 0, | |
| "top_p": 0.9, | |
| "max_tokens": max_tokens, | |
| "ignore_eos": ignore_eos, | |
| } | |
| llm_args_dict = { | |
| "model": model_path, | |
| "max_model_len": max_model_len, | |
| "gpu_memory_utilization": gpu_memory_utilization, | |
| "enforce_eager": enforce_eager, | |
| } | |
| if allowed_local_media_path is not None: | |
| llm_args_dict["allowed_local_media_path"] = allowed_local_media_path | |
| if speculative_config is not None: | |
| llm_args_dict["speculative_config"] = speculative_config | |
| command = [ | |
| VLLM_PYTHON, | |
| run_vllm_file, | |
| "--sampling-params-args", | |
| json.dumps(sampling_params_dict), | |
| "--llm-args", | |
| json.dumps(llm_args_dict), | |
| "--prompts", | |
| json.dumps(prompts), | |
| "--results-file", | |
| results_file, | |
| ] | |
| logger.info("run_vllm.py command:\n {}", command) | |
| # Set environment variables for subprocess | |
| env = os.environ.copy() | |
| if disable_compile_cache: | |
| env["VLLM_DISABLE_COMPILE_CACHE"] = "1" | |
| logger.info("Disabling vLLM compile cache for this test") | |
| result = subprocess.run( # noqa: S603 | |
| command, | |
| stdout=subprocess.PIPE, | |
| stderr=subprocess.STDOUT, | |
| text=True, | |
| check=False, | |
| env=env, | |
| timeout=timeout, | |
| ) | |
| logger.info("run_vllm.py output:\n{}", indent(result.stdout, " ")) | |
| returncode = result.returncode | |
| assert returncode == 0, ( | |
| f"run_vllm.py command exited with non-zero return code: {returncode}" | |
| ) | |
| with Path(results_file).open(encoding="utf-8") as f: | |
| results_dict = json.load(f) | |
| outputs_token_ids = results_dict["outputs"] | |
| metrics_dict = results_dict["metrics"] | |
| logger.info("outputs_token_ids: {}", outputs_token_ids) | |
| logger.info("metrics_dict: {}", metrics_dict) | |
| for output_token_ids in outputs_token_ids: | |
| # If max_tokens is 100 or less, make sure the output length is max_tokens | |
| assert max_tokens > 100 or len(output_token_ids) == max_tokens | |
| assert all(isinstance(token, int) for token in output_token_ids) | |
| if acceptance_thresholds is not None: | |
| for i, thresholdi in enumerate(acceptance_thresholds): | |
| assert f"acceptance_at_token_{i}" in metrics_dict, ( | |
| f"Acceptance at token {i} is not in metrics_dict" | |
| ) | |
| acci = metrics_dict[f"acceptance_at_token_{i}"] | |
| assert acci >= thresholdi, ( | |
| f"Acceptance {acci} at token {i} is less than threshold {thresholdi}" | |
| ) | |