Ретина, ламина, медулла, лобула, грибовидное тело, веерное тело, центральный комплекс, нисходящие нейроны. Обучение памяти тоннеля и считывания MBON, оценка leave-one-bag-out, полигон дальности, 24 теста. Реальный объект на 55 м — 98.9 % кадров, ложных 7.5 трека на км, кадр обрабатывается за 33 мс на CPU. Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
125 lines
5.5 KiB
Python
125 lines
5.5 KiB
Python
"""Подбор ёмкости грибовидного тела по разделяющей способности.
|
||
|
||
Память с малым числом клеток Кеньона насыщается: после нескольких тысяч примеров
|
||
подавлены все синапсы, и новым не выглядит уже ничто — включая настоящее
|
||
препятствие. Скрипт меряет, при каких параметрах память **различает**:
|
||
|
||
* отрицательные примеры — кандидаты отложенного пустого бэга (должны стать знакомы);
|
||
* положительные — реальный объект на ~55 м из `doubleT_obstacle` (должен остаться новым).
|
||
|
||
Считается ROC AUC по новизне. Заодно печатается доля подавленных синапсов —
|
||
прямой индикатор насыщения.
|
||
|
||
python tools/tune_memory.py --device cuda
|
||
"""
|
||
from __future__ import annotations
|
||
|
||
import argparse
|
||
import itertools
|
||
|
||
import numpy as np
|
||
|
||
import _bootstrap as B # noqa: F401
|
||
from flyguard.bag import Bag, find_bags
|
||
from flyguard.mushroom_body import FEATURES, MushroomBody, MushroomBodyConfig, describe
|
||
from flyguard.pipeline import FlyGuard, Params
|
||
|
||
OBSTACLE_BAG = "doubleT_obstacle"
|
||
TRUE_D = (50.0, 62.0)
|
||
|
||
|
||
def collect_bag(path, limit=None, stride=1, params=None):
|
||
fg = FlyGuard(params or Params(), memory=None)
|
||
rows, dists = [], []
|
||
for _, pc in Bag(path).frames(stop=limit, stride=stride):
|
||
res = fg.process(pc)
|
||
if res is None:
|
||
continue
|
||
for c in res.candidates:
|
||
rows.append(describe(c))
|
||
dists.append(c.d)
|
||
X = np.stack(rows).astype(np.float32) if rows else np.zeros((0, len(FEATURES)), np.float32)
|
||
return X, np.asarray(dists, np.float32)
|
||
|
||
|
||
def auc(pos: np.ndarray, neg: np.ndarray) -> float:
|
||
if pos.size == 0 or neg.size == 0:
|
||
return float("nan")
|
||
order = np.argsort(np.concatenate([pos, neg]))
|
||
ranks = np.empty(order.size, np.float64)
|
||
ranks[order] = np.arange(1, order.size + 1)
|
||
r_pos = ranks[:pos.size].sum()
|
||
return float((r_pos - pos.size * (pos.size + 1) / 2) / (pos.size * neg.size))
|
||
|
||
|
||
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("--device", default="cpu")
|
||
ap.add_argument("--limit", type=int, default=None)
|
||
ap.add_argument("--stride", type=int, default=1)
|
||
ap.add_argument("--split-near", type=float, default=None,
|
||
help="переопределить порог ближней зоны разреза")
|
||
ap.add_argument("--collect-only", action="store_true",
|
||
help="только пересобрать кэш дескрипторов и выйти")
|
||
args = ap.parse_args()
|
||
|
||
B.CACHE.mkdir(parents=True, exist_ok=True)
|
||
from pathlib import Path
|
||
cache = Path(args.cache)
|
||
|
||
if cache.exists():
|
||
d = np.load(cache, allow_pickle=True)
|
||
bags = {str(k): d[f"X_{k}"] for k in d["names"]}
|
||
obj_d = d["obj_d"]
|
||
else:
|
||
bags, obj_d = {}, None
|
||
for p in find_bags(args.root):
|
||
over = ({} if args.split_near is None
|
||
else {'split_near': args.split_near})
|
||
X, dd = collect_bag(p, args.limit, args.stride, Params(**over))
|
||
bags[p.name] = X
|
||
if p.name == OBSTACLE_BAG:
|
||
obj_d = dd
|
||
print(f" {p.name:42s} {X.shape[0]:7d} кандидатов")
|
||
np.savez_compressed(cache, names=np.array(list(bags)), obj_d=obj_d,
|
||
**{f"X_{k}": v for k, v in bags.items()})
|
||
print("кэш сохранён:", cache)
|
||
if args.collect_only:
|
||
return
|
||
|
||
X_obs = bags[OBSTACLE_BAG]
|
||
in_band = (obj_d > TRUE_D[0]) & (obj_d < TRUE_D[1])
|
||
X_pos = X_obs[in_band]
|
||
empty_names = [k for k in bags if k != OBSTACLE_BAG]
|
||
print(f"\nположительных (объект ~55 м): {X_pos.shape[0]}, "
|
||
f"пустых бэгов: {len(empty_names)}")
|
||
|
||
grid = itertools.product([2_000, 20_000, 100_000], [0.05, 0.01, 0.002], [0.1, 0.3, 0.7])
|
||
print(f"\n{'n_kc':>8} {'разреж.':>8} {'rate':>6} | {'AUC':>6} | "
|
||
f"{'нов.объект':>10} {'нов.фон':>8} | {'подавл.':>8}")
|
||
print("-" * 72)
|
||
best = None
|
||
for n_kc, sp, rate in grid:
|
||
aucs, novp, novn, sat = [], [], [], []
|
||
for held in empty_names:
|
||
train = np.concatenate([bags[k] for k in empty_names if k != held])
|
||
mb = MushroomBody(MushroomBodyConfig(n_kc=n_kc, sparsity=sp))
|
||
mb.fit_normalizer(train)
|
||
mb.learn(train, rate=rate, device=args.device)
|
||
p = mb.novelty(X_pos)
|
||
n = mb.novelty(bags[held])
|
||
aucs.append(auc(p, n)); novp.append(np.median(p)); novn.append(np.median(n))
|
||
sat.append(float((mb.w_mbon < 0.5).mean()))
|
||
a = float(np.mean(aucs))
|
||
print(f"{n_kc:8d} {sp:8.3f} {rate:6.2f} | {a:6.3f} | "
|
||
f"{np.mean(novp):10.3f} {np.mean(novn):8.3f} | {np.mean(sat):8.1%}")
|
||
if best is None or a > best[0]:
|
||
best = (a, n_kc, sp, rate)
|
||
|
||
print(f"\nлучшее: AUC={best[0]:.3f} при n_kc={best[1]}, разрежённость={best[2]}, rate={best[3]}")
|
||
|
||
|
||
if __name__ == "__main__":
|
||
main()
|