Lidar_Muxa/tools/check_obstacle.py

104 lines
4.7 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.

"""Опорный тест: реальный объект на ~55 м в бэге `doubleT_obstacle`.
Печатает, в скольких кадрах он попал в кандидаты и в подтверждённые треки,
и сколько при этом было посторонних тревог.
python tools/check_obstacle.py [--memory artifacts/mushroom_body.npz] [--mbon artifacts/mbon_folds/mbon_roundT_doubleT.npz]
"""
from __future__ import annotations
import argparse
import numpy as np
import _bootstrap as B # noqa: F401
import _params as PS
from _metrics import auc
from flyguard.bag import Bag
from flyguard.mushroom_body import MushroomBody
from flyguard.pipeline import FlyGuard, Params
TRUE_D = (50.0, 62.0) # объект стоит на 54.7…56.9 м всю запись
def main() -> None:
ap = argparse.ArgumentParser(description=__doc__)
ap.add_argument("--bag", default=str(B.DATA / "for_hackathon" / "doubleT_obstacle"))
ap.add_argument("--memory")
ap.add_argument("--mbon", default="",
help="обученное считывание; для честной цифры берите складку, "
"не видевшую этот бэг: artifacts/mbon_folds/mbon_roundT_doubleT.npz")
ap.add_argument("--limit", type=int, default=200)
ap.add_argument("--verbose", action="store_true")
PS.add_argument(ap)
ap.add_argument("--tracks", action="store_true",
help="сравнить признаки трека у настоящего объекта и у "
"ложных треков: единственная проверка, где метка "
"не синтетическая")
args = ap.parse_args()
memory = MushroomBody.load(args.memory) if args.memory else None
readout = None
if args.mbon:
from flyguard.mbon_readout import MbonReadout
readout = MbonReadout.load(args.mbon)
print(f"считывание MBON: {args.mbon}")
fg = FlyGuard(Params(**PS.apply({}, args.set, Params)), memory=memory,
readout=readout)
bag = Bag(args.bag)
n = cand_hit = track_hit = other = 0
novelties, evid = [], []
trk_X, trk_y = [], []
for k, (_, pc) in enumerate(bag.frames(stop=args.limit)):
res = fg.process(pc)
if res is None:
continue
n += 1
cs = [c for c in res.candidates if TRUE_D[0] < c.d < TRUE_D[1]]
if cs:
cand_hit += 1
novelties.append(max(c.novelty for c in cs))
objs = res.decision.objects
mine = [o for o in objs if TRUE_D[0] < o.distance < TRUE_D[1]]
if mine:
track_hit += 1
evid.append(max(o.confidence for o in mine))
other += len(objs) - len(mine)
if args.tracks:
from flyguard.track_readout import describe_track
cx = fg.cx
for tr in cx.tracks:
dd = tr.distance(cx.s_world)
if not (0.0 < dd <= 220.0) or tr.hits < 2:
continue
trk_X.append(describe_track(tr, cx.s_world))
trk_y.append(int(TRUE_D[0] < dd < TRUE_D[1]))
if args.verbose and k % 20 == 0:
print(f" кадр {k:4d}: канд. в зоне {len(cs)}, "
f"треков всего {len(objs)}, в зоне {len(mine)}, "
f"ближайший {res.decision.distance:.1f} м")
print(f"кадров обработано: {n}")
print(f"объект среди кандидатов: {cand_hit}/{n} = {cand_hit/max(n,1):.1%}"
+ (f" новизна медиана {np.median(novelties):.3f}" if novelties else ""))
print(f"объект подтверждён треком: {track_hit}/{n} = {track_hit/max(n,1):.1%}"
+ (f" уверенность медиана {np.median(evid):.3f}" if evid else ""))
print(f"посторонних подтверждённых объектов: {other} ({other/max(n,1):.2f} на кадр)")
if args.tracks and trk_y:
from flyguard.track_readout import TRACK_FEATURES
X, y = np.stack(trk_X), np.array(trk_y)
idx = {f: k for k, f in enumerate(TRACK_FEATURES)}
print()
print(f"наблюдений трека: {len(y)}, настоящего объекта {int(y.sum())}")
print(f"{'признак':<12}{'объект':>9}{'ложные':>9}{'AUC':>8}")
for f in ("evidence", "w_mean", "p_mean", "p_max", "rays_mean",
"s_std", "u_std", "age"):
v = X[:, idx[f]]
print(f"{f:<12}{np.median(v[y == 1]):9.3f}"
f"{np.median(v[y == 0]):9.3f}{auc(v, y):8.3f}")
if __name__ == "__main__":
main()