SeInE-orientation 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.
- seine_orientation/BayesModel.py +283 -0
- seine_orientation/__init__.py +2 -0
- seine_orientation/_version.py +24 -0
- seine_orientation/process_session.py +269 -0
- seine_orientation/utils/__init__.py +11 -0
- seine_orientation/utils/display_inference_results.py +443 -0
- seine_orientation/utils/utils_analysis.py +239 -0
- seine_orientation/utils/utils_display.py +109 -0
- seine_orientation/utils/utils_model.py +283 -0
- seine_orientation-0.1.0.dist-info/METADATA +13 -0
- seine_orientation-0.1.0.dist-info/RECORD +13 -0
- seine_orientation-0.1.0.dist-info/WHEEL +5 -0
- seine_orientation-0.1.0.dist-info/top_level.txt +1 -0
|
@@ -0,0 +1,283 @@
|
|
|
1
|
+
import numpy as np
|
|
2
|
+
import time
|
|
3
|
+
import itertools
|
|
4
|
+
|
|
5
|
+
from seine.NestedSamplingMethods import (
|
|
6
|
+
run_sampling,
|
|
7
|
+
)
|
|
8
|
+
from seine import HierarchicalModel, functions as prior_fn
|
|
9
|
+
from seine.structures import (
|
|
10
|
+
prior_structure,
|
|
11
|
+
)
|
|
12
|
+
|
|
13
|
+
from .utils import gabor_filter, gabor_response, sine_grating, softplus, ReLU, sigmoid
|
|
14
|
+
|
|
15
|
+
|
|
16
|
+
class HierarchicalBayesInference(HierarchicalModel):
|
|
17
|
+
|
|
18
|
+
def set_gratings(self, measure_points, FoV_range, FoV_steps):
|
|
19
|
+
|
|
20
|
+
self.measure_points = measure_points
|
|
21
|
+
|
|
22
|
+
self.X_FoV, self.Y_FoV = np.meshgrid(
|
|
23
|
+
np.linspace(-FoV_range[0] / 2, FoV_range[0] / 2, FoV_steps),
|
|
24
|
+
np.linspace(-FoV_range[1] / 2, FoV_range[1] / 2, FoV_steps),
|
|
25
|
+
)
|
|
26
|
+
self.FoV_grid = np.dstack((self.X_FoV, self.Y_FoV))
|
|
27
|
+
|
|
28
|
+
FoV_steps = self.X_FoV.shape[0]
|
|
29
|
+
|
|
30
|
+
self.gratings = np.zeros(
|
|
31
|
+
(
|
|
32
|
+
*[len(mp) for mp in self.measure_points],
|
|
33
|
+
FoV_steps,
|
|
34
|
+
FoV_steps,
|
|
35
|
+
)
|
|
36
|
+
)
|
|
37
|
+
for prod in itertools.product(
|
|
38
|
+
enumerate(self.measure_points[0]),
|
|
39
|
+
enumerate(self.measure_points[1]),
|
|
40
|
+
enumerate(self.measure_points[2]),
|
|
41
|
+
):
|
|
42
|
+
idx, elems = zip(*prod)
|
|
43
|
+
phi_0, theta, f = elems
|
|
44
|
+
self.gratings[*idx, ...] = sine_grating(
|
|
45
|
+
self.X_FoV, # [0, ...],
|
|
46
|
+
self.Y_FoV, # [0, ...],
|
|
47
|
+
theta,
|
|
48
|
+
f,
|
|
49
|
+
phi_0,
|
|
50
|
+
square=True,
|
|
51
|
+
)
|
|
52
|
+
|
|
53
|
+
def set_priors(self, priors_init=None, coding="simple"):
|
|
54
|
+
|
|
55
|
+
if priors_init is None:
|
|
56
|
+
|
|
57
|
+
if hasattr(self, "data"):
|
|
58
|
+
fmap = self.data["observed_counts"] / self.data["T"]
|
|
59
|
+
A0_guess, A_guess = np.percentile(fmap, [50, 90])
|
|
60
|
+
A0_guess = np.maximum(A0_guess, 1.0)
|
|
61
|
+
A_guess = np.maximum(A_guess, 2.0)
|
|
62
|
+
else:
|
|
63
|
+
A0_guess, A_guess = 1.0, 3.0
|
|
64
|
+
self.priors_init = {}
|
|
65
|
+
# print("A guesses from data:", A0_guess, A_guess)
|
|
66
|
+
|
|
67
|
+
## define gabor model priors
|
|
68
|
+
if coding in ["simple", "complex"]:
|
|
69
|
+
self.priors_init["theta"] = prior_structure(
|
|
70
|
+
prior_fn.bounded_flat,
|
|
71
|
+
low=-np.pi / 2,
|
|
72
|
+
high=np.pi / 2,
|
|
73
|
+
label=r"$\theta$",
|
|
74
|
+
periodic=True,
|
|
75
|
+
)
|
|
76
|
+
self.priors_init["f"] = prior_structure(
|
|
77
|
+
prior_fn.halfnorm_ppf,
|
|
78
|
+
loc=0.0,
|
|
79
|
+
scale=5.0,
|
|
80
|
+
label=r"$f$",
|
|
81
|
+
)
|
|
82
|
+
if coding == "simple":
|
|
83
|
+
## not required for complex cell model
|
|
84
|
+
self.priors_init["phi_0"] = prior_structure(
|
|
85
|
+
prior_fn.bounded_flat,
|
|
86
|
+
low=-np.pi,
|
|
87
|
+
high=np.pi,
|
|
88
|
+
periodic=True,
|
|
89
|
+
label=r"$\phi_0$",
|
|
90
|
+
)
|
|
91
|
+
else:
|
|
92
|
+
self.priors_init["phi_0"] = prior_structure(
|
|
93
|
+
None,
|
|
94
|
+
value=0.0,
|
|
95
|
+
periodic=True,
|
|
96
|
+
label=r"$\phi_0$",
|
|
97
|
+
)
|
|
98
|
+
|
|
99
|
+
self.priors_init["sigma"] = prior_structure(
|
|
100
|
+
prior_fn.halfnorm_ppf,
|
|
101
|
+
loc=0.05,
|
|
102
|
+
scale=0.1,
|
|
103
|
+
label=r"$\sigma$",
|
|
104
|
+
)
|
|
105
|
+
self.priors_init["gamma"] = prior_structure(
|
|
106
|
+
prior_fn.halfnorm_ppf,
|
|
107
|
+
loc=1.0,
|
|
108
|
+
scale=1.0,
|
|
109
|
+
label=r"$\gamma$",
|
|
110
|
+
)
|
|
111
|
+
self.priors_init["theta_gauss"] = prior_structure(
|
|
112
|
+
prior_fn.bounded_flat,
|
|
113
|
+
low=-np.pi / 2,
|
|
114
|
+
high=np.pi / 2,
|
|
115
|
+
periodic=True,
|
|
116
|
+
label=r"$\theta_{gauss}$",
|
|
117
|
+
)
|
|
118
|
+
|
|
119
|
+
## define nonlinearity priors - should start with "nl_" to be recognized by model response function
|
|
120
|
+
self.priors_init["nl_baseline"] = prior_structure(
|
|
121
|
+
prior_fn.halfnorm_ppf,
|
|
122
|
+
loc=0.0,
|
|
123
|
+
scale=A0_guess,
|
|
124
|
+
label=r"$A_{baseline}$",
|
|
125
|
+
)
|
|
126
|
+
if coding in ["simple", "complex"]:
|
|
127
|
+
self.priors_init["nl_transition"] = prior_structure(
|
|
128
|
+
prior_fn.halfnorm_ppf,
|
|
129
|
+
loc=0.0,
|
|
130
|
+
scale=1.0,
|
|
131
|
+
label=r"$A_{transition}$",
|
|
132
|
+
)
|
|
133
|
+
self.priors_init["nl_amplitude"] = prior_structure(
|
|
134
|
+
prior_fn.halfnorm_ppf,
|
|
135
|
+
loc=A0_guess * 0.2,
|
|
136
|
+
scale=A_guess - A0_guess,
|
|
137
|
+
label=r"$A_{amplitude}$",
|
|
138
|
+
)
|
|
139
|
+
|
|
140
|
+
# parameters for loglikelihood
|
|
141
|
+
self.priors_init["logl_alpha"] = prior_structure(
|
|
142
|
+
prior_fn.halfnorm_ppf,
|
|
143
|
+
loc=0.0,
|
|
144
|
+
scale=1.0,
|
|
145
|
+
label=r"$\alpha_{logl}$",
|
|
146
|
+
)
|
|
147
|
+
|
|
148
|
+
else:
|
|
149
|
+
self.priors_init = priors_init
|
|
150
|
+
|
|
151
|
+
super().set_priors(self.priors_init)
|
|
152
|
+
|
|
153
|
+
def get_model_rate_response(self, params, coding="simple"):
|
|
154
|
+
|
|
155
|
+
## transform parameters, if not already provided in proper format
|
|
156
|
+
if not isinstance(params, dict):
|
|
157
|
+
params = self.get_params_from_p(params)
|
|
158
|
+
|
|
159
|
+
## get (normalized) model response
|
|
160
|
+
if coding == "random":
|
|
161
|
+
response = np.zeros(self.dimensions["shape"])
|
|
162
|
+
else:
|
|
163
|
+
response = gabor_response(
|
|
164
|
+
self.X_FoV,
|
|
165
|
+
self.Y_FoV,
|
|
166
|
+
params,
|
|
167
|
+
self.gratings,
|
|
168
|
+
mode=coding,
|
|
169
|
+
significance_threshold=0.01,
|
|
170
|
+
logger=self.log,
|
|
171
|
+
)
|
|
172
|
+
self.timeit("calculating model and response firing rate")
|
|
173
|
+
|
|
174
|
+
## apply nonlinearity to adjust to scale of firing rates and get final model response in firing rate units
|
|
175
|
+
rate_response = self.apply_nonlinearity(response, params)
|
|
176
|
+
return rate_response
|
|
177
|
+
|
|
178
|
+
def apply_nonlinearity(self, response, params, nonlinearity="softplus"):
|
|
179
|
+
|
|
180
|
+
if nonlinearity == "softplus":
|
|
181
|
+
# print(
|
|
182
|
+
# f"Applying softplus nonlinearity with parameters: baseline={params.get('nl_baseline', 0.0):.3f}, transition={params.get('nl_transition', 1.0):.3f}, amplitude={params.get('nl_amplitude', 1.0):.3f}"
|
|
183
|
+
# )
|
|
184
|
+
rate_response = softplus(
|
|
185
|
+
response * params.get("nl_amplitude", 1.0),
|
|
186
|
+
alpha=params.get("nl_transition", 1.0),
|
|
187
|
+
delta=params.get("nl_baseline", 0.0),
|
|
188
|
+
)
|
|
189
|
+
elif nonlinearity == "ReLU":
|
|
190
|
+
# rate_response =
|
|
191
|
+
raise NotImplementedError("ReLU nonlinearity not implemented yet")
|
|
192
|
+
elif nonlinearity == "sigmoid":
|
|
193
|
+
raise NotImplementedError("Sigmoid nonlinearity not implemented yet")
|
|
194
|
+
else:
|
|
195
|
+
raise ValueError(f"Unknown nonlinearity: {nonlinearity}")
|
|
196
|
+
|
|
197
|
+
self.timeit("applying nonlinearity")
|
|
198
|
+
return rate_response
|
|
199
|
+
|
|
200
|
+
def set_logp_func(self, coding="simple"):
|
|
201
|
+
"""
|
|
202
|
+
some nice description
|
|
203
|
+
"""
|
|
204
|
+
|
|
205
|
+
def get_logp(p_in):
|
|
206
|
+
|
|
207
|
+
self.timeit()
|
|
208
|
+
"""
|
|
209
|
+
build switch between 3 models:
|
|
210
|
+
- gabor type model
|
|
211
|
+
- fourier component model (with n components in each direction)
|
|
212
|
+
- ellipse in rate space
|
|
213
|
+
"""
|
|
214
|
+
|
|
215
|
+
params = self.get_params_from_p(p_in)
|
|
216
|
+
self.timeit("transforming parameters")
|
|
217
|
+
|
|
218
|
+
model_rate_response = self.get_model_rate_response(params, coding)
|
|
219
|
+
if not (model_rate_response.shape == self.data["observed_counts"].shape):
|
|
220
|
+
model_rate_response = model_rate_response[None, ...]
|
|
221
|
+
|
|
222
|
+
logp = self.probability_of_spike_observation(
|
|
223
|
+
model_rate_response,
|
|
224
|
+
model="negative_binomial",
|
|
225
|
+
**params,
|
|
226
|
+
)
|
|
227
|
+
self.timeit("calculating log probability of spike observation")
|
|
228
|
+
|
|
229
|
+
return logp.sum()
|
|
230
|
+
|
|
231
|
+
return get_logp
|
|
232
|
+
|
|
233
|
+
|
|
234
|
+
def run_inference(
|
|
235
|
+
hbm: HierarchicalBayesInference,
|
|
236
|
+
coding="simple",
|
|
237
|
+
n_live=200,
|
|
238
|
+
nP=1,
|
|
239
|
+
dlogz=1.0,
|
|
240
|
+
show_status=True,
|
|
241
|
+
):
|
|
242
|
+
hbm.set_priors(coding=coding)
|
|
243
|
+
my_trafo = hbm.set_prior_transform()
|
|
244
|
+
my_logp = hbm.set_logp_func(coding=coding)
|
|
245
|
+
|
|
246
|
+
results, sampler = run_sampling(
|
|
247
|
+
my_trafo,
|
|
248
|
+
my_logp,
|
|
249
|
+
hbm.parameter_names_all,
|
|
250
|
+
hbm.periodic,
|
|
251
|
+
show_status=show_status,
|
|
252
|
+
n_live=n_live,
|
|
253
|
+
nP=nP,
|
|
254
|
+
dlogz=dlogz,
|
|
255
|
+
)
|
|
256
|
+
return results
|
|
257
|
+
|
|
258
|
+
|
|
259
|
+
def run_model_comparison(data, dwelltime, measure_points, show_status=True, **kwargs):
|
|
260
|
+
|
|
261
|
+
t_start = time.time()
|
|
262
|
+
hbm = HierarchicalBayesInference()
|
|
263
|
+
|
|
264
|
+
hbm.set_gratings(
|
|
265
|
+
measure_points,
|
|
266
|
+
[np.deg2rad(140), np.deg2rad(114)],
|
|
267
|
+
FoV_steps=51,
|
|
268
|
+
)
|
|
269
|
+
|
|
270
|
+
hbm.prepare_data(
|
|
271
|
+
data,
|
|
272
|
+
dwelltime,
|
|
273
|
+
iter_dims=False,
|
|
274
|
+
)
|
|
275
|
+
|
|
276
|
+
results = {}
|
|
277
|
+
for coding in ["random", "simple", "complex"]:
|
|
278
|
+
results[coding] = run_inference(
|
|
279
|
+
hbm, coding=coding, show_status=show_status, **kwargs
|
|
280
|
+
)
|
|
281
|
+
t_end = time.time()
|
|
282
|
+
print(f"Model comparison done after {t_end - t_start:.2f} seconds")
|
|
283
|
+
return results
|
|
@@ -0,0 +1,24 @@
|
|
|
1
|
+
# file generated by vcs-versioning
|
|
2
|
+
# don't change, don't track in version control
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
__all__ = [
|
|
6
|
+
"__version__",
|
|
7
|
+
"__version_tuple__",
|
|
8
|
+
"version",
|
|
9
|
+
"version_tuple",
|
|
10
|
+
"__commit_id__",
|
|
11
|
+
"commit_id",
|
|
12
|
+
]
|
|
13
|
+
|
|
14
|
+
version: str
|
|
15
|
+
__version__: str
|
|
16
|
+
__version_tuple__: tuple[int | str, ...]
|
|
17
|
+
version_tuple: tuple[int | str, ...]
|
|
18
|
+
commit_id: str | None
|
|
19
|
+
__commit_id__: str | None
|
|
20
|
+
|
|
21
|
+
__version__ = version = '0.1.0'
|
|
22
|
+
__version_tuple__ = version_tuple = (0, 1, 0)
|
|
23
|
+
|
|
24
|
+
__commit_id__ = commit_id = 'gffe43a317'
|
|
@@ -0,0 +1,269 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
import h5py, os
|
|
4
|
+
import numpy as np
|
|
5
|
+
from pathlib import Path
|
|
6
|
+
import multiprocessing as mp
|
|
7
|
+
from functools import partial
|
|
8
|
+
from scipy.io import loadmat
|
|
9
|
+
|
|
10
|
+
from pathlib import Path
|
|
11
|
+
from typing import Optional
|
|
12
|
+
|
|
13
|
+
import argparse
|
|
14
|
+
|
|
15
|
+
|
|
16
|
+
def get_folder_path(
|
|
17
|
+
mouse,
|
|
18
|
+
session,
|
|
19
|
+
*,
|
|
20
|
+
type="data",
|
|
21
|
+
dataset="PSD95Mice_Loewel",
|
|
22
|
+
project="tmp",
|
|
23
|
+
):
|
|
24
|
+
return Path(project, type, dataset, mouse, session)
|
|
25
|
+
|
|
26
|
+
|
|
27
|
+
def process_session(
|
|
28
|
+
# define the path towards data, assuming hpc structure
|
|
29
|
+
animal: str,
|
|
30
|
+
session: str,
|
|
31
|
+
fname: str,
|
|
32
|
+
dataset: str = "PSD95Mice_Loewel",
|
|
33
|
+
project: str = "orientation_selectivity",
|
|
34
|
+
spikes_key: str = "/estimates/S_dff",
|
|
35
|
+
response_onset: float = 0.2,
|
|
36
|
+
prefix: str = "SeInE",
|
|
37
|
+
suffix: str = "",
|
|
38
|
+
nP: int = 12,
|
|
39
|
+
save_type: str = "hdf5",
|
|
40
|
+
**kwargs,
|
|
41
|
+
):
|
|
42
|
+
"""
|
|
43
|
+
this assumes some specific structure and naming of data,
|
|
44
|
+
following the HPC project directory conventions.
|
|
45
|
+
"""
|
|
46
|
+
|
|
47
|
+
from .utils.utils_analysis import (
|
|
48
|
+
calculate_firing_maps,
|
|
49
|
+
get_spikes,
|
|
50
|
+
get_unique_stimulus_values,
|
|
51
|
+
)
|
|
52
|
+
from .BayesModel import (
|
|
53
|
+
run_model_comparison,
|
|
54
|
+
)
|
|
55
|
+
|
|
56
|
+
projects_dir = os.environ.get("PROJECT_DIR", "tmp")
|
|
57
|
+
project_dir = Path(projects_dir) / project
|
|
58
|
+
|
|
59
|
+
dir_analysis = get_folder_path(
|
|
60
|
+
animal,
|
|
61
|
+
session,
|
|
62
|
+
type="analysis",
|
|
63
|
+
dataset=dataset,
|
|
64
|
+
project=str(project_dir),
|
|
65
|
+
)
|
|
66
|
+
path_meta = dir_analysis / "CaimanMeta.mat"
|
|
67
|
+
print(path_meta)
|
|
68
|
+
path_detection = dir_analysis / fname
|
|
69
|
+
print(path_detection)
|
|
70
|
+
|
|
71
|
+
dir_data = get_folder_path(
|
|
72
|
+
animal,
|
|
73
|
+
session,
|
|
74
|
+
type="data",
|
|
75
|
+
dataset=dataset,
|
|
76
|
+
project=str(project_dir),
|
|
77
|
+
)
|
|
78
|
+
|
|
79
|
+
## from here, process data
|
|
80
|
+
ld = loadmat(path_meta, variable_names=["CaimanMeta"])
|
|
81
|
+
meta_data = ld["CaimanMeta"]
|
|
82
|
+
|
|
83
|
+
with h5py.File(path_detection, "r") as f:
|
|
84
|
+
S = np.array(f[spikes_key][()])
|
|
85
|
+
|
|
86
|
+
protocols = {"bino": 0, "cont": 1, "ipsi": 2}
|
|
87
|
+
|
|
88
|
+
for protocol, idx in protocols.items():
|
|
89
|
+
idx = protocols[protocol]
|
|
90
|
+
num_frames = np.cumsum(meta_data["num_frames"])
|
|
91
|
+
start_idx = num_frames[idx - 1] if idx > 0 else 0
|
|
92
|
+
end_idx = num_frames[idx]
|
|
93
|
+
|
|
94
|
+
stimuli = meta_data["Stimulus"][idx]
|
|
95
|
+
S_protocol = S[:, start_idx:end_idx]
|
|
96
|
+
|
|
97
|
+
path_stimulus = dir_data.glob(f"*_{protocol}_*").__next__()
|
|
98
|
+
ld = loadmat(path_stimulus)
|
|
99
|
+
stimulus_data = ld["runInfo"]
|
|
100
|
+
|
|
101
|
+
unique_values = get_unique_stimulus_values(stimulus_data)
|
|
102
|
+
measure_points = (
|
|
103
|
+
np.deg2rad(unique_values["phases"]),
|
|
104
|
+
np.deg2rad(unique_values["angles"]),
|
|
105
|
+
1.0 / np.deg2rad(1.0 / unique_values["cycles"]),
|
|
106
|
+
)
|
|
107
|
+
|
|
108
|
+
f = meta_data["frame_rate"]
|
|
109
|
+
spikes = get_spikes(S_protocol, f=f)
|
|
110
|
+
|
|
111
|
+
## calculate baseline rates before / after stimulus presentation
|
|
112
|
+
baseline_pre = spikes[:, : int(stimuli[0, 0])].mean(axis=1) * f
|
|
113
|
+
baseline_post = spikes[:, int(stimuli[-1, 1]) :].mean(axis=1) * f
|
|
114
|
+
|
|
115
|
+
## calculate firing maps during stimulus presentation
|
|
116
|
+
event_counts, dwelltime = calculate_firing_maps(
|
|
117
|
+
stimulus_data=stimulus_data,
|
|
118
|
+
spikes=spikes,
|
|
119
|
+
f=f,
|
|
120
|
+
dt_onset=response_onset,
|
|
121
|
+
dt_offset=response_onset / 2,
|
|
122
|
+
stimulus_frames=stimuli[:, :4],
|
|
123
|
+
collapse_repeats=True,
|
|
124
|
+
)
|
|
125
|
+
|
|
126
|
+
neurons = range(event_counts.shape[-1])
|
|
127
|
+
idx_process = np.array(neurons)
|
|
128
|
+
n_neurons = len(neurons)
|
|
129
|
+
|
|
130
|
+
process_neuron = partial(
|
|
131
|
+
run_model_comparison,
|
|
132
|
+
dwelltime=dwelltime,
|
|
133
|
+
measure_points=measure_points,
|
|
134
|
+
show_status=False,
|
|
135
|
+
)
|
|
136
|
+
# fmaps = (event_counts / dwelltime[..., np.newaxis]).transpose(3, 0, 1, 2)
|
|
137
|
+
fmaps = event_counts.transpose(3, 0, 1, 2)
|
|
138
|
+
|
|
139
|
+
batch_sz = 10 * nP
|
|
140
|
+
nBatch = n_neurons // batch_sz
|
|
141
|
+
|
|
142
|
+
print(
|
|
143
|
+
f"Processing {n_neurons} neurons in {nBatch+1} batches using {nP} processes..."
|
|
144
|
+
)
|
|
145
|
+
results = []
|
|
146
|
+
with mp.Pool(nP) as pool:
|
|
147
|
+
for i in range(nBatch + 1):
|
|
148
|
+
idx_batch = idx_process[
|
|
149
|
+
i * batch_sz : min(n_neurons, (i + 1) * batch_sz)
|
|
150
|
+
]
|
|
151
|
+
outputs = pool.map(
|
|
152
|
+
process_neuron,
|
|
153
|
+
fmaps[idx_batch, ...],
|
|
154
|
+
)
|
|
155
|
+
|
|
156
|
+
for n, entry in zip(idx_batch, outputs):
|
|
157
|
+
results.append({})
|
|
158
|
+
for model in entry.keys():
|
|
159
|
+
## store only relevant output
|
|
160
|
+
results[n][model] = {
|
|
161
|
+
"fmap": fmaps[n, ...],
|
|
162
|
+
"baseline": [baseline_pre[n], baseline_post[n]],
|
|
163
|
+
"dwelltime": dwelltime, ## is saved n_neuron times, but is not large
|
|
164
|
+
"evidence": [
|
|
165
|
+
entry[n][model].logz[-1],
|
|
166
|
+
entry[n][model].logzerr[-1],
|
|
167
|
+
],
|
|
168
|
+
"samples": entry[n][model].samples,
|
|
169
|
+
"weights": entry[n][model].importance_weights(),
|
|
170
|
+
}
|
|
171
|
+
|
|
172
|
+
fname_out = (
|
|
173
|
+
dir_analysis / f"{prefix}_{protocol}_{Path(fname).stem}{suffix}.{save_type}"
|
|
174
|
+
)
|
|
175
|
+
with h5py.File(fname_out, "w") as f:
|
|
176
|
+
for n, entry in enumerate(results):
|
|
177
|
+
grp_neuron = f.create_group(f"neuron_{n}")
|
|
178
|
+
for model in entry.keys():
|
|
179
|
+
grp_model = grp_neuron.create_group(model)
|
|
180
|
+
for key, value in entry[model].items():
|
|
181
|
+
if isinstance(value, list):
|
|
182
|
+
grp_model.create_dataset(key, data=np.array(value))
|
|
183
|
+
else:
|
|
184
|
+
grp_model.create_dataset(key, data=value)
|
|
185
|
+
print(f"Results from {n_neurons} neurons saved to {fname_out}")
|
|
186
|
+
|
|
187
|
+
|
|
188
|
+
def build_parser() -> argparse.ArgumentParser:
|
|
189
|
+
|
|
190
|
+
parser = argparse.ArgumentParser(
|
|
191
|
+
prog="neuron-detection", description="Neuron detection pipeline using CaImAn"
|
|
192
|
+
)
|
|
193
|
+
parser.add_argument(
|
|
194
|
+
"--animal",
|
|
195
|
+
type=str,
|
|
196
|
+
required=True,
|
|
197
|
+
help="Identifier for the animal being processed",
|
|
198
|
+
)
|
|
199
|
+
parser.add_argument(
|
|
200
|
+
"--session",
|
|
201
|
+
type=str,
|
|
202
|
+
required=True,
|
|
203
|
+
help="Identifier for the session being processed",
|
|
204
|
+
)
|
|
205
|
+
parser.add_argument(
|
|
206
|
+
"--fname",
|
|
207
|
+
type=str,
|
|
208
|
+
required=True,
|
|
209
|
+
help="Filename of the imaging data file",
|
|
210
|
+
)
|
|
211
|
+
parser.add_argument(
|
|
212
|
+
"--dataset",
|
|
213
|
+
type=str,
|
|
214
|
+
default="PSD95Mice_Loewel",
|
|
215
|
+
help="Name of the dataset being processed",
|
|
216
|
+
)
|
|
217
|
+
parser.add_argument(
|
|
218
|
+
"--project",
|
|
219
|
+
type=str,
|
|
220
|
+
default="orientation_selectivity",
|
|
221
|
+
help="Name of the project",
|
|
222
|
+
)
|
|
223
|
+
parser.add_argument(
|
|
224
|
+
"--spikes_key",
|
|
225
|
+
type=str,
|
|
226
|
+
default="/estimates/S_dff",
|
|
227
|
+
help="HDF5 key for the spikes data",
|
|
228
|
+
)
|
|
229
|
+
parser.add_argument(
|
|
230
|
+
"--response_onset",
|
|
231
|
+
type=float,
|
|
232
|
+
default=0.2,
|
|
233
|
+
help="Response onset time",
|
|
234
|
+
)
|
|
235
|
+
|
|
236
|
+
parser.add_argument(
|
|
237
|
+
"--prefix",
|
|
238
|
+
type=str,
|
|
239
|
+
default="results_CaImAn",
|
|
240
|
+
help="Prefix of CaImAn result files",
|
|
241
|
+
)
|
|
242
|
+
parser.add_argument(
|
|
243
|
+
"--suffix",
|
|
244
|
+
type=str,
|
|
245
|
+
default="",
|
|
246
|
+
help="Optional suffix for different runs of detection",
|
|
247
|
+
)
|
|
248
|
+
parser.add_argument(
|
|
249
|
+
"--nP",
|
|
250
|
+
type=int,
|
|
251
|
+
default=12,
|
|
252
|
+
help="Number of processes for parallel processing",
|
|
253
|
+
)
|
|
254
|
+
parser.add_argument(
|
|
255
|
+
"--save_type",
|
|
256
|
+
type=str,
|
|
257
|
+
default="hdf5",
|
|
258
|
+
help="Specifies result file type (without trailing '.'). Defaults to hdf5",
|
|
259
|
+
)
|
|
260
|
+
|
|
261
|
+
return parser
|
|
262
|
+
|
|
263
|
+
|
|
264
|
+
def main(argv: list[str] | None = None) -> int:
|
|
265
|
+
parser = build_parser()
|
|
266
|
+
args = parser.parse_args(argv)
|
|
267
|
+
process_session(**args.__dict__)
|
|
268
|
+
|
|
269
|
+
return 0
|