Brainrot_Muxa/tools/evaluate.py

196 lines
9.8 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.

"""Оценка обобщаемости: leave-one-bag-out.
Память тоннеля обучается на всех данных, **кроме** проверяемого бэга, и только
после этого конвейер прогоняется по нему. Иначе цифры лгут: подавлять
конструкции, которые сам же и запомнил, умеет кто угодно, а на приватном тесте
будет новый участок тоннеля.
Отчёт: ложные тревоги на километр пути и доля кадров с тревогой для пустых
бэгов; для `doubleT_obstacle` — ещё и доля кадров, в которых найден настоящий
объект на ~55 м.
python tools/evaluate.py --device cuda --out artifacts/generalisation.json
"""
from __future__ import annotations
import argparse
import json
import time
import numpy as np
import _bootstrap as B # noqa: F401
from flyguard.bag import Bag, find_bags
from flyguard.mushroom_body import MushroomBody, MushroomBodyConfig
from flyguard.pipeline import FlyGuard, Params
OBSTACLE_BAG = "doubleT_obstacle"
TRUE_D = (50.0, 62.0)
def train_excluding(per_bag: dict[str, np.ndarray], extra: np.ndarray | None,
exclude: str, target: float, device: str) -> MushroomBody:
parts = [v for k, v in per_bag.items() if k not in (exclude, OBSTACLE_BAG)]
if extra is not None:
parts.append(extra)
X = np.concatenate(parts).astype(np.float32)
mb = MushroomBody(MushroomBodyConfig())
mb.fit_normalizer(X)
mb.learn(X, rate=mb.auto_rate(X.shape[0], target), device=device)
return mb
def run_bag(bag_path, memory, limit: int, params: Params | None = None,
readout=None) -> dict:
fg = FlyGuard(params or Params(), memory=memory, readout=readout)
bag = Bag(bag_path)
n = alarms = obj_hits = fp_objects = 0
path_m = 0.0
fp_dists, times = [], []
# один и тот же лоток, попавший в треки, виден сотню кадров подряд; для
# эксплуатации важно не это, а сколько РАЗНЫХ ложных объектов возникло —
# именно столько раз поезд затормозил бы напрасно
fp_tracks: set[int] = set()
for _, pc in bag.frames(stop=limit):
res = fg.process(pc)
if res is None:
continue
n += 1
path_m += res.ego.ds if res.ego else 0.0
times.append(res.total_ms)
d = res.decision
is_obstacle_bag = bag.path.name == OBSTACLE_BAG
mine = [o for o in d.objects if TRUE_D[0] < o.distance < TRUE_D[1]] \
if is_obstacle_bag else []
others = [o for o in d.objects if o not in mine]
if mine:
obj_hits += 1
if others:
alarms += 1
fp_objects += len(others)
fp_tracks.update(o.track_id for o in others)
fp_dists.extend(o.distance for o in others)
km = max(path_m / 1000.0, 1e-6)
return dict(bag=bag.path.name, frames=n, path_m=path_m,
alarm_frames=alarms, alarm_rate=alarms / max(n, 1),
fp_objects=fp_objects, fp_tracks=len(fp_tracks),
fp_per_km=(len(fp_tracks) / km) if path_m > 5 else float("nan"),
fp_median_d=float(np.median(fp_dists)) if fp_dists else float("nan"),
obj_rate=obj_hits / max(n, 1) if bag.path.name == OBSTACLE_BAG else None,
ms_p50=float(np.median(times)) if times else 0.0,
ms_p95=float(np.percentile(times, 95)) if times else 0.0)
def main() -> None:
ap = argparse.ArgumentParser(description=__doc__)
ap.add_argument("--root", default=str(B.DATA / "for_hackathon"))
ap.add_argument("--cache", default=str(B.CACHE / "tune_candidates.npz"))
ap.add_argument("--extra-cache", default=str(B.CACHE / "new_data_candidates.npz"))
ap.add_argument("--limit", type=int, default=250)
ap.add_argument("--target", type=float, default=0.4)
ap.add_argument("--device", default="cpu")
ap.add_argument("--split-adv", type=float, default=None,
help="порог разделения фигуры и фона; 0 — выключить")
ap.add_argument("--split-gap", type=float, default=None,
help="порог разреза по контрасту ламины, м; 0 — выключить")
ap.add_argument("--split-near", type=float, default=None,
help="ближе этой дальности не резать, м")
ap.add_argument("--split-top", type=int, default=None,
help="сколько фигур выносить из одной компоненты; 0 — все")
ap.add_argument("--no-acc", action="store_true", help="выключить накопитель")
ap.add_argument("--no-hab", action="store_true", help="выключить привыкание")
ap.add_argument("--mbon", default="", help="путь к обученному считыванию MBON")
ap.add_argument("--mbon-dir", default="",
help="каталог с моделями по складкам (mbon_<бэг>.npz): "
"для каждого бэга берётся модель, его не видевшая")
ap.add_argument("--mbon-blend", type=float, default=None,
help="1 — только модель, 0 — только ручная формула")
ap.add_argument("--mbon-power", type=float, default=None,
help="резкость вероятности модели")
ap.add_argument("--warn", type=float, default=None,
help="порог улики для тревоги (по умолчанию 0.5)")
ap.add_argument("--clear", type=float, default=None,
help="нижний порог гистерезиса (по умолчанию 0.3)")
ap.add_argument("--min-hits", type=int, default=None,
help="наблюдений, без которых трек не считается")
ap.add_argument("--out", default=str(B.ARTIFACTS / "generalisation.json"))
args = ap.parse_args()
d = np.load(args.cache, allow_pickle=True)
per_bag = {str(k): d[f"X_{k}"].astype(np.float32) for k in d["names"]}
extra = None
from pathlib import Path
if Path(args.extra_cache).exists():
extra = np.load(args.extra_cache)["X"].astype(np.float32)
print(f"дополнительно в обучение: {extra.shape[0]} кандидатов из new_data")
over = {}
if args.split_adv is not None:
over["split_adv"] = args.split_adv
if args.split_gap is not None:
over["split_gap"] = args.split_gap
if args.split_near is not None:
over["split_near"] = args.split_near
if args.split_top is not None:
over["split_top"] = args.split_top
if args.no_acc:
over["enable_accumulator"] = False
if args.no_hab:
over["enable_habituation"] = False
if args.mbon_blend is not None:
over["mbon_blend"] = args.mbon_blend
if args.mbon_power is not None:
over["mbon_power"] = args.mbon_power
if args.warn is not None:
over["warn_evidence"] = args.warn
over["clear_evidence"] = args.warn * 0.6 if args.clear is None else args.clear
elif args.clear is not None:
over["clear_evidence"] = args.clear
if args.min_hits is not None:
over["min_hits"] = args.min_hits
params = Params(**over)
readout = None
folds = {}
if args.mbon_dir:
from pathlib import Path as _P
from flyguard.mbon_readout import MbonReadout
for f in _P(args.mbon_dir).glob("mbon_*.npz"):
folds[f.stem[len("mbon_"):]] = MbonReadout.load(f)
print(f"считывание MBON по складкам: {args.mbon_dir} "
f"({len(folds)} моделей)")
elif args.mbon:
from flyguard.mbon_readout import MbonReadout
readout = MbonReadout.load(args.mbon)
print(f"считывание MBON: {args.mbon}")
rows = []
for p in find_bags(args.root):
t0 = time.time()
mem = train_excluding(per_bag, extra, p.name, args.target, args.device)
rd = folds.get(p.name, readout) if folds else readout
if folds and p.name not in folds and p.name != OBSTACLE_BAG:
print(f" внимание: для {p.name} нет своей складки")
r = run_bag(p, mem, args.limit, params, readout=rd)
r["train_size"] = int(mem.n_seen)
rows.append(r)
obj = f"объект {r['obj_rate']:6.1%} | " if r["obj_rate"] is not None else ""
print(f"{r['bag']:40s} кадров {r['frames']:4d} путь {r['path_m']:6.0f} м | "
f"{obj}тревог {r['alarm_rate']:6.1%} | ложных треков {r['fp_tracks']:3d} "
f"({r['fp_per_km']:6.1f} на км) | {r['ms_p50']:5.1f}/{r['ms_p95']:5.1f} мс | "
f"{time.time()-t0:5.0f} с", flush=True)
B.ARTIFACTS.mkdir(parents=True, exist_ok=True)
with open(args.out, "w", encoding="utf-8") as f:
json.dump(rows, f, ensure_ascii=False, indent=1)
empty = [r for r in rows if r["bag"] != OBSTACLE_BAG]
print("\nсводка по пустым бэгам:")
print(f" доля кадров с ложной тревогой: {np.mean([r['alarm_rate'] for r in empty]):.2%}")
fp_km = [r["fp_per_km"] for r in empty if np.isfinite(r["fp_per_km"])]
if fp_km:
print(f" разных ложных треков на километр: {np.mean(fp_km):.1f} "
f"(медиана {np.median(fp_km):.1f})")
print(f" суммарный путь: {sum(r['path_m'] for r in empty):.0f} м")
print("сохранено:", args.out)
if __name__ == "__main__":
main()