"""Обучение памяти тоннеля — грибовидного тела. Учится **без единой метки**: конвейер прогоняется по проездам пустого тоннеля, все выданные геометрией кандидаты объявляются «знакомой обстановкой», и синапсы KC→MBON на них депрессируются. После этого лотки, ниши, гермозатворы, кромки платформ и стрелочные приводы перестают быть новостью, а незнакомая форма — нет. python tools/train_mushroom_body.py --out artifacts/mushroom_body.npz python tools/train_mushroom_body.py --exclude roundT_doubleT --device cuda Датасет кандидатов кэшируется, поэтому подбор параметров памяти не требует повторного прогона конвейера по 100 ГБ данных. """ from __future__ import annotations import argparse import time import numpy as np import _bootstrap as B # noqa: F401 from flyguard.bag import Bag, find_bags from flyguard.mushroom_body import MushroomBody, MushroomBodyConfig, describe from flyguard.pipeline import FlyGuard, Params # бэг с реальным препятствием в обучение не идёт: память обязана считать его новым HOLDOUT = {"doubleT_obstacle"} def collect(bag_path, params: Params, limit: int | None, stride: int): bag = Bag(bag_path) fg = FlyGuard(params, memory=None) rows, meta = [], [] for k, (_, pc) in enumerate(bag.frames(stop=limit, stride=stride)): res = fg.process(pc) if res is None: continue for c in res.candidates: rows.append(describe(c)) meta.append((c.d, c.u, c.h, c.n_rays)) return rows, meta def main() -> None: ap = argparse.ArgumentParser(description=__doc__) ap.add_argument("--root", default=str(B.DATA / "for_hackathon")) ap.add_argument("--extra", action="append", default=[], help="дополнительные бэги (например, data/new_data)") ap.add_argument("--out", default=str(B.ARTIFACTS / "mushroom_body.npz")) ap.add_argument("--cache", default=str(B.CACHE / "candidates.npz")) ap.add_argument("--exclude", action="append", default=[]) ap.add_argument("--extra-cache", action="append", default=[], help="готовые наборы дескрипторов (npz с массивом X), " "например data/cache/new_data_candidates.npz") ap.add_argument("--limit", type=int, default=None) ap.add_argument("--stride", type=int, default=1) _d = MushroomBodyConfig() # умолчания берутся из самой модели ap.add_argument("--rate", type=float, default=0.0, help="темп депрессии; 0 — подобрать по размеру выборки") ap.add_argument("--target", type=float, default=0.4, help="во сколько e-раз ослабляется типичная клетка Кеньона; " "0.4 даёт лучшее разделение (tools/tune_memory.py)") ap.add_argument("--n-kc", type=int, default=_d.n_kc) ap.add_argument("--claws", type=int, default=_d.claws) ap.add_argument("--sparsity", type=float, default=_d.sparsity) ap.add_argument("--device", default="cpu") ap.add_argument("--reuse-cache", action="store_true") args = ap.parse_args() B.ARTIFACTS.mkdir(parents=True, exist_ok=True) B.CACHE.mkdir(parents=True, exist_ok=True) params = Params() skip = HOLDOUT | set(args.exclude) bags = [b for b in find_bags(args.root) if b.name not in skip] bags += [__import__("pathlib").Path(p) for p in args.extra] if args.reuse_cache and __import__("pathlib").Path(args.cache).exists(): d = np.load(args.cache, allow_pickle=True) X = d["X"] names = list(d["names"]) print(f"кэш: {X.shape[0]} кандидатов из {len(names)} бэгов") else: all_rows, names, per_bag = [], [], [] for b in bags: t0 = time.time() rows, _ = collect(b, params, args.limit, args.stride) all_rows.extend(rows) names.append(b.name) per_bag.append(len(rows)) print(f" {b.name:42s} кандидатов {len(rows):7d} за {time.time()-t0:6.1f} с") if not all_rows: raise SystemExit("кандидатов не собрано — нечему учиться") X = np.stack(all_rows).astype(np.float32) np.savez_compressed(args.cache, X=X, names=np.array(names), per_bag=np.array(per_bag)) print(f"кэш сохранён: {args.cache}") for path in args.extra_cache: d = np.load(path, allow_pickle=True) extra = d["X"].astype(np.float32) if extra.shape[1] != X.shape[1]: raise SystemExit(f"{path}: {extra.shape[1]} признаков вместо {X.shape[1]} — " "набор собран другой версией дескриптора, пересоберите") print(f" + {path}: {extra.shape[0]} кандидатов") X = np.concatenate([X, extra]) print(f"обучающая выборка: {X.shape[0]} кандидатов, {X.shape[1]} признаков") cfg = MushroomBodyConfig(n_kc=args.n_kc, claws=args.claws, sparsity=args.sparsity) mb = MushroomBody(cfg) mb.fit_normalizer(X) rate = args.rate if args.rate > 0 else mb.auto_rate(X.shape[0], args.target) t0 = time.time() mb.learn(X, rate=rate, device=args.device) print(f"обучено на {X.shape[0]} примерах за {time.time()-t0:.2f} с " f"(устройство {args.device}, темп депрессии {rate:.4f})") nov = mb.novelty(X) frac = (mb.w_mbon < 0.5).mean() print(f"клеток Кеньона с подавленным синапсом: {frac:.1%}") print("новизна обучающей выборки: " + " ".join(f"p{q}={np.percentile(nov, q):.3f}" for q in (5, 25, 50, 75, 95))) mb.save(args.out) print("сохранено:", args.out) if __name__ == "__main__": main()