327 lines
13 KiB
Python
327 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 = 180) -> 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
|
|
report_path = os.getenv("CRASH_RESULTS_PATH", "/home/ubuntu/fallout-venice/src/config/crash_results.json")
|
|
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()
|