Files
creator/backend/train.py
2026-07-04 12:21:45 +02:00

199 lines
9.3 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""make train: Parameter-Optimierung auf Mini-Themen (Baseline → Screening → Koordinaten-Suche).
Jeder Trial ist ein Subprozess (train_lauf.py) mit CREATOR_PARAMS im ENV — so binden die
Module die überschriebenen Werte beim Import. Metriken sind deterministisch (qa_report
ohne LLM); gegen Judge-/Lauf-Rauschen gilt: Baseline mit Wiederholung liefert die
Rausch-Schwelle, und eine Übernahme braucht einen BESTÄTIGUNGSLAUF (sonst Random Walk).
CLI: python3 train.py [--trials 40] [--stunden 8] [--sitzung NAME]
Ergebnis: storage/train/<sitzung>/{trials.jsonl, report.md, beste_params.json}
"""
import argparse
import asyncio
import hashlib
import json
import sys
import time
from datetime import datetime, timezone
from pathlib import Path
from config import STORAGE_DIR
from train_params import PARAMS, schritte
HAUPT_THEMA = ("train-sort", "benchmarks/sortierverfahren")
VALIDIER_THEMA = ("train-foto", "benchmarks/fotografie")
# Score-Gewichte: Qualität + Auswahl dominieren (Entwicklungsphase), Kosten ziehen ab.
W_NOTE, W_AUSWAHL, W_ZEIT, W_TOKEN = 4.0, 4.0, 1.0, 1.0
def score(m: dict, basis: dict) -> float:
"""Skalarer Vergleichswert eines Trials. note 010; auswahl aus den MECE-Quoten;
Zeit/Tokens normiert auf die Baseline (1.0 = Baseline-Kosten)."""
q = m.get("quoten") or {}
qa_ = m.get("quoten_artefakte") or {}
auswahl = 10.0 * max(0.0, 1.0 - min(1.0, (
q.get("dubletten_verdacht", 0) + q.get("luecken", 0) + q.get("fremd", 0)
+ qa_.get("sub_dubletten_verdacht", 0) + qa_.get("verwaiste", 0))))
zeit = (m.get("dauer_min") or 0) / max(basis.get("dauer_min") or 1, 0.1)
tok = _tokens(m) / max(_tokens(basis), 1)
return round(W_NOTE * (m.get("note") or 0) + W_AUSWAHL * auswahl
- W_ZEIT * 10 * zeit - W_TOKEN * 10 * tok, 2)
def _tokens(m: dict) -> int:
t = m.get("tokens") or {}
return int(t.get("input") or 0) + int(t.get("output") or 0)
class Trainer:
def __init__(self, sitzung: Path, max_trials: int, max_stunden: float, runner=None):
self.dir = sitzung
self.dir.mkdir(parents=True, exist_ok=True)
self.cache_pfad = self.dir / "trials.jsonl"
self.cache: dict[str, dict] = {}
if self.cache_pfad.exists(): # Resume: bezahlte Trials nie wiederholen
for line in self.cache_pfad.read_text(encoding="utf-8").splitlines():
e = json.loads(line)
self.cache[e["key"]] = e["metrics"]
self.max_trials = max_trials
self.deadline = time.monotonic() + max_stunden * 3600
self.gezahlt = 0
self.runner = runner or self._subprozess
self.log = []
# ── Trial-Ausführung ────────────────────────────────────────────────────────────
def _key(self, params: dict, thema: tuple, tag: str = "") -> str:
raw = json.dumps({"p": params, "t": thema[0], "tag": tag}, sort_keys=True)
return hashlib.md5(raw.encode()).hexdigest()[:12]
async def trial(self, params: dict, thema: tuple = HAUPT_THEMA, tag: str = "") -> dict | None:
"""tag unterscheidet bewusste Wiederholungen (Baseline n=2, Bestätigung)."""
key = self._key(params, thema, tag)
if key in self.cache:
return self.cache[key]
if self.gezahlt >= self.max_trials or time.monotonic() > self.deadline:
return None
self.gezahlt += 1
metrics = await self.runner(params, thema)
if metrics is not None:
with open(self.cache_pfad, "a", encoding="utf-8") as f:
f.write(json.dumps({"key": key, "params": params, "thema": thema[0],
"tag": tag, "metrics": metrics}, ensure_ascii=False) + "\n")
self.cache[key] = metrics
return metrics
async def _subprozess(self, params: dict, thema: tuple) -> dict | None:
out = self.dir / f"metrics-{self._key(params, thema)}.json"
env = {"CREATOR_PARAMS": json.dumps(params)}
import os
proc = await asyncio.create_subprocess_exec(
sys.executable, "train_lauf.py", thema[0], thema[1], str(out),
env={**os.environ, **env})
rc = await proc.wait()
if rc != 0 or not out.exists():
self._log(f"Trial fehlgeschlagen (rc={rc}, params={params})")
return None
return json.loads(out.read_text(encoding="utf-8"))
def _log(self, msg: str) -> None:
line = f"{datetime.now(timezone.utc).isoformat()[11:19]} {msg}"
print(line, flush=True)
self.log.append(line)
# ── Trainings-Phasen ────────────────────────────────────────────────────────────
async def run(self) -> dict:
# Phase 0: Baseline zweimal → Score-Basis + Rausch-Schwelle
self._log("Baseline (2 Läufe)…")
b1 = await self.trial({}, tag="baseline-1")
b2 = await self.trial({}, tag="baseline-2")
if not b1 or not b2:
self._log("Baseline unvollständig — Abbruch.")
return {}
self.basis = b1
s1, s2 = score(b1, b1), score(b2, b1)
self.rauschen = max(abs(s1 - s2), 0.5) # Mindest-Schwelle gegen Glücks-Übernahmen
best_params: dict = {}
best_score = max(s1, s2)
self._log(f"Baseline-Score {s1}/{s2}, Rausch-Schwelle {self.rauschen}")
# Phase 1: Screening — je Parameter ±1 Schritt, Effekt vs. Rauschen
effekte: list[tuple[float, str, float]] = [] # (|effekt|, name, bester_wert)
for name in PARAMS:
lo, hi = schritte(name)
for wert in dict.fromkeys((lo, hi)): # lo==hi am Rand nur einmal
if wert == PARAMS[name]["default"]:
continue
m = await self.trial({**best_params, name: wert})
if m is None:
continue
delta = score(m, self.basis) - best_score
self._log(f"Screening {name}={wert}: Δ{delta:+.2f}")
if delta > self.rauschen:
effekte.append((delta, name, wert))
effekte.sort(reverse=True)
self._log(f"Wirksam: {[(n, w) for _, n, w in effekte]}")
# Phase 2: Koordinaten-Suche über ALLE wirksamen Parameter (keine feste Obergrenze),
# Übernahme nur nach Bestätigungslauf
for _, name, start_wert in effekte:
wert = start_wert
p = PARAMS[name]
richtung = p["step"] if wert > p["default"] else -p["step"]
while True:
kandidat = {**best_params, name: wert}
m = await self.trial(kandidat)
if m is None:
break
delta = score(m, self.basis) - best_score
if delta <= self.rauschen:
break
m2 = await self.trial(kandidat, tag="bestaetigung")
if m2 is None or score(m2, self.basis) - best_score <= self.rauschen:
self._log(f"{name}={wert}: nicht bestätigt — verworfen")
break
best_params, best_score = kandidat, min(score(m, self.basis), score(m2, self.basis))
self._log(f"ÜBERNOMMEN {name}={wert} → Score {best_score}")
naechster = round(wert + richtung, 4)
if not p["min"] <= naechster <= p["max"]:
break
wert = naechster
# Validierung auf dem zweiten Thema
if best_params:
v_base = await self.trial({}, thema=VALIDIER_THEMA, tag="val-base")
v_best = await self.trial(best_params, thema=VALIDIER_THEMA, tag="val-best")
if v_base and v_best:
self._log(f"Validierung {VALIDIER_THEMA[0]}: Baseline {score(v_base, v_base)}"
f" → Best {score(v_best, v_base)}")
self._schreibe_report(best_params, best_score)
return best_params
def _schreibe_report(self, best_params: dict, best_score: float) -> None:
from fsutil import atomic_write_json, atomic_write_text
atomic_write_json(self.dir / "beste_params.json", best_params, indent=1)
report = ["# Trainings-Report", "",
f"Trials bezahlt: {self.gezahlt}/{self.max_trials}",
f"Bester Score: {best_score} (Baseline-Rauschen {self.rauschen})",
f"Beste Parameter: `{json.dumps(best_params, ensure_ascii=False)}`",
"", "Nutzung: `CREATOR_PARAMS=$(cat beste_params.json) make dev` —",
"Übernahme nach config.py bleibt eine manuelle Entscheidung.", "", "## Log", ""]
report += [f"- {l}" for l in self.log]
atomic_write_text(self.dir / "report.md", "\n".join(report))
print(f"\nReport: {self.dir / 'report.md'}")
def main() -> None:
ap = argparse.ArgumentParser()
ap.add_argument("--trials", type=int, default=40)
ap.add_argument("--stunden", type=float, default=12.0)
ap.add_argument("--sitzung", default=datetime.now(timezone.utc).strftime("%Y%m%d-%H%M"))
args = ap.parse_args()
trainer = Trainer(STORAGE_DIR / "train" / args.sitzung, args.trials, args.stunden)
asyncio.run(trainer.run())
if __name__ == "__main__":
main()