diff --git a/scripts/definition_smoke.py b/scripts/definition_smoke.py new file mode 100644 index 0000000..97d4c42 --- /dev/null +++ b/scripts/definition_smoke.py @@ -0,0 +1,354 @@ +"""Smoke test engine definitions on a deployed Lilo control plane. + +For each definition this creates a training client, runs a cross-entropy step, +publishes the weights and samples from them, then runs an importance-sampling +step on the sampled response and samples once more. The model is unloaded +afterwards so the trainer is released. + +Usage:: + + export TINKER_BASE_URL=https://...modal.run + export TINKER_API_KEY=tml-lilo-... + uv run scripts/definition_smoke.py # default definition + uv run scripts/definition_smoke.py --list # show all definitions + uv run scripts/definition_smoke.py --parallel \ + --definition-id qwen3_5_9b_miles_lora_16k \ + --definition-id qwen3_5_4b_full_64k + +Any id from ``lilo.providers.modal.definitions`` works; the client type +(full or LoRA) and the LoRA target flags are read from the definition module so +the request matches what the deployment accepts. The definition id is passed as +``base_model`` so non-cataloged definitions can be targeted directly. +Definitions run sequentially unless ``--parallel`` is set. +""" + +from __future__ import annotations + +import argparse +import json +import math +import os +import time +import traceback +from concurrent.futures import ThreadPoolExecutor +from pathlib import Path +from typing import Any + +import httpx +import tinker +from tinker import types + +from lilo.backends.miles_config import lora_target_flags +from lilo.client import create_full_training_client +from lilo.providers.modal.app import DEFINITIONS, module_for + +DEFAULT_DEFINITION = "qwen3_8_27b_miles_lora_64k" +TIMEOUT = 60 * 60 +PROMPT = "Question: What is two plus two?\nAnswer:" + + +def _finite_metrics(metrics: dict[str, Any]) -> dict[str, float]: + numeric = { + key: float(value) + for key, value in metrics.items() + if isinstance(value, int | float) + } + if not numeric or any(not math.isfinite(value) for value in numeric.values()): + raise RuntimeError(f"non-finite metrics: {metrics}") + return numeric + + +def _sft_datum(tokenizer) -> types.Datum: + prompt = tokenizer.encode(PROMPT, add_special_tokens=True) + completion = tokenizer.encode(" 4", add_special_tokens=False) + tokens = prompt + completion + return types.Datum( + model_input=types.ModelInput.from_ints(tokens[:-1]), + loss_fn_inputs={ + "target_tokens": tokens[1:], + "weights": [0.0] * (len(prompt) - 1) + [1.0] * len(completion), + }, + ) + + +def _rl_datum( + prompt: list[int], + response: list[int], + logprobs: list[float], + reward: float, +) -> types.Datum: + prompt_targets = len(prompt) - 1 + return types.Datum( + model_input=types.ModelInput.from_ints(prompt + response[:-1]), + loss_fn_inputs={ + "target_tokens": prompt[1:] + response, + "logprobs": [0.0] * prompt_targets + logprobs, + "advantages": [0.0] * prompt_targets + [reward] * len(response), + }, + ) + + +def _step(training, data: list[types.Datum], loss_fn: str, lr: float) -> dict: + started = time.perf_counter() + forward_result = training.forward_backward(data, loss_fn).result(timeout=TIMEOUT) + forward_done = time.perf_counter() + if len(forward_result.loss_fn_outputs) != len(data): + raise RuntimeError( + f"expected {len(data)} loss outputs, got " + f"{len(forward_result.loss_fn_outputs)}" + ) + optimizer_result = training.optim_step(types.AdamParams(learning_rate=lr)).result( + timeout=TIMEOUT + ) + finished = time.perf_counter() + optimizer_metrics = _finite_metrics(optimizer_result.metrics) + if optimizer_metrics.get("update_successful:mean") != 1.0: + raise RuntimeError(f"optimizer step skipped the update: {optimizer_metrics}") + return { + "loss_fn": loss_fn, + "metrics": _finite_metrics(forward_result.metrics), + "optimizer_metrics": optimizer_metrics, + "forward_backward_seconds": forward_done - started, + "optimizer_seconds": finished - forward_done, + } + + +def _publish_and_sample( + training, tokenizer, prompt: list[int], max_tokens: int +) -> tuple[dict, list[int], list[float]]: + started = time.perf_counter() + sampling = training.save_weights_and_get_sampling_client() + published = time.perf_counter() + result = sampling.sample( + prompt=types.ModelInput.from_ints(prompt), + num_samples=1, + sampling_params=types.SamplingParams(max_tokens=max_tokens, temperature=1.0), + ).result(timeout=TIMEOUT) + finished = time.perf_counter() + if len(result.sequences) != 1: + raise RuntimeError(f"expected one sequence, got {len(result.sequences)}") + sequence = result.sequences[0] + tokens = list(sequence.tokens) + logprobs = [float(value) for value in sequence.logprobs or ()] + if not tokens or len(logprobs) != len(tokens): + raise RuntimeError("sample did not return matching tokens and logprobs") + if any(not math.isfinite(value) for value in logprobs): + raise RuntimeError(f"non-finite sample logprobs: {logprobs}") + text = tokenizer.decode(tokens) + return ( + { + "publish_seconds": published - started, + "sample_seconds": finished - published, + "output_tokens": len(tokens), + "text": text, + }, + tokens, + logprobs, + ) + + +def _unload(base_url: str, api_key: str, model_id: str) -> None: + headers = {"X-API-Key": api_key} + with httpx.Client(base_url=base_url, headers=headers, timeout=60) as client: + response = client.post("/api/v1/unload_model", json={"model_id": model_id}) + if response.status_code == 404: + return + response.raise_for_status() + request_id = response.json()["request_id"] + while True: + response = client.post( + "/api/v1/retrieve_future", json={"request_id": request_id} + ) + if response.status_code != 408: + response.raise_for_status() + return + time.sleep(1) + + +def _create_training_client( + service: tinker.ServiceClient, definition: Any, rank: int | None +) -> tuple[tinker.TrainingClient, dict[str, Any]]: + definition_id = definition.DEFINITION_ID + if definition.PARAMETERIZATION == "full": + training = create_full_training_client(service, definition_id) + return training, {"parameterization": "full"} + train_attn, train_mlp, train_unembed = lora_target_flags(definition.TARGET_MODULES) + if rank is None: + rank = min(16, definition.MAX_LORA_RANK) + training = service.create_lora_training_client( + base_model=definition_id, + rank=rank, + train_attn=train_attn, + train_mlp=train_mlp, + train_unembed=train_unembed, + ) + return training, { + "parameterization": "lora", + "rank": rank, + "train_attn": train_attn, + "train_mlp": train_mlp, + "train_unembed": train_unembed, + } + + +def _run_definition( + definition_id: str, + *, + base_url: str, + api_key: str, + rank: int | None, + max_tokens: int, +) -> dict: + report: dict[str, Any] = { + "definition_id": definition_id, + "status": "running", + "started_at": time.time(), + "phases": {}, + } + training = None + try: + definition = module_for(definition_id) + service = tinker.ServiceClient(base_url=base_url, api_key=api_key) + started = time.perf_counter() + training, spec = _create_training_client(service, definition, rank) + info = training.get_info() + if info.is_lora != (spec["parameterization"] == "lora"): + raise RuntimeError(f"expected {spec['parameterization']} model, got {info}") + tokenizer = training.get_tokenizer() + report["phases"]["provision"] = { + "seconds": time.perf_counter() - started, + "model_id": str(training.model_id), + "base_model": info.model_name, + "lora_rank": info.lora_rank, + **spec, + } + print(json.dumps({"definition": definition_id, **report["phases"]})) + + report["phases"]["sft_step"] = _step( + training, [_sft_datum(tokenizer)], "cross_entropy", 1e-4 + ) + print(json.dumps({"definition": definition_id, "sft_step": "ok"})) + + prompt = tokenizer.encode(PROMPT, add_special_tokens=True) + sample, tokens, logprobs = _publish_and_sample( + training, tokenizer, prompt, max_tokens + ) + report["phases"]["sample_1"] = sample + print(json.dumps({"definition": definition_id, "sample_1": sample["text"]})) + + reward = 1.0 if "4" in sample["text"] else -1.0 + report["phases"]["rl_step"] = _step( + training, + [_rl_datum(prompt, tokens, logprobs, reward)], + "importance_sampling", + 1e-5, + ) + report["phases"]["rl_step"]["reward"] = reward + print(json.dumps({"definition": definition_id, "rl_step": "ok"})) + + sample, _, _ = _publish_and_sample(training, tokenizer, prompt, max_tokens) + report["phases"]["sample_2"] = sample + print(json.dumps({"definition": definition_id, "sample_2": sample["text"]})) + + report["status"] = "passed" + except Exception as exc: # noqa: BLE001 - report every failure per definition + report["status"] = "failed" + report["error"] = f"{type(exc).__name__}: {exc}" + report["traceback"] = traceback.format_exc() + finally: + report["finished_at"] = time.time() + report["elapsed_seconds"] = report["finished_at"] - report["started_at"] + if training is not None: + try: + _unload(base_url, api_key, str(training.model_id)) + report["cleanup"] = {"model_unloaded": True} + except (httpx.HTTPError, RuntimeError, TimeoutError, ValueError) as exc: + error = f"{type(exc).__name__}: {exc}" + report["cleanup"] = {"model_unloaded": False, "error": error} + report["status"] = "failed" + report.setdefault("error", f"cleanup failed: {error}") + return report + + +def main() -> None: + parser = argparse.ArgumentParser(description=__doc__.split("\n\n")[0]) + known = [definition.DEFINITION_ID for definition in DEFINITIONS] + parser.add_argument( + "--definition-id", + action="append", + dest="definition_ids", + choices=known, + metavar="ID", + help=f"engine definition id (repeatable); default {DEFAULT_DEFINITION}", + ) + parser.add_argument( + "--list", action="store_true", help="print known definitions and exit" + ) + parser.add_argument("--base-url", default=os.environ.get("TINKER_BASE_URL")) + parser.add_argument( + "--rank", + type=int, + help="LoRA rank; default min(16, definition MAX_LORA_RANK)", + ) + parser.add_argument("--max-tokens", type=int, default=32) + parser.add_argument("--parallel", action="store_true") + parser.add_argument( + "--output", + type=Path, + default=Path("scripts/results/definition_smoke.json"), + ) + args = parser.parse_args() + if args.list: + for definition in DEFINITIONS: + print( + f"{definition.DEFINITION_ID:40} {definition.PARAMETERIZATION:5} " + f"{definition.GPUS}x{definition.GPU_TYPE} " + f"ctx={definition.MAX_CONTEXT_LENGTH}" + + ("" if definition.CATALOG_VISIBLE else " (not cataloged)") + ) + return + if not args.base_url: + parser.error("--base-url or TINKER_BASE_URL is required") + api_key = os.environ.get("TINKER_API_KEY") + if not api_key: + parser.error("TINKER_API_KEY is required") + definition_ids = args.definition_ids or [DEFAULT_DEFINITION] + + def run(definition_id: str) -> dict: + return _run_definition( + definition_id, + base_url=args.base_url, + api_key=api_key, + rank=args.rank, + max_tokens=args.max_tokens, + ) + + if args.parallel: + with ThreadPoolExecutor(max_workers=len(definition_ids)) as pool: + reports = list(pool.map(run, definition_ids)) + else: + reports = [run(definition_id) for definition_id in definition_ids] + + stamp = f"{time.strftime('%Y%m%d%H%M%S')}.{os.getpid()}" + output = args.output.with_name(f"{args.output.stem}.{stamp}{args.output.suffix}") + output.parent.mkdir(parents=True, exist_ok=True) + output.write_text( + json.dumps({"base_url": args.base_url, "results": reports}, indent=2) + "\n", + encoding="utf-8", + ) + for report in reports: + line = { + "definition": report["definition_id"], + "status": report["status"], + "elapsed_seconds": round(report["elapsed_seconds"], 1), + } + if report["status"] != "passed": + line["error"] = report.get("error") + print(json.dumps(line)) + print(output) + if any(report["status"] != "passed" for report in reports): + raise SystemExit(1) + + +if __name__ == "__main__": + main() diff --git a/src/lilo/backends/miles_config.py b/src/lilo/backends/miles_config.py index ee90897..cf108af 100644 --- a/src/lilo/backends/miles_config.py +++ b/src/lilo/backends/miles_config.py @@ -18,6 +18,23 @@ "linear_fc2": ("down_proj",), "output_layer": ("lm_head",), } +_ATTN_LEAVES = frozenset( + {"linear_qkv", "linear_q", "linear_k", "linear_v", "linear_proj"} +) +_MLP_LEAVES = frozenset( + {"linear_fc1", "linear_fc1_gate", "linear_fc1_up", "linear_fc2"} +) +_UNEMBED_LEAVES = frozenset({"output_layer"}) + + +def lora_target_flags(target_modules: tuple[str, ...]) -> tuple[bool, bool, bool]: + """Return ``(train_attn, train_mlp, train_unembed)`` implied by target modules.""" + leaves = {module.rsplit(".", 1)[-1] for module in target_modules} + return ( + bool(leaves & _ATTN_LEAVES), + bool(leaves & _MLP_LEAVES), + bool(leaves & _UNEMBED_LEAVES), + ) @dataclass(frozen=True, slots=True) diff --git a/src/lilo/backends/miles_lora.py b/src/lilo/backends/miles_lora.py index 2e88f30..6a59f83 100644 --- a/src/lilo/backends/miles_lora.py +++ b/src/lilo/backends/miles_lora.py @@ -25,7 +25,11 @@ ModelSpec, SamplerPublication, ) -from .miles_config import MilesBackendConfig, parse_backend_config +from .miles_config import ( + MilesBackendConfig, + lora_target_flags, + parse_backend_config, +) from .miles_runtime.data import build_outputs, pad_slot_rows, prepare_batch from .miles_runtime.profiling import RankProfiler, StepPhaseTimer, TorchProfileConfig from .miles_runtime.runtime import MilesRuntime @@ -552,18 +556,7 @@ def _validate_job(self, state: MilesJobState) -> None: ) if state.seed is not None: raise ValueError("Miles multi-LoRA does not support per-model seeds") - leaves = {module.rsplit(".", 1)[-1] for module in self.config.target_modules} - configured = ( - bool( - leaves - & {"linear_qkv", "linear_q", "linear_k", "linear_v", "linear_proj"} - ), - bool( - leaves - & {"linear_fc1", "linear_fc1_gate", "linear_fc1_up", "linear_fc2"} - ), - bool(leaves & {"output_layer"}), - ) + configured = lora_target_flags(self.config.target_modules) requested = (state.train_attn, state.train_mlp, state.train_unembed) if requested != configured: raise ValueError(