Ретина, ламина, медулла, лобула, грибовидное тело, веерное тело, центральный комплекс, нисходящие нейроны. Обучение памяти тоннеля и считывания MBON, оценка leave-one-bag-out, полигон дальности, 24 теста. Реальный объект на 55 м — 98.9 % кадров, ложных 7.5 трека на км, кадр обрабатывается за 33 мс на CPU.
183 lines
9.1 KiB
Python
183 lines
9.1 KiB
Python
"""Оценка обобщаемости: 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("--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
|
||
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()
|