"""Подбор ёмкости грибовидного тела по разделяющей способности. Память с малым числом клеток Кеньона насыщается: после нескольких тысяч примеров подавлены все синапсы, и новым не выглядит уже ничто — включая настоящее препятствие. Скрипт меряет, при каких параметрах память **различает**: * отрицательные примеры — кандидаты отложенного пустого бэга (должны стать знакомы); * положительные — реальный объект на ~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()