k3-a40-bootstrap / drive_pod.py
patdev's picture
Sauvegarde 22/08 : resultats, rapport, docs, scripts, Dockerfile image v3
2454f65 verified
Raw History Blame Contribute Delete
5.16 kB
"""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=<dtype|''>,spec=<off|dspark|mtp|on>[,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()