Brainrot_Muxa/tools/train_mushroom_body.py

130 lines
6.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.

"""Обучение памяти тоннеля — грибовидного тела.
Учится **без единой метки**: конвейер прогоняется по проездам пустого тоннеля,
все выданные геометрией кандидаты объявляются «знакомой обстановкой», и синапсы
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="auto",
help="устройство вычислений ('auto', 'cuda', '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()