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 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)