Brainrot_Muxa/tools/evaluate.py

235 lines
13 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("--warn-far", type=float, default=None,
help="порог тревоги на дальнем краю; ниже обычного — послабление далёким трекам")
ap.add_argument("--warn-far-from", type=float, default=None,
help="с какой дальности порог начинает падать, м")
ap.add_argument("--novelty-floor", type=float, default=None,
help="ниже этой новизны трек не считается")
ap.add_argument("--min-rays", type=int, default=None,
help="сколько лучей минимум образуют кандидата")
ap.add_argument("--leak-far", type=float, default=None,
help="утечка улики на дальнем краю; ниже обычной — далёкий трек прощает промахи")
ap.add_argument("--leak-far-from", type=float, default=None,
help="с какой дальности утечка начинает падать, м")
ap.add_argument("--nov-fade-from", type=float, default=None,
help="с какой дальности гасить вклад знакомости; 0 — не гасить")
ap.add_argument("--nov-fade-to", type=float, default=None,
help="к какой дальности вклад знакомости обнуляется")
ap.add_argument("--mbon-prior-from", type=float, default=None,
help="с какой дальности поправлять оценку модели на распространённость предметов; 0 — не поправлять")
ap.add_argument("--no-memory", action="store_true",
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
if args.warn_far is not None:
over["warn_far"] = args.warn_far
if args.warn_far_from is not None:
over["warn_far_from"] = args.warn_far_from
if args.novelty_floor is not None:
over["novelty_floor"] = args.novelty_floor
if args.min_rays is not None:
over["min_rays"] = args.min_rays
if args.leak_far is not None:
over["leak_far"] = args.leak_far
if args.leak_far_from is not None:
over["leak_far_from"] = args.leak_far_from
if args.nov_fade_from is not None:
over["nov_fade_from"] = args.nov_fade_from
if args.nov_fade_to is not None:
over["nov_fade_to"] = args.nov_fade_to
if args.mbon_prior_from is not None:
over["mbon_prior_from"] = args.mbon_prior_from
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 = (None if args.no_memory else
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) if mem is not None else 0
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()