Ретина, ламина, медулла, лобула, грибовидное тело, веерное тело, центральный комплекс, нисходящие нейроны. Обучение памяти тоннеля и считывания MBON, оценка leave-one-bag-out, полигон дальности, 24 теста. Реальный объект на 55 м — 98.9 % кадров, ложных 7.5 трека на км, кадр обрабатывается за 33 мс на CPU. Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
236 lines
13 KiB
Python
236 lines
13 KiB
Python
"""MBON с обучением с учителем — считывание, а не новая сеть.
|
||
|
||
Грибовидное тело в `mushroom_body.py` учится **без меток**: синапсы KC→MBON
|
||
депрессируются на всём, что тоннель показывает часто, и выход MBON означает
|
||
«незнакомо». Это ровно familiarity suppression MBON-α′3 и ровно то, что нужно,
|
||
когда меток нет.
|
||
|
||
Но у мухи та же схема умеет и другое. При обучении с подкреплением
|
||
дофаминергические нейроны PPL1/PAM депрессируют KC→MBON **избирательно** — те
|
||
клетки, что были активны вместе с наказанием, — и выход MBON начинает означать
|
||
не «незнакомо», а «это предвещает удар». Один и тот же нейропиль, один и тот же
|
||
разрежённый код, другой учитель.
|
||
|
||
Здесь сделано именно это. Слои не меняются:
|
||
|
||
* вход — те же признаки кандидата плюс опора веерного тела (23 «проекционных
|
||
нейрона»);
|
||
* KC — та же случайная разрежённая проекция по 6 «когтей» на клетку;
|
||
* APL — то же глобальное торможение «победитель забирает всё», но отклик
|
||
остаётся **градуальным** и делится на общую активность (дивизивная
|
||
нормировка), а не превращается в единицы и нули;
|
||
* MBON — один выход, веса которого обучены различать предмет и обстановку.
|
||
|
||
Метки берутся не из разметки (её нет), а из физики: `flyguard.synth` вставляет
|
||
предмет трассировкой лучей, и кандидат считается предметом, если его ядро
|
||
состоит из лучей, в которые предмет действительно записан
|
||
(`tools/make_training_set.py`).
|
||
|
||
Зачем это нужно поверх ручной формулы веса улики. В `central_complex._quality`
|
||
шесть множителей, придуманных руками, а в дескрипторе 23 признака: **тень и
|
||
интенсивность в вес улики не входят вообще**, хотя окклюзионная тень за
|
||
предметом на большой дальности во много раз крупнее самого предмета.
|
||
|
||
Стоимость в инференсе — одно умножение матрицы 23 × n_kc на кандидата, доли
|
||
миллисекунды на CPU. GPU нужен только на обучении, и то не обязателен.
|
||
"""
|
||
from __future__ import annotations
|
||
|
||
from dataclasses import dataclass
|
||
from pathlib import Path
|
||
|
||
import numpy as np
|
||
|
||
from .lobula import Candidate
|
||
from .mushroom_body import describe
|
||
|
||
|
||
@dataclass
|
||
class MbonConfig:
|
||
"""Параметры считывания."""
|
||
|
||
# Ёмкость выбирается замером (`tools/train_mbon.py --sweep-kc`), а не на
|
||
# глаз: отбор «победитель забирает всё» стоит дороже матмула, и в худшем
|
||
# кадре с 64 кандидатами 20 000 клеток это 7.3 мс, 8 000 — 2.8 мс, 4 000 —
|
||
# 1.3 мс. Умолчание взято средним: платить временем кадра за ёмкость имеет
|
||
# смысл, только если развёртка показала, что она что-то даёт.
|
||
n_kc: int = 8_000 # клеток Кеньона в этом контуре
|
||
claws: int = 6 # входов на клетку (из коннектома)
|
||
sparsity: float = 0.0125 # доля активных после торможения APL — 100 клеток
|
||
seed: int = 20260921
|
||
|
||
|
||
def describe_full(c: Candidate) -> np.ndarray:
|
||
"""Вектор кандидата для считывания: дескриптор памяти плюс опора накопителя."""
|
||
acc = float(c.extra.get("acc_support", 0.0)) if c.extra else 0.0
|
||
return np.append(describe(c), np.float32(acc)).astype(np.float32)
|
||
|
||
|
||
class MbonReadout:
|
||
"""Разрежённый код клеток Кеньона + обученный линейный выход MBON."""
|
||
|
||
def __init__(self, cfg: MbonConfig | None = None, n_pn: int = 23):
|
||
self.cfg = cfg or MbonConfig()
|
||
self.n_pn = int(n_pn)
|
||
self.n_active = max(1, int(round(self.cfg.n_kc * self.cfg.sparsity)))
|
||
rng = np.random.default_rng(self.cfg.seed)
|
||
idx = np.stack([rng.choice(self.n_pn, size=self.cfg.claws, replace=False)
|
||
for _ in range(self.cfg.n_kc)])
|
||
sign = rng.choice((-1.0, 1.0), size=idx.shape)
|
||
w = np.zeros((self.cfg.n_kc, self.n_pn), np.float32)
|
||
np.put_along_axis(w, idx, sign.astype(np.float32), axis=1)
|
||
self.W = w
|
||
self.mean = np.zeros(self.n_pn, np.float32)
|
||
self.scale = np.ones(self.n_pn, np.float32)
|
||
self.w_mbon = np.zeros(self.cfg.n_kc, np.float32)
|
||
self.bias = np.float32(0.0)
|
||
# калибровка выхода: sigmoid(gain·z + shift). Градиентный спуск по
|
||
# разрежённому коду хорошо упорядочивает кандидатов, но масштаб логита
|
||
# зависит от ёмкости и числа эпох, а нам нужна осмысленная вероятность —
|
||
# она идёт множителем в вес улики.
|
||
self.gain = np.float32(1.0)
|
||
self.shift = np.float32(0.0)
|
||
|
||
# ------------------------------------------------------------------ код
|
||
|
||
def fit_normalizer(self, X: np.ndarray) -> None:
|
||
X = np.atleast_2d(np.asarray(X, np.float32))
|
||
self.mean = X.mean(0).astype(np.float32)
|
||
s = X.std(0).astype(np.float32)
|
||
self.scale = np.where(s > 1e-6, s, 1.0).astype(np.float32)
|
||
|
||
def encode(self, X: np.ndarray, device: str | None = None):
|
||
"""Признаки → (индексы активных клеток, их нормированный отклик).
|
||
|
||
Отклик остаётся градуальным и делится на свою сумму: это дивизивная
|
||
нормировка APL. Двоичный код тут заметно хуже — он выбрасывает
|
||
«насколько» клетка возбуждена, а для решения это как раз важно.
|
||
|
||
Нормировка приводит средний отклик к единице, а не сумму: иначе каждая
|
||
из сотни активных клеток даёт вклад 0.01, логиты выходят микроскопические
|
||
и обучение сваливается в разумное ранжирование при бессмысленной
|
||
калибровке (замерено: AUC 0.97 при доле верных 0.37).
|
||
"""
|
||
X = np.atleast_2d(np.asarray(X, np.float32))
|
||
if X.shape[1] != self.n_pn:
|
||
raise ValueError(f"считывание обучено на {self.n_pn} признаках, "
|
||
f"а дескриптор даёт {X.shape[1]}")
|
||
k = self.n_active
|
||
chunk = max(1, int(2 ** 26 // max(self.cfg.n_kc, 1)))
|
||
|
||
if device and device != "cpu":
|
||
import torch
|
||
out_i = np.empty((X.shape[0], k), np.int64)
|
||
out_v = np.empty((X.shape[0], k), np.float32)
|
||
with torch.no_grad():
|
||
m = torch.as_tensor(self.mean, device=device)
|
||
s = torch.as_tensor(self.scale, device=device)
|
||
w = torch.as_tensor(self.W, device=device).T.contiguous()
|
||
for i in range(0, X.shape[0], chunk):
|
||
t = torch.as_tensor(X[i:i + chunk], device=device)
|
||
y = torch.relu(((t - m) / s) @ w)
|
||
v, a = torch.topk(y, k, dim=1)
|
||
v = v * (k / v.sum(1, keepdim=True).clamp_min(1e-6))
|
||
out_i[i:i + chunk] = a.cpu().numpy()
|
||
out_v[i:i + chunk] = v.cpu().numpy()
|
||
return out_i, out_v
|
||
|
||
out_i = np.empty((X.shape[0], k), np.int64)
|
||
out_v = np.empty((X.shape[0], k), np.float32)
|
||
for i in range(0, X.shape[0], chunk):
|
||
z = (X[i:i + chunk] - self.mean) / self.scale
|
||
y = np.maximum(z @ self.W.T, 0.0)
|
||
a = np.argpartition(-y, k - 1, axis=1)[:, :k]
|
||
v = np.take_along_axis(y, a, axis=1)
|
||
v *= k / np.maximum(v.sum(1, keepdims=True), 1e-6)
|
||
out_i[i:i + chunk] = a
|
||
out_v[i:i + chunk] = v
|
||
return out_i, out_v
|
||
|
||
# ------------------------------------------------------------------ выход
|
||
|
||
def score(self, X: np.ndarray) -> np.ndarray:
|
||
"""Вероятность «это посторонний предмет», 0…1."""
|
||
a, v = self.encode(X)
|
||
z = self.logit(X, code=(a, v))
|
||
return 1.0 / (1.0 + np.exp(-z))
|
||
|
||
def logit(self, X: np.ndarray, code=None) -> np.ndarray:
|
||
a, v = code if code is not None else self.encode(X)
|
||
z = self.bias + (self.w_mbon[a] * v).sum(axis=1)
|
||
return self.gain * z + self.shift
|
||
|
||
def score_of(self, c: Candidate) -> float:
|
||
return float(self.score(describe_full(c)[None, :])[0])
|
||
|
||
def annotate(self, cands: list[Candidate]) -> list[Candidate]:
|
||
if not cands:
|
||
return cands
|
||
X = np.stack([describe_full(c) for c in cands])
|
||
for c, p in zip(cands, self.score(X)):
|
||
c.extra["mbon"] = float(p)
|
||
return cands
|
||
|
||
# ------------------------------------------------------------------ обучение
|
||
|
||
def learn(self, X: np.ndarray, y: np.ndarray, *, epochs: int = 60,
|
||
lr: float = 4.0, l2: float = 1e-5, device: str | None = None,
|
||
verbose: bool = False) -> None:
|
||
"""Логистическая регрессия по разрежённому коду — депрессия с учителем.
|
||
|
||
Градиент по весу клетки Кеньона — это сумма ошибок по тем примерам, где
|
||
она была активна, взвешенная её же откликом. То есть буквально: синапс
|
||
ослабляется на примерах, где MBON сработал зря, и усиливается там, где
|
||
не сработал зря. У мухи это делает дофамин.
|
||
"""
|
||
y = np.asarray(y, np.float32)
|
||
a, v = self.encode(X, device=device)
|
||
n, k = a.shape
|
||
flat = a.ravel()
|
||
for ep in range(epochs):
|
||
z = self.bias + (self.w_mbon[a] * v).sum(axis=1)
|
||
p = 1.0 / (1.0 + np.exp(-z))
|
||
g = (p - y) / n
|
||
grad = np.bincount(flat, weights=np.repeat(g, k) * v.ravel(),
|
||
minlength=self.cfg.n_kc).astype(np.float32)
|
||
self.w_mbon -= lr * (grad + l2 * self.w_mbon)
|
||
self.bias -= np.float32(lr * g.sum())
|
||
if verbose and (ep + 1) % 50 == 0:
|
||
eps = 1e-7
|
||
loss = -(y * np.log(p + eps) + (1 - y) * np.log(1 - p + eps)).mean()
|
||
print(f" эпоха {ep + 1:4d}: логистическая потеря {loss:.4f}")
|
||
self._calibrate(self.bias + (self.w_mbon[a] * v).sum(axis=1), y)
|
||
|
||
def _calibrate(self, z: np.ndarray, y: np.ndarray, iters: int = 400) -> None:
|
||
"""Шкалирование Платта: подобрать наклон и сдвиг по обучающей выборке."""
|
||
g, sh = 1.0, 0.0
|
||
for _ in range(iters):
|
||
p = 1.0 / (1.0 + np.exp(-(g * z + sh)))
|
||
e = p - y
|
||
g -= 2.0 * float((e * z).mean()) / max(float((z * z).mean()), 1e-6)
|
||
sh -= 2.0 * float(e.mean())
|
||
self.gain, self.shift = np.float32(g), np.float32(sh)
|
||
|
||
# ------------------------------------------------------------------ хранение
|
||
|
||
def save(self, path: str | Path) -> None:
|
||
np.savez_compressed(path, w_mbon=self.w_mbon, bias=self.bias,
|
||
gain=self.gain, shift=self.shift,
|
||
mean=self.mean, scale=self.scale,
|
||
n_kc=self.cfg.n_kc, claws=self.cfg.claws,
|
||
sparsity=self.cfg.sparsity, seed=self.cfg.seed,
|
||
n_pn=self.n_pn)
|
||
|
||
@staticmethod
|
||
def load(path: str | Path) -> "MbonReadout":
|
||
d = np.load(path, allow_pickle=False)
|
||
cfg = MbonConfig(n_kc=int(d["n_kc"]), claws=int(d["claws"]),
|
||
sparsity=float(d["sparsity"]), seed=int(d["seed"]))
|
||
m = MbonReadout(cfg, n_pn=int(d["n_pn"]))
|
||
m.w_mbon = d["w_mbon"].astype(np.float32)
|
||
m.bias = np.float32(d["bias"])
|
||
m.gain = np.float32(d["gain"])
|
||
m.shift = np.float32(d["shift"])
|
||
m.mean = d["mean"].astype(np.float32)
|
||
m.scale = d["scale"].astype(np.float32)
|
||
return m
|