Files
fallout-venice/tools/llm_crash_test.py
T

328 lines
13 KiB
Python

"""
LLM Crash Test — Fallout: Venice of Wasteland
Test de performance des deux LLM (MJ 14b + PNJ 7b) en conditions réelles.
Modes:
--mode simultane : les deux LLM tournent en parallèle (threads)
--mode decale : push progressif PNJ → MJ (pipeline producteur/consommateur)
--mode solo-mj : seulement le 14b
--mode solo-pnj : seulement le 7b
Usage:
python llm_crash_test.py --mode simultane --rounds 5
python llm_crash_test.py --mode decale --rounds 10 --pnj-count 3
"""
import os, sys, time, json, argparse, threading, queue
import urllib.request
from datetime import datetime
OLLAMA_URL = os.getenv("OLLAMA_URL", "http://localhost:11434")
MODEL_MJ = os.getenv("MODEL_MJ", "qwen2.5:14b")
MODEL_PNJ = os.getenv("MODEL_PNJ", "qwen2.5:7b")
# ---------------------------------------------------------------------------
# Prompts de test réalistes (contexte Fallout Venice)
# ---------------------------------------------------------------------------
PNJ_PROMPTS = [
"Tu es Remy Tureaud, marchand créole de Pearl River. Réponds en 2 phrases à : 'T'as des piles ?'",
"Tu es Sœur Eulalie, soigneuse de L'Union. Réponds en 2 phrases à : 'J'ai une blessure par balle.'",
"Tu es Jacques-Henri, garde de la Régie. Réponds en 2 phrases à : 'Laisse-moi passer le checkpoint.'",
"Tu es Mama Voodoo, goule du Grand Krewe. Réponds en 2 phrases à : 'C'est quoi ce tatouage ?'",
"Tu es un Écumeur anonyme. Réponds en 2 phrases à : 'On veut juste passer.'",
]
MJ_PROMPTS = [
"Résume en 3 phrases la situation politique entre L'Union et la CdA à New Orleans post-apo.",
"Décris en 3 phrases une rencontre de nuit dans les bayous de Pearl River.",
"En 3 phrases, quelles sont les conséquences d'un blocus sur la route de Baton Rouge ?",
"Décris en 3 phrases l'ambiance d'un marché noir dans le Vieux Carré de NOLA.",
"En 3 phrases, comment réagit le Grand Krewe si des étrangers fouillent leurs ruines ?",
]
# ---------------------------------------------------------------------------
# LLM call
# ---------------------------------------------------------------------------
def call_llm(model: str, prompt: str, timeout: int = 60) -> dict:
"""Appel Ollama, retourne {tokens, elapsed, error, text}."""
t0 = time.time()
body = json.dumps({
"model": model,
"prompt": prompt,
"stream": False,
"options": {"num_predict": 150, "temperature": 0.7},
}).encode()
req = urllib.request.Request(
f"{OLLAMA_URL}/api/generate", data=body, method="POST",
headers={"Content-Type": "application/json"}
)
try:
with urllib.request.urlopen(req, timeout=timeout) as r:
data = json.loads(r.read())
elapsed = time.time() - t0
tokens = data.get("eval_count", 0)
return {
"model": model,
"tokens": tokens,
"elapsed": round(elapsed, 2),
"tok_s": round(tokens / elapsed, 1) if elapsed > 0 else 0,
"error": None,
"text": data.get("response", "")[:120],
}
except Exception as e:
return {
"model": model,
"tokens": 0,
"elapsed": round(time.time() - t0, 2),
"tok_s": 0,
"error": str(e)[:80],
"text": "",
}
# ---------------------------------------------------------------------------
# Modes
# ---------------------------------------------------------------------------
def run_solo(model: str, prompts: list[str], rounds: int) -> list[dict]:
results = []
for i in range(rounds):
prompt = prompts[i % len(prompts)]
print(f" [{model}] Round {i+1}/{rounds}...", end="", flush=True)
r = call_llm(model, prompt)
print(f" {r['tok_s']} tok/s | {r['elapsed']}s | {'ERR: '+r['error'] if r['error'] else 'OK'}")
results.append(r)
return results
def run_simultane(rounds: int, pnj_count: int = 2) -> dict:
"""Les deux LLM tournent en parallèle via threads."""
print(f"\n[SIMULTANE] {MODEL_MJ} + {MODEL_PNJ} x{pnj_count} PNJ — {rounds} rounds\n")
all_results = {"mj": [], "pnj": []}
lock = threading.Lock()
def mj_worker(round_idx: int):
prompt = MJ_PROMPTS[round_idx % len(MJ_PROMPTS)]
r = call_llm(MODEL_MJ, prompt)
with lock:
all_results["mj"].append(r)
status = f"ERR: {r['error']}" if r['error'] else f"{r['tok_s']} tok/s"
print(f" [MJ 14b] R{round_idx+1}{status} ({r['elapsed']}s)")
def pnj_worker(round_idx: int, pnj_idx: int):
prompt = PNJ_PROMPTS[(round_idx * pnj_count + pnj_idx) % len(PNJ_PROMPTS)]
r = call_llm(MODEL_PNJ, prompt)
with lock:
all_results["pnj"].append(r)
status = f"ERR: {r['error']}" if r['error'] else f"{r['tok_s']} tok/s"
print(f" [PNJ 7b] R{round_idx+1} PNJ{pnj_idx+1}{status} ({r['elapsed']}s)")
t_start = time.time()
for i in range(rounds):
threads = []
t = threading.Thread(target=mj_worker, args=(i,))
threads.append(t)
for j in range(pnj_count):
t = threading.Thread(target=pnj_worker, args=(i, j))
threads.append(t)
for t in threads:
t.start()
for t in threads:
t.join()
print()
total = time.time() - t_start
return {**all_results, "total_elapsed": round(total, 1)}
def run_decale(rounds: int, pnj_count: int = 3) -> dict:
"""Pipeline producteur/consommateur : les PNJ accumulent des infos → poussé vers le MJ."""
print(f"\n[DÉCALÉ] {pnj_count} PNJ 7b → buffer → MJ 14b — {rounds} rounds\n")
pnj_queue: queue.Queue = queue.Queue()
all_results = {"mj": [], "pnj": [], "latency_pnj_to_mj": []}
stop_flag = threading.Event()
def pnj_producer():
idx = 0
while not stop_flag.is_set():
prompt = PNJ_PROMPTS[idx % len(PNJ_PROMPTS)]
r = call_llm(MODEL_PNJ, prompt)
r["produced_at"] = time.time()
pnj_queue.put(r)
with threading.Lock():
status = f"ERR: {r['error']}" if r['error'] else f"{r['tok_s']} tok/s"
print(f" [PNJ 7b ] prod #{idx+1}{status} ({r['elapsed']}s) | queue={pnj_queue.qsize()}")
idx += 1
def mj_consumer():
consumed = 0
while consumed < rounds:
# Attendre qu'il y ait au moins 1 item dans la queue (ou timeout)
try:
pnj_result = pnj_queue.get(timeout=90)
except queue.Empty:
print(" [MJ 14b] Timeout attente PNJ")
break
latency = time.time() - pnj_result["produced_at"]
# Construire le prompt MJ à partir du résultat PNJ
pnj_text = pnj_result.get("text", "(vide)")
mj_prompt = (
f"Un PNJ vient de parler : \"{pnj_text}\"\n"
f"En tant que MJ, en 2 phrases, quelle est la conséquence narrative ?"
)
r = call_llm(MODEL_MJ, mj_prompt)
r["latency_from_pnj"] = round(latency, 2)
all_results["mj"].append(r)
all_results["latency_pnj_to_mj"].append(round(latency + r["elapsed"], 2))
status = f"ERR: {r['error']}" if r['error'] else f"{r['tok_s']} tok/s"
print(f" [MJ 14b ] cons #{consumed+1}{status} ({r['elapsed']}s) | pipeline={round(latency+r['elapsed'],1)}s")
consumed += 1
pnj_queue.task_done()
print()
# Lancer N threads PNJ + 1 thread MJ
pnj_threads = [threading.Thread(target=pnj_producer, daemon=True) for _ in range(pnj_count)]
mj_thread = threading.Thread(target=mj_consumer)
t_start = time.time()
for t in pnj_threads:
t.start()
mj_thread.start()
mj_thread.join()
stop_flag.set()
# Récupérer tous les résultats PNJ dans la queue restante
while not pnj_queue.empty():
all_results["pnj"].append(pnj_queue.get())
total = time.time() - t_start
return {**all_results, "total_elapsed": round(total, 1)}
# ---------------------------------------------------------------------------
# Rapport
# ---------------------------------------------------------------------------
def print_report(results: dict, mode: str):
def stats(lst: list[dict]) -> dict:
if not lst:
return {}
ok = [r for r in lst if not r.get("error")]
err = len(lst) - len(ok)
if not ok:
return {"calls": len(lst), "errors": err}
tok_s = [r["tok_s"] for r in ok]
elaps = [r["elapsed"] for r in ok]
return {
"calls": len(lst),
"errors": err,
"avg_tok_s": round(sum(tok_s) / len(tok_s), 1),
"min_tok_s": min(tok_s),
"max_tok_s": max(tok_s),
"avg_elapsed": round(sum(elaps) / len(elaps), 1),
"total_tokens": sum(r["tokens"] for r in ok),
}
sep = "=" * 60
print(f"\n{sep}")
print(f" RAPPORT — MODE {mode.upper()}")
print(f" {datetime.now().strftime('%Y-%m-%d %H:%M:%S')}")
print(sep)
if "mj" in results and results["mj"]:
s = stats(results["mj"])
print(f"\n MJ 14b ({MODEL_MJ})")
for k, v in s.items():
print(f" {k:20s} : {v}")
if "pnj" in results and results["pnj"]:
s = stats(results["pnj"])
print(f"\n PNJ 7b ({MODEL_PNJ})")
for k, v in s.items():
print(f" {k:20s} : {v}")
if "latency_pnj_to_mj" in results and results["latency_pnj_to_mj"]:
lat = results["latency_pnj_to_mj"]
print(f"\n Pipeline PNJ→MJ (latence totale)")
print(f" avg_pipeline_s : {round(sum(lat)/len(lat), 1)}")
print(f" min_pipeline_s : {min(lat)}")
print(f" max_pipeline_s : {max(lat)}")
print(f"\n Durée totale : {results.get('total_elapsed', '?')}s")
# Score de viabilité
mj_ok = not any(r.get("error") for r in results.get("mj", []))
pnj_ok = not any(r.get("error") for r in results.get("pnj", []))
mj_speed = stats(results.get("mj", [])).get("avg_tok_s", 0)
pnj_speed = stats(results.get("pnj", [])).get("avg_tok_s", 0)
print(f"\n DIAGNOSTIC")
print(f" MJ stable : {'' if mj_ok else '✗ ERREURS'}")
print(f" PNJ stable : {'' if pnj_ok else '✗ ERREURS'}")
if mj_speed:
viable = "✓ VIABLE" if mj_speed > 8 else "⚠ LENT (< 8 tok/s)" if mj_speed > 3 else "✗ TROP LENT"
print(f" MJ débit : {mj_speed} tok/s → {viable}")
if pnj_speed:
viable = "✓ VIABLE" if pnj_speed > 15 else "⚠ LENT" if pnj_speed > 5 else "✗ TROP LENT"
print(f" PNJ débit : {pnj_speed} tok/s → {viable}")
print(sep)
# Export JSON
RESULTS_PATH = os.getenv("CRASH_RESULTS_PATH", "/home/ubuntu/fallout-venice/src/config/crash_results.json")
report_path = RESULTS_PATH
try:
with open(report_path, "w") as f:
json.dump({"mode": mode, "results": results,
"stats": {"mj": stats(results.get("mj",[])),
"pnj": stats(results.get("pnj",[]))}}, f, indent=2)
print(f"\n Rapport JSON : {report_path}")
except Exception:
pass
# ---------------------------------------------------------------------------
# CLI
# ---------------------------------------------------------------------------
def main():
parser = argparse.ArgumentParser(description="LLM Crash Test — Fallout Venice")
parser.add_argument("--mode", choices=["simultane","decale","solo-mj","solo-pnj"],
default="simultane")
parser.add_argument("--rounds", type=int, default=5, help="Nombre de rounds MJ")
parser.add_argument("--pnj-count", type=int, default=2, help="PNJ parallèles (modes simultane/decale)")
args = parser.parse_args()
print(f"=== LLM CRASH TEST — {args.mode.upper()} ===")
print(f"Ollama : {OLLAMA_URL}")
print(f"MJ : {MODEL_MJ}")
print(f"PNJ : {MODEL_PNJ}")
print(f"Rounds : {args.rounds} | PNJ parallèles : {args.pnj_count}")
# Vérifier que Ollama répond
try:
with urllib.request.urlopen(f"{OLLAMA_URL}/api/tags", timeout=5) as r:
models = [m["name"] for m in json.loads(r.read()).get("models", [])]
print(f"Modèles disponibles : {', '.join(models[:5])}")
except Exception as e:
print(f"[ERR] Ollama inaccessible : {e}")
sys.exit(1)
t_global = time.time()
if args.mode == "simultane":
results = run_simultane(args.rounds, args.pnj_count)
elif args.mode == "decale":
results = run_decale(args.rounds, args.pnj_count)
elif args.mode == "solo-mj":
r = run_solo(MODEL_MJ, MJ_PROMPTS, args.rounds)
results = {"mj": r, "total_elapsed": round(time.time() - t_global, 1)}
elif args.mode == "solo-pnj":
r = run_solo(MODEL_PNJ, PNJ_PROMPTS, args.rounds)
results = {"pnj": r, "total_elapsed": round(time.time() - t_global, 1)}
print_report(results, args.mode)
if __name__ == "__main__":
main()