Ретина, ламина, медулла, лобула, грибовидное тело, веерное тело, центральный комплекс, нисходящие нейроны. Обучение памяти тоннеля и считывания MBON, оценка leave-one-bag-out, полигон дальности, 24 теста. Реальный объект на 55 м — 98.9 % кадров, ложных 7.5 трека на км, кадр обрабатывается за 33 мс на CPU. Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
108 lines
4.3 KiB
Python
108 lines
4.3 KiB
Python
"""Абляция: вклад каждого механизма, измеренный в одинаковых условиях.
|
||
|
||
Каждый вариант отличается от полного ровно одним отключённым механизмом.
|
||
Память тоннеля во всех вариантах обучается **без проверяемого бэга**
|
||
(leave-one-bag-out), иначе сравнение было бы нечестным.
|
||
|
||
Меряются две величины, которые и определяют полезность системы:
|
||
доля кадров с ложной тревогой на пустых проездах и доля кадров, в которых
|
||
найден реальный объект на ~55 м в `doubleT_obstacle`.
|
||
|
||
python tools/ablation.py --device cuda
|
||
"""
|
||
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.pipeline import FlyGuard, Params
|
||
from evaluate import OBSTACLE_BAG, TRUE_D, train_excluding
|
||
|
||
VARIANTS: dict[str, dict] = {
|
||
"полная система": {},
|
||
"− память тоннеля": {"enable_memory": False},
|
||
"− ось пути (прямой коридор)": {"use_corridor": False},
|
||
"− признаки формы": {"use_shape": False},
|
||
"− накопление улик": {"use_tracking": False},
|
||
}
|
||
|
||
|
||
def run(bag_path, params: Params, memory, limit: int) -> tuple[int, int, int, int]:
|
||
fg = FlyGuard(params, memory=memory)
|
||
n = alarm = obj = 0
|
||
fp_tracks: set[int] = set()
|
||
is_obs = Bag(bag_path).path.name == OBSTACLE_BAG
|
||
for _, pc in Bag(bag_path).frames(stop=limit):
|
||
res = fg.process(pc)
|
||
if res is None:
|
||
continue
|
||
n += 1
|
||
mine = [o for o in res.decision.objects
|
||
if is_obs and TRUE_D[0] < o.distance < TRUE_D[1]]
|
||
others = [o for o in res.decision.objects if o not in mine]
|
||
if mine:
|
||
obj += 1
|
||
if others:
|
||
alarm += 1
|
||
fp_tracks.update(o.track_id for o in others)
|
||
return n, alarm, obj, len(fp_tracks)
|
||
|
||
|
||
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=200)
|
||
ap.add_argument("--device", default="cpu")
|
||
ap.add_argument("--out", default=str(B.ARTIFACTS / "ablation.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"]}
|
||
from pathlib import Path
|
||
extra = (np.load(args.extra_cache)["X"].astype(np.float32)
|
||
if Path(args.extra_cache).exists() else None)
|
||
|
||
bags = find_bags(args.root)
|
||
memories = {p.name: train_excluding(per_bag, extra, p.name, 0.4, args.device)
|
||
for p in bags}
|
||
|
||
print(f"{'вариант':30s} | {'ложных кадров':>14s} {'ложн. треков':>13s} | "
|
||
f"{'объект 55 м':>12s}")
|
||
print("-" * 78)
|
||
rows = []
|
||
for label, over in VARIANTS.items():
|
||
t0 = time.time()
|
||
params = Params(**over)
|
||
tot_n = tot_alarm = tot_fp = 0
|
||
obj_n = obj_hit = 0
|
||
for p in bags:
|
||
mem = memories[p.name] if params.enable_memory else None
|
||
n, alarm, obj, fp = run(p, params, mem, args.limit)
|
||
if p.name == OBSTACLE_BAG:
|
||
obj_n, obj_hit = n, obj
|
||
else:
|
||
tot_n += n
|
||
tot_alarm += alarm
|
||
tot_fp += fp
|
||
row = dict(variant=label, alarm_rate=tot_alarm / max(tot_n, 1),
|
||
fp_tracks=tot_fp, obj_rate=obj_hit / max(obj_n, 1),
|
||
seconds=round(time.time() - t0, 1))
|
||
rows.append(row)
|
||
print(f"{label:30s} | {row['alarm_rate']:13.1%} {tot_fp:13d} | "
|
||
f"{row['obj_rate']:11.1%}", 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)
|
||
print("\nсохранено:", args.out)
|
||
|
||
|
||
if __name__ == "__main__":
|
||
main()
|