forked from Dan4ick/Lidar_Muxa
99 lines
4.4 KiB
Python
99 lines
4.4 KiB
Python
"""Сквозной прогон конвейера по бэгу: решения, треки, тайминги.
|
|
|
|
python tools/run_pipeline.py --bag data/for_hackathon/doubleT_obstacle --verbose
|
|
python tools/run_pipeline.py --all --memory artifacts/mushroom_body.npz
|
|
"""
|
|
from __future__ import annotations
|
|
|
|
import argparse
|
|
|
|
import numpy as np
|
|
|
|
import _bootstrap as B # noqa: F401
|
|
from flyguard.bag import Bag, find_bags
|
|
from flyguard.mushroom_body import MushroomBody
|
|
from flyguard.mbon_readout import MbonReadout
|
|
from flyguard.pipeline import FlyGuard, Params
|
|
from flyguard.track_readout import TrackReadout
|
|
|
|
|
|
def run(bag_path, params: Params, memory, readout=None, track_readout=None,
|
|
limit: int | None = None, verbose: bool = False) -> dict:
|
|
bag = Bag(bag_path)
|
|
fg = FlyGuard(params, memory=memory, readout=readout, track_readout=track_readout)
|
|
stages: dict[str, list[float]] = {}
|
|
n_det = n_frames = 0
|
|
dists, speeds, ncand = [], [], []
|
|
|
|
for _, pc in bag.frames(stop=limit):
|
|
res = fg.process(pc)
|
|
if res is None:
|
|
continue
|
|
n_frames += 1
|
|
for k, v in res.timings.items():
|
|
stages.setdefault(k, []).append(v)
|
|
ncand.append(len(res.candidates))
|
|
speeds.append(res.ego.speed if res.ego else 0.0)
|
|
d = res.decision
|
|
if d.detected:
|
|
n_det += 1
|
|
dists.append(d.distance)
|
|
if verbose and n_frames % 10 == 0:
|
|
obj = (f"{d.distance:6.1f} м conf={d.confidence:.2f} "
|
|
f"{'ЭКСТРЕННО' if d.emergency else 'предупр.'}"
|
|
if d.detected else "путь свободен")
|
|
print(f" кадр {n_frames:4d} v={res.ego.speed*3.6:5.1f} км/ч "
|
|
f"канд.={len(res.candidates):3d} треков={len(fg.cx.tracks):3d} | {obj} "
|
|
f"| {res.total_ms:5.1f} мс")
|
|
|
|
tot = np.array([sum(v[i] for v in stages.values()) for i in range(n_frames)]) \
|
|
if n_frames else np.zeros(1)
|
|
res = dict(name=bag.path.name, frames=n_frames, det_rate=n_det / max(n_frames, 1),
|
|
p50=float(np.median(tot)), p95=float(np.percentile(tot, 95)),
|
|
cands=float(np.mean(ncand)) if ncand else 0.0,
|
|
v=float(np.median(speeds) * 3.6) if speeds else 0.0,
|
|
d_med=float(np.median(dists)) if dists else float("nan"))
|
|
print(f"{res['name']:40s} кадров {res['frames']:4d} | тревога в {res['det_rate']:6.1%} "
|
|
f"кадров (медиана {res['d_med']:6.1f} м) | канд./кадр {res['cands']:5.1f} | "
|
|
f"v={res['v']:5.1f} км/ч | {res['p50']:5.1f}/{res['p95']:5.1f} мс (p50/p95)")
|
|
if verbose:
|
|
for k in sorted(stages, key=lambda k: -np.median(stages[k])):
|
|
print(f" {k:12s} {np.median(stages[k]):6.2f} мс")
|
|
return res
|
|
|
|
|
|
def main() -> None:
|
|
ap = argparse.ArgumentParser(description=__doc__)
|
|
ap.add_argument("--bag")
|
|
ap.add_argument("--all", action="store_true")
|
|
ap.add_argument("--limit", type=int, default=150)
|
|
ap.add_argument("--memory", default=None)
|
|
ap.add_argument("--readout", default=None, help="модель MBON (mbon_readout.npz)")
|
|
ap.add_argument("--track-readout", default=None, help="модель TrackReadout (track_readout.npz)")
|
|
ap.add_argument("--device", default="auto", choices=["auto", "cuda", "cpu"],
|
|
help="устройство вычислений ('auto', 'cuda', 'cpu')")
|
|
ap.add_argument("--fov", type=float, default=30.0)
|
|
ap.add_argument("--verbose", action="store_true")
|
|
args = ap.parse_args()
|
|
|
|
memory = MushroomBody.load(args.memory) if args.memory else None
|
|
readout = MbonReadout.load(args.readout) if args.readout else None
|
|
track_readout = TrackReadout.load(args.track_readout) if args.track_readout else None
|
|
|
|
if not args.all and not args.bag:
|
|
args.all = True
|
|
|
|
params = Params(fov_deg=args.fov, device=args.device)
|
|
bag_root = (B.DATA / "for_hackathon") if (B.DATA / "for_hackathon").exists() else B.DATA
|
|
bags = find_bags(bag_root) if args.all else ([args.bag] if args.bag else [])
|
|
if not bags:
|
|
print(f"Внимание: бэги не найдены в {bag_root}. Убедитесь, что каталог смонтирован в FLYGUARD_DATA.")
|
|
return
|
|
|
|
for b in bags:
|
|
run(b, params, memory, readout=readout, track_readout=track_readout,
|
|
limit=args.limit, verbose=args.verbose)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
main()
|