flydnet 0.1.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- flydnet/__init__.py +24 -0
- flydnet/__main__.py +28 -0
- flydnet/circuit.py +234 -0
- flydnet/data.py +145 -0
- flydnet/datasets.py +93 -0
- flydnet/encoders.py +71 -0
- flydnet/layers.py +437 -0
- flydnet/plasticity.py +210 -0
- flydnet/readout.py +45 -0
- flydnet/visual.py +127 -0
- flydnet-0.1.0.dist-info/METADATA +477 -0
- flydnet-0.1.0.dist-info/RECORD +15 -0
- flydnet-0.1.0.dist-info/WHEEL +5 -0
- flydnet-0.1.0.dist-info/licenses/LICENSE +21 -0
- flydnet-0.1.0.dist-info/top_level.txt +1 -0
flydnet/__init__.py
ADDED
|
@@ -0,0 +1,24 @@
|
|
|
1
|
+
"""flydnet — 초파리 커넥톰을 배선으로 쓰는 스파이킹 신경망 층
|
|
2
|
+
|
|
3
|
+
import flydnet as fd
|
|
4
|
+
mb = fd.Circuit.from_flywire() # 버섯체: PN → KC → MBON (+APL)
|
|
5
|
+
enc = fd.RateEncoder(784, len(mb.groups["PN"])) # 텐서 → PN 발화율
|
|
6
|
+
layer = fd.ConnectomeLayer(mb, inputs="PN", outputs="KC")
|
|
7
|
+
feats = layer(enc(images)) # (B, n_KC) 발화율 → 리드아웃 학습
|
|
8
|
+
|
|
9
|
+
# 역전파 학습: 배선 고정, 연결 세기 학습 (대리 기울기)
|
|
10
|
+
layer = fd.ConnectomeLayer(mb, "PN", "KC", dt=0.5, t_ms=50, input_mode="regular", trainable=True)
|
|
11
|
+
model = torch.nn.Sequential(enc, layer, torch.nn.Linear(layer.n_out, 10))
|
|
12
|
+
"""
|
|
13
|
+
from . import data
|
|
14
|
+
from .data import data_dir, set_data_dir, download, data_status
|
|
15
|
+
from .circuit import Circuit, MUSHROOM_BODY
|
|
16
|
+
from .encoders import RateEncoder, GlomerularEncoder
|
|
17
|
+
from .layers import ConnectomeLayer, SpikeFn, DEFAULT_PARAMS
|
|
18
|
+
from .readout import extract, train_linear
|
|
19
|
+
from .plasticity import DopamineReadout, AssocReadout
|
|
20
|
+
from .datasets import synthetic_odors, door_odors, biconditional_mixtures
|
|
21
|
+
from .visual import (visual_circuit, column_map, drifting_grating, direction_offsets,
|
|
22
|
+
VISUAL_SYSTEM, PHOTORECEPTORS, COLUMNAR, LPTC, MOTION_PATHWAY)
|
|
23
|
+
|
|
24
|
+
__version__ = "0.1.0"
|
flydnet/__main__.py
ADDED
|
@@ -0,0 +1,28 @@
|
|
|
1
|
+
"""명령줄: python -m flydnet [status | download [flywire|door ...]]
|
|
2
|
+
|
|
3
|
+
python -m flydnet # 데이터 상태 (어디서 무엇을 찾았는지)
|
|
4
|
+
python -m flydnet download # FlyWire v783 + DoOR 데이터 받기 (없는 파일만, 약 130 MB)
|
|
5
|
+
python -m flydnet download flywire # 한 묶음만
|
|
6
|
+
"""
|
|
7
|
+
import sys
|
|
8
|
+
|
|
9
|
+
from . import __version__
|
|
10
|
+
from .data import data_status, download
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
def main(argv=None):
|
|
14
|
+
argv = list(sys.argv[1:] if argv is None else argv)
|
|
15
|
+
cmd = argv.pop(0) if argv else "status"
|
|
16
|
+
if cmd == "status":
|
|
17
|
+
print(f"flydnet {__version__}")
|
|
18
|
+
data_status()
|
|
19
|
+
elif cmd == "download":
|
|
20
|
+
download(argv or ("flywire", "door"))
|
|
21
|
+
else:
|
|
22
|
+
print(__doc__)
|
|
23
|
+
return 1
|
|
24
|
+
return 0
|
|
25
|
+
|
|
26
|
+
|
|
27
|
+
if __name__ == "__main__":
|
|
28
|
+
sys.exit(main())
|
flydnet/circuit.py
ADDED
|
@@ -0,0 +1,234 @@
|
|
|
1
|
+
"""FlyWire 커넥톰에서 뉴런 묶음(회로)을 잘라내 신경망 층의 '배선'으로 쓰는 Circuit"""
|
|
2
|
+
from __future__ import annotations
|
|
3
|
+
|
|
4
|
+
from pathlib import Path
|
|
5
|
+
|
|
6
|
+
import numpy as np
|
|
7
|
+
import pandas as pd
|
|
8
|
+
|
|
9
|
+
# 이름 → (주석 열, 값 또는 값 목록). 주석 열: super_class, cell_class, cell_sub_class, cell_type
|
|
10
|
+
# 버섯체 기본 회로: 투사 뉴런 → Kenyon 세포 → MBON, APL이 전체를 억제
|
|
11
|
+
MUSHROOM_BODY = {
|
|
12
|
+
"PN": ("cell_class", "ALPN"),
|
|
13
|
+
"KC": ("cell_class", "Kenyon_Cell"),
|
|
14
|
+
"APL": ("cell_type", "APL"),
|
|
15
|
+
"MBON": ("cell_class", "MBON"),
|
|
16
|
+
}
|
|
17
|
+
|
|
18
|
+
|
|
19
|
+
def _repair(pre: np.ndarray, post: np.ndarray, N: int, rng, max_iter: int = 200) -> np.ndarray:
|
|
20
|
+
"""섞은 뒤 생긴 중복 연결(같은 pre→post 두 번)과 자기 연결(pre == post)을 없앰.
|
|
21
|
+
문제 연결의 post를 무작위 다른 연결과 맞바꿈 → 뉴런별 연결 수는 그대로 유지.
|
|
22
|
+
(중복을 그냥 두면 합쳐져서 무작위 회로의 뉴런이 실제보다 적은 입력을 받게 됨)"""
|
|
23
|
+
post = post.copy()
|
|
24
|
+
for _ in range(max_iter):
|
|
25
|
+
bad = pd.Series(pre * N + post).duplicated().values | (pre == post)
|
|
26
|
+
if not bad.any():
|
|
27
|
+
return post
|
|
28
|
+
for i, j in zip(np.nonzero(bad)[0], rng.integers(0, len(post), bad.sum())):
|
|
29
|
+
post[i], post[j] = post[j], post[i] # 하나씩 (한꺼번에 하면 같은 j가 겹칠 때 값이 사라짐)
|
|
30
|
+
raise RuntimeError(f"중복 연결을 {max_iter}번 안에 없애지 못함 ({bad.sum()}개 남음)")
|
|
31
|
+
|
|
32
|
+
|
|
33
|
+
class Circuit:
|
|
34
|
+
"""뉴런 N개와 부호 있는 시냅스 목록 (pre → post, weight = ±시냅스 수)
|
|
35
|
+
|
|
36
|
+
groups: 이름 → 이 회로 안에서의 뉴런 번호 배열 (예: groups["KC"])
|
|
37
|
+
meta: 뉴런별 주석 (cell_type, cell_sub_class 등), 행 순서 = 회로 번호
|
|
38
|
+
pos: 뉴런별 위치 (N, 3) nm (FlyWire 주석의 대표 점), 없으면 None
|
|
39
|
+
"""
|
|
40
|
+
|
|
41
|
+
def __init__(self, root_ids, groups: dict[str, np.ndarray], pre, post, weight, name="circuit",
|
|
42
|
+
meta: pd.DataFrame | None = None, pos=None):
|
|
43
|
+
self.root_ids = np.asarray(root_ids, dtype=np.int64)
|
|
44
|
+
self.groups = {k: np.asarray(v, dtype=np.int64) for k, v in groups.items()}
|
|
45
|
+
self.pre = np.asarray(pre, dtype=np.int64)
|
|
46
|
+
self.post = np.asarray(post, dtype=np.int64)
|
|
47
|
+
self.weight = np.asarray(weight, dtype=np.float32)
|
|
48
|
+
self.name = name
|
|
49
|
+
self.meta = meta.reset_index(drop=True) if meta is not None else None
|
|
50
|
+
self.pos = np.asarray(pos, dtype=np.float32) if pos is not None else None
|
|
51
|
+
|
|
52
|
+
@property
|
|
53
|
+
def N(self) -> int:
|
|
54
|
+
return len(self.root_ids)
|
|
55
|
+
|
|
56
|
+
@property
|
|
57
|
+
def n_edges(self) -> int:
|
|
58
|
+
return len(self.pre)
|
|
59
|
+
|
|
60
|
+
def group_of(self) -> np.ndarray:
|
|
61
|
+
"""뉴런별 그룹 이름 (그룹이 겹치지 않는다고 가정)"""
|
|
62
|
+
g = np.empty(self.N, dtype=object)
|
|
63
|
+
for k, v in self.groups.items():
|
|
64
|
+
g[v] = k
|
|
65
|
+
return g
|
|
66
|
+
|
|
67
|
+
@classmethod
|
|
68
|
+
def from_flywire(cls, groups: dict = MUSHROOM_BODY, side: str | None = "right",
|
|
69
|
+
data_dir: str | Path | None = None, group_by: str | None = None, annotations: str = "flywire_annotations.tsv",
|
|
70
|
+
connectivity: str = "Connectivity_783.parquet",
|
|
71
|
+
completeness: str = "Completeness_783.csv") -> "Circuit":
|
|
72
|
+
"""주석으로 고른 뉴런들 사이의 연결만 남긴 회로 (induced subgraph)
|
|
73
|
+
group_by: 주석 열 이름 (예: "cell_type")이면 각 그룹을 그 값마다 다시 나눔 → 그룹 이름 = 값
|
|
74
|
+
(값이 없는 뉴런은 "<그룹 이름>?"). 무작위 대조군(shuffled)도 이 단위로 섞임
|
|
75
|
+
data_dir: None이면 flydnet.data_dir("flywire") (환경변수 → ~/.flydnet/config.json → ~/.flydnet/data)"""
|
|
76
|
+
from .data import require
|
|
77
|
+
d = require("flywire", data_dir)
|
|
78
|
+
all_ids = pd.read_csv(d / completeness, index_col=0).index.values.astype(np.int64)
|
|
79
|
+
ann = pd.read_csv(d / annotations, sep="\t", low_memory=False,
|
|
80
|
+
usecols=["root_id", "super_class", "cell_class", "cell_sub_class", "cell_type", "side",
|
|
81
|
+
"pos_x", "pos_y", "pos_z"])
|
|
82
|
+
ann = ann[ann.root_id.isin(all_ids)]
|
|
83
|
+
if side:
|
|
84
|
+
ann = ann[ann.side == side]
|
|
85
|
+
|
|
86
|
+
picked, gidx = [], {}
|
|
87
|
+
for name, (col, val) in groups.items():
|
|
88
|
+
vals = [val] if isinstance(val, str) else list(val)
|
|
89
|
+
sel = ann[ann[col].isin(vals)]
|
|
90
|
+
parts = sel.groupby(sel[group_by].fillna(f"{name}?"), sort=True) if group_by else [(name, sel)]
|
|
91
|
+
for sub, s in parts:
|
|
92
|
+
if sub in gidx:
|
|
93
|
+
raise ValueError(f"그룹 이름이 겹침: {sub}")
|
|
94
|
+
gidx[sub] = np.arange(len(picked), len(picked) + len(s))
|
|
95
|
+
picked.extend(s.root_id.values)
|
|
96
|
+
ids = np.array(picked, dtype=np.int64)
|
|
97
|
+
if len(np.unique(ids)) != len(ids):
|
|
98
|
+
raise ValueError("그룹끼리 뉴런이 겹침")
|
|
99
|
+
|
|
100
|
+
# 전체 뇌 번호 → 회로 번호
|
|
101
|
+
glob = pd.Series(np.arange(len(all_ids)), index=all_ids)[ids].values
|
|
102
|
+
local = np.full(len(all_ids), -1, np.int64); local[glob] = np.arange(len(ids))
|
|
103
|
+
|
|
104
|
+
df = pd.read_parquet(d / connectivity, columns=["Presynaptic_Index", "Postsynaptic_Index",
|
|
105
|
+
"Connectivity", "Excitatory"])
|
|
106
|
+
pre, post = local[df.Presynaptic_Index.values], local[df.Postsynaptic_Index.values]
|
|
107
|
+
keep = (pre >= 0) & (post >= 0)
|
|
108
|
+
w = (df.Connectivity.values * df.Excitatory.values)[keep]
|
|
109
|
+
a = ann.set_index("root_id").loc[ids]
|
|
110
|
+
meta = a[["super_class", "cell_class", "cell_sub_class", "cell_type"]].reset_index()
|
|
111
|
+
return cls(ids, gidx, pre[keep], post[keep], w, name=f"FlyWire {'/'.join(groups)} ({side or 'both'})",
|
|
112
|
+
meta=meta, pos=a[["pos_x", "pos_y", "pos_z"]].values)
|
|
113
|
+
|
|
114
|
+
def shuffled(self, seed: int = 0, pairs=None, exclude=None, local=None, merge=None) -> "Circuit":
|
|
115
|
+
"""무작위 배선 대조군: (보내는 그룹, 받는 그룹) 쌍마다 받는 뉴런을 섞음.
|
|
116
|
+
그룹 간 연결 수·시냅스 수·뉴런별 입출력 개수는 그대로, '누가 누구에게'만 무작위.
|
|
117
|
+
|
|
118
|
+
pairs: 섞을 쌍만 지정 (예: ["PN>KC"]). None이면 전부
|
|
119
|
+
exclude: 섞지 않을 쌍 (예: ["PN>KC", "KC>KC"])
|
|
120
|
+
local: (xy, radius) 이면 받는 뉴런 위치 xy (N, 2)를 radius 크기 칸으로 나눠 같은 칸 안에서만 섞음
|
|
121
|
+
→ 시야 위치 대응(retinotopy)은 유지, 그 안의 미세 배선(누가 정확히 누구에게)만 무작위.
|
|
122
|
+
위치가 NaN인 뉴런은 한 칸으로 묶음
|
|
123
|
+
merge: {그룹: 별칭} 섞을 때 같은 별칭 그룹들을 하나로 봄 (예: T4a~d → "T4").
|
|
124
|
+
받는 뉴런이 원래 다른 아형으로 가던 연결도 받게 됨 → 아형별 배선 차이(방향 구조)가 사라짐.
|
|
125
|
+
pairs/exclude는 별칭이 아닌 원래 그룹 이름 기준
|
|
126
|
+
"""
|
|
127
|
+
rng = np.random.default_rng(seed)
|
|
128
|
+
g = self.group_of()
|
|
129
|
+
post = self.post.copy()
|
|
130
|
+
key = pd.Series(g[self.pre] + ">" + g[self.post])
|
|
131
|
+
if pairs is not None and (unknown := set(pairs) - set(key.unique())):
|
|
132
|
+
raise ValueError(f"회로에 없는 연결 쌍: {sorted(unknown)}")
|
|
133
|
+
pair_of = key.values
|
|
134
|
+
if merge:
|
|
135
|
+
ga = np.array([merge.get(x, x) for x in g], dtype=object)
|
|
136
|
+
key = pd.Series(ga[self.pre] + ">" + ga[self.post])
|
|
137
|
+
if local is not None:
|
|
138
|
+
xy, radius = local
|
|
139
|
+
b = np.floor(np.asarray(xy)[self.post] / radius)
|
|
140
|
+
b = np.where(np.isnan(b).any(1, keepdims=True), np.inf, b)
|
|
141
|
+
key = key + "|" + pd.Series(b[:, 0]).astype(str) + "," + pd.Series(b[:, 1]).astype(str)
|
|
142
|
+
stuck = []
|
|
143
|
+
for k, idx in key.groupby(key).groups.items():
|
|
144
|
+
idx = np.asarray(idx)
|
|
145
|
+
if pairs is not None or exclude is not None:
|
|
146
|
+
pk = pair_of[idx]
|
|
147
|
+
sel = np.ones(len(idx), bool)
|
|
148
|
+
if pairs is not None:
|
|
149
|
+
sel &= np.isin(pk, list(pairs))
|
|
150
|
+
if exclude is not None:
|
|
151
|
+
sel &= ~np.isin(pk, list(exclude))
|
|
152
|
+
idx = idx[sel]
|
|
153
|
+
if len(idx) == 0:
|
|
154
|
+
continue
|
|
155
|
+
try:
|
|
156
|
+
post[idx] = _repair(self.pre[idx], post[idx][rng.permutation(len(idx))], self.N, rng)
|
|
157
|
+
except RuntimeError: # 뉴런 몇 개뿐인 쌍은 중복 없이 섞을 수 없을 때가 있음
|
|
158
|
+
stuck.append((k, len(idx)))
|
|
159
|
+
if stuck:
|
|
160
|
+
import warnings
|
|
161
|
+
warnings.warn(f"중복 없이 섞을 수 없어 원래 배선으로 둔 연결 쌍 {len(stuck)}개 "
|
|
162
|
+
f"(연결 {sum(n for _, n in stuck):,}개 / 전체 {self.n_edges:,}개): "
|
|
163
|
+
f"{', '.join(k for k, _ in stuck[:5])}{' …' if len(stuck) > 5 else ''}")
|
|
164
|
+
what = "" if pairs is None and exclude is None else \
|
|
165
|
+
f" {'+'.join(pairs) if pairs is not None else 'all'}{' -' + '-'.join(exclude) if exclude else ''}"
|
|
166
|
+
if local is not None:
|
|
167
|
+
what += f" local r={local[1]:g}"
|
|
168
|
+
if merge:
|
|
169
|
+
what += f" merge {'+'.join(sorted(set(merge.values())))}"
|
|
170
|
+
return Circuit(self.root_ids, self.groups, self.pre, post, self.weight,
|
|
171
|
+
name=f"{self.name} [shuffled{what}]", meta=self.meta, pos=self.pos)
|
|
172
|
+
|
|
173
|
+
def subset(self, groups) -> "Circuit":
|
|
174
|
+
"""지정한 그룹들의 뉴런만 남긴 회로 (그 사이 연결만)"""
|
|
175
|
+
groups = [g for g in groups if g in self.groups]
|
|
176
|
+
keep = np.concatenate([self.groups[g] for g in groups])
|
|
177
|
+
new = np.full(self.N, -1, np.int64); new[keep] = np.arange(len(keep))
|
|
178
|
+
m = (new[self.pre] >= 0) & (new[self.post] >= 0)
|
|
179
|
+
gidx, o = {}, 0
|
|
180
|
+
for g in groups:
|
|
181
|
+
gidx[g] = np.arange(o, o + len(self.groups[g])); o += len(self.groups[g])
|
|
182
|
+
return Circuit(self.root_ids[keep], gidx, new[self.pre[m]], new[self.post[m]], self.weight[m],
|
|
183
|
+
name=f"{self.name} [그룹 {len(groups)}개]",
|
|
184
|
+
meta=self.meta.iloc[keep] if self.meta is not None else None,
|
|
185
|
+
pos=self.pos[keep] if self.pos is not None else None)
|
|
186
|
+
|
|
187
|
+
def normalized(self) -> "Circuit":
|
|
188
|
+
"""받는 뉴런마다 입력 시냅스 수 합(|weight|)이 1이 되도록 나눈 회로.
|
|
189
|
+
입력 비율(누가 얼마나 주는지)은 그대로, 입력이 많은 뉴런과 적은 뉴런의 총입력 크기만 맞춤"""
|
|
190
|
+
tot = np.bincount(self.post, weights=np.abs(self.weight), minlength=self.N)
|
|
191
|
+
w = (self.weight / tot[self.post]).astype(np.float32)
|
|
192
|
+
return Circuit(self.root_ids, self.groups, self.pre, self.post, w, name=f"{self.name} [정규화]",
|
|
193
|
+
meta=self.meta, pos=self.pos)
|
|
194
|
+
|
|
195
|
+
def with_sign(self, pre_groups, sign: int) -> "Circuit":
|
|
196
|
+
"""pre_groups 뉴런이 보내는 연결을 모두 흥분(+1) 또는 억제(-1)로 바꾼 회로.
|
|
197
|
+
예: 광수용체의 히스타민은 받는 뉴런을 억제하지만 신경전달물질 예측에는 흥분으로 잡힘"""
|
|
198
|
+
if sign not in (1, -1):
|
|
199
|
+
raise ValueError("sign은 1 또는 -1")
|
|
200
|
+
pre_groups = [pre_groups] if isinstance(pre_groups, str) else list(pre_groups)
|
|
201
|
+
m = np.isin(self.pre, np.concatenate([self.groups[g] for g in pre_groups]))
|
|
202
|
+
w = self.weight.copy(); w[m] = sign * np.abs(w[m])
|
|
203
|
+
return Circuit(self.root_ids, self.groups, self.pre, self.post, w,
|
|
204
|
+
name=f"{self.name} [{'+'.join(pre_groups)} {'흥분' if sign > 0 else '억제'}]",
|
|
205
|
+
meta=self.meta, pos=self.pos)
|
|
206
|
+
|
|
207
|
+
def summary(self) -> pd.DataFrame:
|
|
208
|
+
"""그룹 간 연결 요약 (시냅스 수, 흥분/억제)"""
|
|
209
|
+
g = self.group_of()
|
|
210
|
+
df = pd.DataFrame({"pre": g[self.pre], "post": g[self.post], "w": self.weight})
|
|
211
|
+
return df.groupby(["pre", "post"]).agg(edges=("w", "size"), exc_syn=("w", lambda x: x[x > 0].sum()),
|
|
212
|
+
inh_syn=("w", lambda x: -x[x < 0].sum()))
|
|
213
|
+
|
|
214
|
+
def __repr__(self):
|
|
215
|
+
gs = ", ".join(f"{k} {len(v)}" for k, v in self.groups.items())
|
|
216
|
+
return f"<Circuit '{self.name}' | {self.N:,} neurons ({gs}) | {self.n_edges:,} edges>"
|
|
217
|
+
|
|
218
|
+
# 저장: 텐서·문자열·리스트만 써서 torch.load(weights_only=True)로 안전하게 읽힘
|
|
219
|
+
def to_dict(self) -> dict:
|
|
220
|
+
import torch
|
|
221
|
+
meta = None
|
|
222
|
+
if self.meta is not None:
|
|
223
|
+
meta = {c: [None if pd.isna(v) else str(v) for v in self.meta[c]] for c in self.meta.columns}
|
|
224
|
+
return dict(root_ids=torch.from_numpy(self.root_ids), pre=torch.from_numpy(self.pre),
|
|
225
|
+
post=torch.from_numpy(self.post), weight=torch.from_numpy(self.weight),
|
|
226
|
+
groups={k: torch.from_numpy(v) for k, v in self.groups.items()}, name=self.name, meta=meta,
|
|
227
|
+
pos=torch.from_numpy(self.pos) if self.pos is not None else None)
|
|
228
|
+
|
|
229
|
+
@classmethod
|
|
230
|
+
def from_dict(cls, d: dict) -> "Circuit":
|
|
231
|
+
meta = pd.DataFrame(d["meta"]) if d.get("meta") is not None else None
|
|
232
|
+
return cls(d["root_ids"].numpy(), {k: v.numpy() for k, v in d["groups"].items()}, d["pre"].numpy(),
|
|
233
|
+
d["post"].numpy(), d["weight"].numpy(), name=d["name"], meta=meta,
|
|
234
|
+
pos=d["pos"].numpy() if d.get("pos") is not None else None)
|
flydnet/data.py
ADDED
|
@@ -0,0 +1,145 @@
|
|
|
1
|
+
"""데이터 위치 설정과 다운로드
|
|
2
|
+
|
|
3
|
+
데이터 묶음 두 가지
|
|
4
|
+
flywire : FlyWire v783 연결·뉴런 목록 (Shiu et al. 2024 모델 저장소) + 세포 유형 주석 (Schlegel et al. 2024)
|
|
5
|
+
door : DoOR 2.0 냄새 반응 (Münch & Galizia 2016, CC BY-SA 4.0)
|
|
6
|
+
|
|
7
|
+
위치를 찾는 순서 (묶음마다)
|
|
8
|
+
1. 함수에 직접 준 경로
|
|
9
|
+
2. 환경변수 FLYDNET_FLYWIRE / FLYDNET_DOOR (FLYDNET_DATA도 flywire로 인정 — 이전 버전 호환)
|
|
10
|
+
3. 설정 파일 ~/.flydnet/config.json ← fd.set_data_dir()로 저장
|
|
11
|
+
4. 기본값 ~/.flydnet/data/<묶음>
|
|
12
|
+
|
|
13
|
+
import flydnet as fd
|
|
14
|
+
fd.download() # 없는 파일만 받음 (약 140 MB)
|
|
15
|
+
fd.set_data_dir(flywire=r"D:\\my\\flywire") # 이미 받아 둔 곳을 쓰려면
|
|
16
|
+
fd.data_status() # 어디서 무엇을 찾았는지
|
|
17
|
+
"""
|
|
18
|
+
from __future__ import annotations
|
|
19
|
+
|
|
20
|
+
import json
|
|
21
|
+
import os
|
|
22
|
+
import urllib.request
|
|
23
|
+
from pathlib import Path
|
|
24
|
+
|
|
25
|
+
CONFIG = Path.home() / ".flydnet" / "config.json"
|
|
26
|
+
DEFAULT_ROOT = Path.home() / ".flydnet" / "data"
|
|
27
|
+
|
|
28
|
+
# 파일 이름 → (URL, 바이트 크기). 크기로 다운로드가 온전한지 확인
|
|
29
|
+
SOURCES = {
|
|
30
|
+
"flywire": {
|
|
31
|
+
"Connectivity_783.parquet":
|
|
32
|
+
("https://github.com/philshiu/Drosophila_brain_model/raw/main/Connectivity_783.parquet", 100804642),
|
|
33
|
+
"Completeness_783.csv":
|
|
34
|
+
("https://github.com/philshiu/Drosophila_brain_model/raw/main/Completeness_783.csv", 3327347),
|
|
35
|
+
"flywire_annotations.tsv":
|
|
36
|
+
("https://github.com/flyconnectome/flywire_annotations/raw/main/supplemental_files/"
|
|
37
|
+
"Supplemental_file1_neuron_annotations.tsv", 31720298),
|
|
38
|
+
},
|
|
39
|
+
"door": {
|
|
40
|
+
name: (f"https://raw.githubusercontent.com/ropensci/DoOR.data/master/data/{name}", None)
|
|
41
|
+
for name in ("door_response_matrix.csv", "door_mappings.csv", "odor.csv")
|
|
42
|
+
},
|
|
43
|
+
}
|
|
44
|
+
CITATIONS = {
|
|
45
|
+
"flywire": "FlyWire: Dorkenwald et al. 2024, Schlegel et al. 2024 (Nature); "
|
|
46
|
+
"연결 파일: Shiu et al. 2024 (Nature), github.com/philshiu/Drosophila_brain_model (MIT)",
|
|
47
|
+
"door": "DoOR 2.0: Münch & Galizia 2016 (Sci Rep 6:21841), github.com/ropensci/DoOR.data (CC BY-SA 4.0)",
|
|
48
|
+
}
|
|
49
|
+
_ENV = {"flywire": ("FLYDNET_FLYWIRE", "FLYDNET_DATA"), "door": ("FLYDNET_DOOR",)}
|
|
50
|
+
|
|
51
|
+
|
|
52
|
+
def _config() -> dict:
|
|
53
|
+
try:
|
|
54
|
+
return json.loads(CONFIG.read_text(encoding="utf-8"))
|
|
55
|
+
except (FileNotFoundError, json.JSONDecodeError):
|
|
56
|
+
return {}
|
|
57
|
+
|
|
58
|
+
|
|
59
|
+
def data_dir(kind: str = "flywire", path: str | Path | None = None) -> Path:
|
|
60
|
+
"""묶음(kind)의 데이터 폴더. 파일이 실제로 있는지는 확인하지 않음 (require()가 확인)"""
|
|
61
|
+
if kind not in SOURCES:
|
|
62
|
+
raise ValueError(f"kind는 {list(SOURCES)} 중 하나")
|
|
63
|
+
if path is not None:
|
|
64
|
+
return Path(path)
|
|
65
|
+
for env in _ENV[kind]:
|
|
66
|
+
if os.environ.get(env):
|
|
67
|
+
return Path(os.environ[env])
|
|
68
|
+
if kind in _config():
|
|
69
|
+
return Path(_config()[kind])
|
|
70
|
+
return DEFAULT_ROOT / kind
|
|
71
|
+
|
|
72
|
+
|
|
73
|
+
def set_data_dir(flywire: str | Path | None = None, door: str | Path | None = None):
|
|
74
|
+
"""이 컴퓨터에서 쓸 데이터 위치를 ~/.flydnet/config.json에 저장 (None인 항목은 그대로)"""
|
|
75
|
+
cfg = _config()
|
|
76
|
+
for kind, p in (("flywire", flywire), ("door", door)):
|
|
77
|
+
if p is not None:
|
|
78
|
+
cfg[kind] = str(Path(p).resolve())
|
|
79
|
+
CONFIG.parent.mkdir(parents=True, exist_ok=True)
|
|
80
|
+
CONFIG.write_text(json.dumps(cfg, indent=2, ensure_ascii=False), encoding="utf-8")
|
|
81
|
+
return cfg
|
|
82
|
+
|
|
83
|
+
|
|
84
|
+
def missing(kind: str = "flywire", path=None) -> list[str]:
|
|
85
|
+
d = data_dir(kind, path)
|
|
86
|
+
return [f for f in SOURCES[kind] if not (d / f).exists()]
|
|
87
|
+
|
|
88
|
+
|
|
89
|
+
def require(kind: str = "flywire", path=None) -> Path:
|
|
90
|
+
"""데이터 폴더를 돌려주되, 파일이 없으면 무엇을 하면 되는지 알려주는 오류"""
|
|
91
|
+
d = data_dir(kind, path)
|
|
92
|
+
lack = missing(kind, path)
|
|
93
|
+
if lack:
|
|
94
|
+
raise FileNotFoundError(
|
|
95
|
+
f"{kind} 데이터 파일이 없음: {lack}\n 찾은 위치: {d}\n"
|
|
96
|
+
f" 받기: flydnet.download('{kind}') 또는 python -m flydnet download {kind}\n"
|
|
97
|
+
f" 이미 있으면: flydnet.set_data_dir({kind}=r'경로') 또는 환경변수 {_ENV[kind][0]}")
|
|
98
|
+
return d
|
|
99
|
+
|
|
100
|
+
|
|
101
|
+
def download(kinds=("flywire", "door"), path=None, overwrite: bool = False, quiet: bool = False) -> dict:
|
|
102
|
+
"""없는 파일만 받음. 반환: {묶음: 폴더}"""
|
|
103
|
+
kinds = [kinds] if isinstance(kinds, str) else list(kinds)
|
|
104
|
+
out = {}
|
|
105
|
+
for kind in kinds:
|
|
106
|
+
d = data_dir(kind, path if len(kinds) == 1 else None)
|
|
107
|
+
d.mkdir(parents=True, exist_ok=True)
|
|
108
|
+
for name, (url, size) in SOURCES[kind].items():
|
|
109
|
+
f = d / name
|
|
110
|
+
if f.exists() and not overwrite and (size is None or f.stat().st_size == size):
|
|
111
|
+
if not quiet:
|
|
112
|
+
print(f" 있음 {f}")
|
|
113
|
+
continue
|
|
114
|
+
tmp = f.with_suffix(f.suffix + ".part")
|
|
115
|
+
if not quiet:
|
|
116
|
+
print(f" 받는 중 {name} ← {url}", flush=True)
|
|
117
|
+
_fetch(url, tmp, quiet)
|
|
118
|
+
if size is not None and tmp.stat().st_size != size:
|
|
119
|
+
got = tmp.stat().st_size
|
|
120
|
+
tmp.unlink()
|
|
121
|
+
raise IOError(f"{name} 크기가 다름 (받은 {got:,} B, 기대 {size:,} B) — 다시 시도")
|
|
122
|
+
tmp.replace(f)
|
|
123
|
+
if not quiet:
|
|
124
|
+
print(f"[{kind}] {d}\n 출처: {CITATIONS[kind]}")
|
|
125
|
+
out[kind] = d
|
|
126
|
+
return out
|
|
127
|
+
|
|
128
|
+
|
|
129
|
+
def _fetch(url: str, dest: Path, quiet: bool):
|
|
130
|
+
with urllib.request.urlopen(url, timeout=60) as r, open(dest, "wb") as f:
|
|
131
|
+
total = int(r.headers.get("Content-Length") or 0)
|
|
132
|
+
done, step = 0, max(total // 10, 1)
|
|
133
|
+
while chunk := r.read(1 << 20):
|
|
134
|
+
f.write(chunk)
|
|
135
|
+
done += len(chunk)
|
|
136
|
+
if not quiet and total and done // step != (done - len(chunk)) // step:
|
|
137
|
+
print(f" {done / total * 100:3.0f}%", flush=True)
|
|
138
|
+
|
|
139
|
+
|
|
140
|
+
def data_status() -> dict:
|
|
141
|
+
"""묶음별 위치와 빠진 파일"""
|
|
142
|
+
st = {k: dict(dir=str(data_dir(k)), missing=missing(k)) for k in SOURCES}
|
|
143
|
+
for k, v in st.items():
|
|
144
|
+
print(f"{k:<8} {'준비됨' if not v['missing'] else '빠짐: ' + ', '.join(v['missing'])} ({v['dir']})")
|
|
145
|
+
return st
|
flydnet/datasets.py
ADDED
|
@@ -0,0 +1,93 @@
|
|
|
1
|
+
"""시험용 데이터 생성"""
|
|
2
|
+
import torch
|
|
3
|
+
|
|
4
|
+
|
|
5
|
+
def synthetic_odors(n_classes: int, n_glomeruli: int, n_train: int, n_test: int, protos_per_class: int = 1,
|
|
6
|
+
active_frac: float = 0.2, noise: float = 0.5, add_noise: float = 0.1, seed: int = 0):
|
|
7
|
+
"""합성 냄새 분류 과제
|
|
8
|
+
|
|
9
|
+
냄새 원형: 사구체마다 active_frac 확률로 켜짐, 세기 U(0.2, 1)
|
|
10
|
+
클래스: 원형 protos_per_class개의 묶음 (예: '먹이' = 사과 냄새 또는 효모 냄새).
|
|
11
|
+
2개 이상이면 사구체 공간에서 선형 분리가 어려워짐 → 확장 층(KC)이 필요한 과제
|
|
12
|
+
샘플: 클래스의 원형 하나를 골라 × exp(N(0, noise)) (사구체별 세기 흔들림) + |N(0, add_noise)| (배경)
|
|
13
|
+
반환: (Xtr, ytr, Xte, yte), X는 (n, n_glomeruli), 클래스당 n_train / n_test 개
|
|
14
|
+
"""
|
|
15
|
+
g = torch.Generator().manual_seed(seed)
|
|
16
|
+
P = n_classes * protos_per_class
|
|
17
|
+
on = torch.rand(P, n_glomeruli, generator=g) < active_frac
|
|
18
|
+
protos = on * (0.2 + 0.8 * torch.rand(P, n_glomeruli, generator=g))
|
|
19
|
+
|
|
20
|
+
def sample(n):
|
|
21
|
+
y = torch.arange(n_classes).repeat_interleave(n)
|
|
22
|
+
k = y * protos_per_class + torch.randint(0, protos_per_class, (len(y),), generator=g)
|
|
23
|
+
x = protos[k] * torch.exp(noise * torch.randn(len(y), n_glomeruli, generator=g))
|
|
24
|
+
x = x + (add_noise * torch.randn(len(y), n_glomeruli, generator=g)).abs()
|
|
25
|
+
return x, y
|
|
26
|
+
|
|
27
|
+
Xtr, ytr = sample(n_train)
|
|
28
|
+
Xte, yte = sample(n_test)
|
|
29
|
+
return Xtr, ytr, Xte, yte
|
|
30
|
+
|
|
31
|
+
|
|
32
|
+
# 같은 계열인데 이름이 다른 DoOR 화학 계열
|
|
33
|
+
_CLASS_ALIASES = {"arom": "aromatics", "terpenes": "terpene", "sulfid": "sulfide"}
|
|
34
|
+
|
|
35
|
+
|
|
36
|
+
def door_odors(glomeruli, data_dir=None, min_measured: int = 20):
|
|
37
|
+
"""DoOR 2.0 실제 냄새 반응 → 사구체 벡터 (Münch & Galizia 2016, CC BY-SA 4.0)
|
|
38
|
+
|
|
39
|
+
data_dir에 door_response_matrix.csv, door_mappings.csv, odor.csv 필요. None이면 flydnet.data_dir("door")
|
|
40
|
+
(없으면 flydnet.download("door"), 출처 https://github.com/ropensci/DoOR.data 의 data/ 폴더)
|
|
41
|
+
|
|
42
|
+
glomeruli: 사구체 이름 순서 (예: GlomerularEncoder.glomeruli)
|
|
43
|
+
반환 dict:
|
|
44
|
+
X (n_odors, n_glomeruli) 반응 − 자발 발화(SFR), 0 아래는 0. 측정 안 된 칸도 0
|
|
45
|
+
measured (n_odors, n_glomeruli) 측정 여부
|
|
46
|
+
names, classes (화학 계열), inchikey
|
|
47
|
+
min_measured: 이만큼 이상의 사구체가 측정된 냄새만
|
|
48
|
+
"""
|
|
49
|
+
import numpy as np
|
|
50
|
+
import pandas as pd
|
|
51
|
+
from .data import require
|
|
52
|
+
|
|
53
|
+
d = require("door", data_dir)
|
|
54
|
+
R = pd.read_csv(d / "door_response_matrix.csv", sep=";")
|
|
55
|
+
M = pd.read_csv(d / "door_mappings.csv", sep=";")
|
|
56
|
+
O = pd.read_csv(d / "odor.csv", sep=";").drop_duplicates("InChIKey").set_index("InChIKey")
|
|
57
|
+
sfr = R.loc["SFR"]
|
|
58
|
+
R = R.drop(index="SFR")
|
|
59
|
+
|
|
60
|
+
col = {g: j for j, g in enumerate(glomeruli)}
|
|
61
|
+
X = np.zeros((len(R), len(glomeruli)), np.float32)
|
|
62
|
+
meas = np.zeros_like(X, dtype=bool)
|
|
63
|
+
for rec, glo in M[["receptor", "code"]].dropna().itertuples(index=False):
|
|
64
|
+
if rec in R.columns and glo in col:
|
|
65
|
+
v = R[rec].values
|
|
66
|
+
ok = ~np.isnan(v)
|
|
67
|
+
X[ok, col[glo]] = np.maximum(v[ok] - sfr[rec], 0)
|
|
68
|
+
meas[ok, col[glo]] = True
|
|
69
|
+
|
|
70
|
+
keep = meas.sum(1) >= min_measured
|
|
71
|
+
info = O.reindex(R.index[keep])
|
|
72
|
+
classes = info["Class"].fillna("other").replace(_CLASS_ALIASES).values
|
|
73
|
+
return dict(X=torch.tensor(X[keep]), measured=torch.tensor(meas[keep]), names=info["Name"].fillna("").values,
|
|
74
|
+
classes=classes, inchikey=R.index[keep].values)
|
|
75
|
+
|
|
76
|
+
|
|
77
|
+
def biconditional_mixtures(X0, sets, n: int, noise: float = 0.5, add_noise: float = 0.05, sat: float = 0.3,
|
|
78
|
+
generator=None):
|
|
79
|
+
"""조건부 구별 과제 샘플: 냄새 4개 (a, b, c, d) 묶음마다 AB, CD → 1 (보상), AC, BD → 0 (무보상)
|
|
80
|
+
|
|
81
|
+
혼합물 = 포화(Σ 성분 × exp(N(0, noise)) + |N(0, add_noise)|), 포화(x) = x / (x + sat)
|
|
82
|
+
반환: (x, y, 묶음 번호), 묶음·혼합물마다 n개
|
|
83
|
+
"""
|
|
84
|
+
xs, ys, ps = [], [], []
|
|
85
|
+
G = X0.shape[1]
|
|
86
|
+
for p, (a, b, c, d) in enumerate(sets):
|
|
87
|
+
for comp, label in (((a, b), 1), ((c, d), 1), ((a, c), 0), ((b, d), 0)):
|
|
88
|
+
base = sum(X0[k] * torch.exp(noise * torch.randn(n, G, generator=generator)) for k in comp)
|
|
89
|
+
x = base + (add_noise * torch.randn(n, G, generator=generator)).abs()
|
|
90
|
+
xs.append(x / (x + sat))
|
|
91
|
+
ys += [label] * n
|
|
92
|
+
ps += [p] * n
|
|
93
|
+
return torch.cat(xs), torch.tensor(ys), torch.tensor(ps)
|
flydnet/encoders.py
ADDED
|
@@ -0,0 +1,71 @@
|
|
|
1
|
+
"""텐서 → 입력 뉴런 발화율(Hz) 변환. '설탕 뉴런 150Hz 자극'을 일반 데이터로 확장한 것"""
|
|
2
|
+
import numpy as np
|
|
3
|
+
import torch
|
|
4
|
+
import torch.nn as nn
|
|
5
|
+
|
|
6
|
+
|
|
7
|
+
def _to_rates(x: torch.Tensor, max_rate: float) -> torch.Tensor:
|
|
8
|
+
"""음수는 0, 샘플마다 최댓값 = max_rate (전체 세기가 달라도 같은 패턴이면 같은 발화율)"""
|
|
9
|
+
x = x.clamp_min(0)
|
|
10
|
+
return x / x.amax(1, keepdim=True).clamp_min(1e-8) * max_rate
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
class RateEncoder(nn.Module):
|
|
14
|
+
"""입력 특징 n_in개를 입력 뉴런 n_out개의 발화율로 바꿈
|
|
15
|
+
|
|
16
|
+
- n_in == n_out 이고 projection=None: 특징 하나 = 뉴런 하나
|
|
17
|
+
- 아니면 고정 무작위 희소 투영: 뉴런마다 특징 k개를 모아 받음 (사구체가 여러 수용체 입력을 모으듯)
|
|
18
|
+
출력은 샘플마다 최댓값이 max_rate가 되도록 정규화 (음수는 0)
|
|
19
|
+
"""
|
|
20
|
+
|
|
21
|
+
def __init__(self, n_in: int, n_out: int, max_rate: float = 100.0, k: int = 20, seed: int = 0,
|
|
22
|
+
projection: str | None = "random"):
|
|
23
|
+
super().__init__()
|
|
24
|
+
self.n_in, self.n_out, self.max_rate = n_in, n_out, max_rate
|
|
25
|
+
if projection is None:
|
|
26
|
+
if n_in != n_out:
|
|
27
|
+
raise ValueError("projection=None이면 n_in == n_out 이어야 함")
|
|
28
|
+
self.P = None
|
|
29
|
+
else:
|
|
30
|
+
g = torch.Generator().manual_seed(seed)
|
|
31
|
+
cols = torch.stack([torch.randperm(n_in, generator=g)[:k] for _ in range(n_out)])
|
|
32
|
+
P = torch.zeros(n_out, n_in)
|
|
33
|
+
P.scatter_(1, cols, 1.0 / k)
|
|
34
|
+
self.register_buffer("P", P)
|
|
35
|
+
|
|
36
|
+
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
|
37
|
+
x = x.flatten(1).float()
|
|
38
|
+
if self.P is not None:
|
|
39
|
+
x = x @ self.P.T.to(x.device)
|
|
40
|
+
return _to_rates(x, self.max_rate)
|
|
41
|
+
|
|
42
|
+
|
|
43
|
+
class GlomerularEncoder(nn.Module):
|
|
44
|
+
"""냄새 = 사구체별 활성 벡터 (B, n_glomeruli) → 투사 뉴런(PN) 발화율 (B, n_PN)
|
|
45
|
+
|
|
46
|
+
실제 더듬이엽 구조를 따름: 같은 사구체의 단일 사구체형 PN들은 같은 발화율을 받음.
|
|
47
|
+
다중 사구체형 PN은 입력 0 (오른쪽 버섯체에서 PN→KC 시냅스의 96%가 단일 사구체형).
|
|
48
|
+
사구체 이름은 cell_type의 '_' 앞부분 (예: DM1_lPN → DM1)
|
|
49
|
+
"""
|
|
50
|
+
|
|
51
|
+
def __init__(self, circuit, group: str = "PN", max_rate: float = 100.0):
|
|
52
|
+
super().__init__()
|
|
53
|
+
if circuit.meta is None:
|
|
54
|
+
raise ValueError("circuit.meta(세포 주석)가 필요함 — Circuit.from_flywire()로 만든 회로를 쓸 것")
|
|
55
|
+
m = circuit.meta.iloc[circuit.groups[group]]
|
|
56
|
+
uni = m.cell_sub_class.astype(str).eq("uniglomerular").values
|
|
57
|
+
glom = m.cell_type.astype(str).str.split("_").str[0].values
|
|
58
|
+
self.glomeruli = sorted(set(glom[uni]))
|
|
59
|
+
col = {gname: j for j, gname in enumerate(self.glomeruli)}
|
|
60
|
+
P = torch.zeros(len(m), len(self.glomeruli))
|
|
61
|
+
for i in np.nonzero(uni)[0]:
|
|
62
|
+
P[i, col[glom[i]]] = 1.0
|
|
63
|
+
self.register_buffer("P", P)
|
|
64
|
+
self.max_rate = max_rate
|
|
65
|
+
|
|
66
|
+
@property
|
|
67
|
+
def n_glomeruli(self) -> int:
|
|
68
|
+
return len(self.glomeruli)
|
|
69
|
+
|
|
70
|
+
def forward(self, odor: torch.Tensor) -> torch.Tensor:
|
|
71
|
+
return _to_rates(odor.float() @ self.P.T.to(odor.device), self.max_rate)
|