"""Enchaine les variantes sur le pod sans jamais le recreer. Le pod recharge le bootstrap a chaud des qu'il change sur le Hub. On modifie donc les DEFAUTS du bloc EXPERIENCE (VL_KV / VL_SPEC), on republie, on attend que vLLM redemarre avec les nouveaux flags (verifie dans son journal, pas suppose), puis on mesure avec pod_bench.py. Les poids restent sur le disque : un redemarrage coute ~2-3 min, pas un retelechargement. python drive_pod.py --variante kv=turboquant_k3v4_nc,spec=off """ from __future__ import annotations import argparse import json import os import re import subprocess import sys import time import urllib.request ICI = os.path.dirname(os.path.abspath(__file__)) BOOT = r"g:\Environements\HuggingFace\vllm_bootstrap.sh" BENCH = r"g:\Environements\HuggingFace\pod_bench.py" POD = open(os.path.join(ICI, "podid")).read().strip() BASE = f"https://{POD}-8080.proxy.runpod.net" SSH_HOST, SSH_PORT = open(os.path.join(ICI, "podssh")).read().split() # "ip port", ecrit par le moniteur KEY = os.path.join(ICI, "ssh", "pod_key") def ssh(cmd: str, timeout: int = 60) -> str: r = subprocess.run(["ssh", "-i", KEY, "-o", "StrictHostKeyChecking=no", "-o", "UserKnownHostsFile=/dev/null", "-o", "ConnectTimeout=20", "-o", "LogLevel=ERROR", "-p", SSH_PORT, f"root@{SSH_HOST}", cmd], capture_output=True, text=True, timeout=timeout) return r.stdout def publier(kv: str, spec: str, n_dspark: int) -> None: s = open(BOOT, encoding="utf-8", newline="").read() s = re.sub(r'^: "\$\{VL_KV:=[^}]*\}"', f': "${{VL_KV:={kv}}}"', s, flags=re.M) s = re.sub(r'^: "\$\{VL_SPEC:=[^}]*\}"', f': "${{VL_SPEC:={spec}}}"', s, flags=re.M) s = re.sub(r'^: "\$\{VL_DSPARK_N:=[^}]*\}"', f': "${{VL_DSPARK_N:={n_dspark}}}"', s, flags=re.M) open(BOOT, "w", encoding="utf-8", newline="").write(s) # (pas de `bash -n` ici : depuis Python, `bash` resout vers le bash WSL, # qui ne lit pas les chemins Windows ; la syntaxe est verifiee a la main) subprocess.run(["hf", "upload", "patdev/k3-a40-bootstrap", BOOT, "vllm_bootstrap.sh", "--commit-message", f"experience kv={kv or 'bf16'} spec={spec}"], check=True, capture_output=True) print(f" publie : VL_KV={kv or '(bf16)'} VL_SPEC={spec}", flush=True) def attendre(kv: str, spec: str, budget: int = 900) -> bool: """Vrai quand vLLM a redemarre AVEC les flags voulus et repond.""" t0 = time.time() motif_kv = f"'kv_cache_dtype': '{kv}'" if kv else None motif_spec = {"dspark": "dspark", "mtp": "'method': 'mtp'", "on": "ngram"}.get(spec) while time.time() - t0 < budget: log = ssh("grep -h 'non-default args' /tmp/vllm.log 2>/dev/null | tail -1 | cut -c1-6000; " "echo ---; grep -c 'Application startup complete' /tmp/vllm.log 2>/dev/null") args, _, pret = log.partition("---") ok_kv = (motif_kv in args) if motif_kv else ("kv_cache_dtype" not in args) ok_spec = (motif_spec in args) if motif_spec else ("speculative_config" not in args) if ok_kv and ok_spec and pret.strip() not in ("", "0"): try: urllib.request.urlopen(BASE + "/v1/models", timeout=20) print(f" pret en {time.time() - t0:.0f}s", flush=True) return True except Exception: # noqa: BLE001 pass time.sleep(20) print(" DELAI depasse ; journal :", ssh("tail -n 5 /tmp/vllm.log | cut -c1-200")) return False def mesurer(etiquette: str, seqs: str, max_tokens: int) -> dict: r = subprocess.run([sys.executable, BENCH, "--base", BASE, "--seqs", seqs, "--max-tokens", str(max_tokens)], capture_output=True, text=True, timeout=1800) print(r.stdout, flush=True) lignes = [l for l in r.stdout.splitlines() if l.startswith("[")] res = json.loads(lignes[-1]) if lignes else [] kv_tok = ssh("grep -h 'GPU KV cache size' /tmp/vllm.log | tail -1 | grep -oE '[0-9,]+ tokens'") out = {"variante": etiquette, "kv_jetons": kv_tok.strip(), "debits": res} chemin = os.path.join(ICI, "mesures_pod.jsonl") with open(chemin, "a", encoding="utf-8") as f: f.write(json.dumps(out) + "\n") return out def main() -> None: ap = argparse.ArgumentParser() ap.add_argument("--variante", required=True, help="kv=,spec=[,n=8]") ap.add_argument("--seqs", default="1,4,8,16,32") ap.add_argument("--max-tokens", type=int, default=256) ap.add_argument("--sans-publier", action="store_true", help="mesurer l'etat courant") a = ap.parse_args() kv, spec, n = "", "off", 8 for part in a.variante.split(","): k, _, v = part.partition("=") if k == "kv": kv = v elif k == "spec": spec = v elif k == "n": n = int(v) if not a.sans_publier: publier(kv, spec, n) if not attendre(kv, spec): sys.exit(2) print(json.dumps(mesurer(a.variante, a.seqs, a.max_tokens), ensure_ascii=False)) if __name__ == "__main__": main()