Lidar_Muxa/tools/ablation.py
Данил Омелечко b8e95eccbe ML-ядро детектора: конвейер на схемах мозга дрозофилы
Ретина, ламина, медулла, лобула, грибовидное тело, веерное тело,
центральный комплекс, нисходящие нейроны. Обучение памяти тоннеля и
считывания MBON, оценка leave-one-bag-out, полигон дальности, 24 теста.

Реальный объект на 55 м — 98.9 % кадров, ложных 7.5 трека на км,
кадр обрабатывается за 33 мс на CPU.
2026-09-21 17:24:25 +03:00

108 lines
4.3 KiB
Python
Raw Permalink 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), иначе сравнение было бы нечестным.
Меряются две величины, которые и определяют полезность системы:
доля кадров с ложной тревогой на пустых проездах и доля кадров, в которых
найден реальный объект на ~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()