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

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

125 lines
5.5 KiB
Python
Raw 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` (должен остаться новым).
Считается 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()