Brainrot_Muxa/tools/build_brain_atlas.py

177 lines
7.8 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.

"""Атлас нейронов FlyWire: соматические координаты + привязка к стадиям FlyGuard.
Берёт публичные выгрузки коннектома FAFB v783 (Codex) и сводит их в один
компактный файл: положение каждого нейрона во фронтальной проекции и номер
стадии конвейера, которой он соответствует. Дальше вид мозга рисуется как
облако из 139 тысяч точек, подсвеченное живой активностью — без какой-либо
симуляции: активность берётся из наших же стадий, а коннектом даёт только
анатомию и принадлежность клеток.
Привязка не выдумана: FlyGuard с самого начала собран из конкретных типов
клеток, и все они есть в выгрузке поимённо — LC11 (127 нейронов), LPLC2 (210),
HS (6), VS (16), T4 (6243), T5 (6002), клетки Кеньона (5177), MBON (96),
APL (2), гигантское волокно DNp01 (2).
Входные файлы (скачиваются сами, около 9 МБ):
coordinates.csv.gz root_id → положение сомы
classification.csv.gz super_class / class / sub_class / сторона
neurons.csv.gz нейропиль (group)
consolidated_cell_types.csv.gz имя типа клетки
python tools/build_brain_atlas.py
Данные FlyWire распространяются под CC-BY 4.0 (Dorkenwald et al., Nature 2024;
Schlegel et al., Nature 2024).
"""
from __future__ import annotations
import argparse
import csv
import gzip
import re
import urllib.request
from pathlib import Path
import numpy as np
import _bootstrap as B # noqa: F401
BASE = "https://storage.googleapis.com/flywire-data/codex/data/fafb/783"
FILES = ("coordinates", "classification", "neurons", "consolidated_cell_types")
# Порядок важен: стадии проверяются сверху вниз, первая подошедшая выигрывает.
# Именованный тип клетки сильнее нейропиля, нейропиль сильнее общего класса.
STAGES = (
"retina", # 0 R1–R8, омматидиальная решётка
"lamina", # 1 L1/L2, ON/OFF и центр-окружение
"medulla", # 2 T4/T5, элементарные детекторы движения
"lptc", # 3 HS/VS, широкопольный поток → собственная скорость
"looming", # 4 LPLC2, надвигание
"lobula", # 5 LC11, мелкий объект
"mushroom", # 6 KC / APL / MBON, новизна
"central", # 7 PB/FB/EB/NO, накопление улик
"descending", # 8 нисходящие, решение
"other", # 9 остальной мозг — контекст, рисуется тускло
)
OTHER = len(STAGES) - 1
CX_NEUROPILS = {"PB", "FB", "EB", "NO", "AB"}
def _download(dest: Path) -> None:
dest.mkdir(parents=True, exist_ok=True)
for name in FILES:
p = dest / f"{name}.csv.gz"
if p.exists() and p.stat().st_size > 1000:
continue
print(f" качаю {name}.csv.gz …", flush=True)
urllib.request.urlretrieve(f"{BASE}/{name}.csv.gz", p)
def _read(dest: Path, name: str):
with gzip.open(dest / f"{name}.csv.gz", "rt", encoding="utf-8") as f:
yield from csv.DictReader(f)
def classify(cell_type: str, group: str, sub_class: str, super_class: str) -> int:
"""Номер стадии FlyGuard для одного нейрона."""
t = cell_type or ""
if t in ("DNp01", "DNp02", "DNp11") or super_class == "descending":
return STAGES.index("descending")
if t == "LC11":
return STAGES.index("lobula")
if t == "LPLC2":
return STAGES.index("looming")
if re.fullmatch(r"HS[ENS]|VS\d+", t):
return STAGES.index("lptc")
if re.fullmatch(r"T[45][a-d]", t):
return STAGES.index("medulla")
if t in ("L1", "L2") or sub_class == "lamina_monopolar":
return STAGES.index("lamina")
if sub_class == "photo_receptor":
return STAGES.index("retina")
if t.startswith(("KC", "MBON")) or t == "APL":
return STAGES.index("mushroom")
parts = set(group.split(".")) if group else set()
if any(p.startswith("MB_") for p in parts):
return STAGES.index("mushroom")
if parts & CX_NEUROPILS:
return STAGES.index("central")
if "LA" in parts:
return STAGES.index("lamina")
# Остальная оптическая доля: медулла отвечает за движение, лобула — за форму.
if "LOP" in parts:
return STAGES.index("lptc")
if "LO" in parts:
return STAGES.index("lobula")
if "ME" in parts:
return STAGES.index("medulla")
return OTHER
def main() -> None:
ap = argparse.ArgumentParser(description=__doc__)
ap.add_argument("--src", default=str(B.DATA / "flywire"))
ap.add_argument("--out", default=str(B.ROOT / "ros2_ws" / "src" / "flyguard" /
"flyguard" / "data" / "brain_atlas.npz"))
ap.add_argument("--width", type=int, default=1180)
ap.add_argument("--height", type=int, default=620)
ap.add_argument("--margin", type=int, default=18)
args = ap.parse_args()
src = Path(args.src)
_download(src)
pos: dict[str, tuple[int, int, int]] = {}
for r in _read(src, "coordinates"):
rid = r["root_id"]
if rid in pos:
continue # берём первую точку на нейрон
x, y, z = (int(v) for v in r["position"].strip("[]").split())
pos[rid] = (x, y, z)
group = {r["root_id"]: r["group"] for r in _read(src, "neurons")}
ctype = {r["root_id"]: r["primary_type"]
for r in _read(src, "consolidated_cell_types")}
rid_list, stage_list, side_list = [], [], []
for r in _read(src, "classification"):
rid = r["root_id"]
if rid not in pos:
continue
rid_list.append(rid)
stage_list.append(classify(ctype.get(rid, ""), group.get(rid, ""),
r["sub_class"], r["super_class"]))
side_list.append({"left": 0, "right": 1}.get(r["side"], 2))
xyz = np.array([pos[r] for r in rid_list], np.float64)
stage = np.array(stage_list, np.int8)
side = np.array(side_list, np.int8)
# Фронтальная проекция: x — влево-вправо, y — вверх-вниз, z — вглубь.
# Масштаб общий по обеим осям, иначе мозг растянется.
W, H, m = args.width, args.height, args.margin
lo, hi = xyz[:, :2].min(0), xyz[:, :2].max(0)
k = min((W - 2 * m) / (hi[0] - lo[0]), (H - 2 * m) / (hi[1] - lo[1]))
px = np.rint((xyz[:, 0] - lo[0]) * k + (W - (hi[0] - lo[0]) * k) / 2)
py = np.rint((xyz[:, 1] - lo[1]) * k + (H - (hi[1] - lo[1]) * k) / 2)
z = xyz[:, 2]
depth = np.rint(255 * (z - z.min()) / max(np.ptp(z), 1.0))
np.savez_compressed(
args.out,
px=px.astype(np.int16), py=py.astype(np.int16),
stage=stage, side=side, depth=depth.astype(np.uint8),
stages=np.array(STAGES), shape=np.array([H, W], np.int32))
print(f"\nнейронов в атласе: {len(rid_list)} холст {W}×{H}")
for i, name in enumerate(STAGES):
n = int((stage == i).sum())
print(f" {i} {name:<11s} {n:6d}")
print(f"\nсохранено: {args.out} "
f"({Path(args.out).stat().st_size / 1024:.0f} КБ)")
if __name__ == "__main__":
main()