bmtool 0.8.0__tar.gz → 0.8.2__tar.gz
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.
- {bmtool-0.8.0 → bmtool-0.8.2}/PKG-INFO +1 -1
- bmtool-0.8.2/bmtool/analysis/netcon_reports.py +228 -0
- {bmtool-0.8.0 → bmtool-0.8.2}/bmtool/bmplot/connections.py +15 -2
- {bmtool-0.8.0 → bmtool-0.8.2}/bmtool/stimulus/core.py +126 -98
- {bmtool-0.8.0 → bmtool-0.8.2}/bmtool/util/util.py +6 -0
- {bmtool-0.8.0 → bmtool-0.8.2}/bmtool.egg-info/PKG-INFO +1 -1
- {bmtool-0.8.0 → bmtool-0.8.2}/setup.py +1 -1
- bmtool-0.8.0/bmtool/analysis/netcon_reports.py +0 -94
- {bmtool-0.8.0 → bmtool-0.8.2}/LICENSE +0 -0
- {bmtool-0.8.0 → bmtool-0.8.2}/README.md +0 -0
- {bmtool-0.8.0 → bmtool-0.8.2}/bmtool/SLURM.py +0 -0
- {bmtool-0.8.0 → bmtool-0.8.2}/bmtool/__init__.py +0 -0
- {bmtool-0.8.0 → bmtool-0.8.2}/bmtool/__main__.py +0 -0
- {bmtool-0.8.0 → bmtool-0.8.2}/bmtool/analysis/__init__.py +0 -0
- {bmtool-0.8.0 → bmtool-0.8.2}/bmtool/analysis/entrainment.py +0 -0
- {bmtool-0.8.0 → bmtool-0.8.2}/bmtool/analysis/lfp.py +0 -0
- {bmtool-0.8.0 → bmtool-0.8.2}/bmtool/analysis/spikes.py +0 -0
- {bmtool-0.8.0 → bmtool-0.8.2}/bmtool/bmplot/__init__.py +0 -0
- {bmtool-0.8.0 → bmtool-0.8.2}/bmtool/bmplot/entrainment.py +0 -0
- {bmtool-0.8.0 → bmtool-0.8.2}/bmtool/bmplot/lfp.py +0 -0
- {bmtool-0.8.0 → bmtool-0.8.2}/bmtool/bmplot/netcon_reports.py +0 -0
- {bmtool-0.8.0 → bmtool-0.8.2}/bmtool/bmplot/spikes.py +0 -0
- {bmtool-0.8.0 → bmtool-0.8.2}/bmtool/connectors.py +0 -0
- {bmtool-0.8.0 → bmtool-0.8.2}/bmtool/debug/__init__.py +0 -0
- {bmtool-0.8.0 → bmtool-0.8.2}/bmtool/debug/commands.py +0 -0
- {bmtool-0.8.0 → bmtool-0.8.2}/bmtool/debug/debug.py +0 -0
- {bmtool-0.8.0 → bmtool-0.8.2}/bmtool/graphs.py +0 -0
- {bmtool-0.8.0 → bmtool-0.8.2}/bmtool/manage.py +0 -0
- {bmtool-0.8.0 → bmtool-0.8.2}/bmtool/plot_commands.py +0 -0
- {bmtool-0.8.0 → bmtool-0.8.2}/bmtool/singlecell.py +0 -0
- {bmtool-0.8.0 → bmtool-0.8.2}/bmtool/stimulus/__init__.py +0 -0
- {bmtool-0.8.0 → bmtool-0.8.2}/bmtool/stimulus/assemblies.py +0 -0
- {bmtool-0.8.0 → bmtool-0.8.2}/bmtool/stimulus/generators.py +0 -0
- {bmtool-0.8.0 → bmtool-0.8.2}/bmtool/synapses.py +0 -0
- {bmtool-0.8.0 → bmtool-0.8.2}/bmtool/util/__init__.py +0 -0
- {bmtool-0.8.0 → bmtool-0.8.2}/bmtool/util/commands.py +0 -0
- {bmtool-0.8.0 → bmtool-0.8.2}/bmtool/util/neuron/__init__.py +0 -0
- {bmtool-0.8.0 → bmtool-0.8.2}/bmtool/util/neuron/celltuner.py +0 -0
- {bmtool-0.8.0 → bmtool-0.8.2}/bmtool.egg-info/SOURCES.txt +0 -0
- {bmtool-0.8.0 → bmtool-0.8.2}/bmtool.egg-info/dependency_links.txt +0 -0
- {bmtool-0.8.0 → bmtool-0.8.2}/bmtool.egg-info/entry_points.txt +0 -0
- {bmtool-0.8.0 → bmtool-0.8.2}/bmtool.egg-info/requires.txt +0 -0
- {bmtool-0.8.0 → bmtool-0.8.2}/bmtool.egg-info/top_level.txt +0 -0
- {bmtool-0.8.0 → bmtool-0.8.2}/pyproject.toml +0 -0
- {bmtool-0.8.0 → bmtool-0.8.2}/setup.cfg +0 -0
|
@@ -0,0 +1,228 @@
|
|
|
1
|
+
import h5py
|
|
2
|
+
import numpy as np
|
|
3
|
+
import xarray as xr
|
|
4
|
+
from typing import Union, List, Dict, Any
|
|
5
|
+
|
|
6
|
+
from ..util.util import load_nodes_from_config
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
def load_synapse_report(
|
|
10
|
+
h5_file_path: str,
|
|
11
|
+
config_path: str,
|
|
12
|
+
edge_name: str,
|
|
13
|
+
source_groupby: Union[str, List[str]],
|
|
14
|
+
target_groupby: Union[str, List[str]],
|
|
15
|
+
) -> xr.Dataset:
|
|
16
|
+
"""
|
|
17
|
+
Load and process a synapse report from a bmtk simulation into an xarray.
|
|
18
|
+
|
|
19
|
+
Parameters:
|
|
20
|
+
-----------
|
|
21
|
+
h5_file_path : str
|
|
22
|
+
Path to the h5 file containing the synapse report
|
|
23
|
+
config_path : str
|
|
24
|
+
Path to the simulation configuration file
|
|
25
|
+
edge_name : str
|
|
26
|
+
Edge name in format 'source_to_target' (e.g., 'thalamic_tone_to_LA')
|
|
27
|
+
This determines which source and target networks to load for population mapping
|
|
28
|
+
source_groupby : str or List[str]
|
|
29
|
+
Node property column name(s) to use for labeling source synapses.
|
|
30
|
+
Examples: 'pop_name', ['pop_name', 'model_type']
|
|
31
|
+
target_groupby : str or List[str]
|
|
32
|
+
Node property column name(s) to use for labeling target synapses.
|
|
33
|
+
Examples: 'pop_name', ['pop_name', 'model_type']
|
|
34
|
+
|
|
35
|
+
Returns:
|
|
36
|
+
--------
|
|
37
|
+
xarray.Dataset
|
|
38
|
+
An xarray containing the synapse report data with proper population labeling.
|
|
39
|
+
For each column in source_groupby/target_groupby, separate coordinates are created:
|
|
40
|
+
'source_{column}', 'target_{column}', etc.
|
|
41
|
+
A 'connection_label' coordinate is also created with pipe-delimited values
|
|
42
|
+
(e.g., 'Pyr|biophys->PV|biophys').
|
|
43
|
+
|
|
44
|
+
Examples:
|
|
45
|
+
---------
|
|
46
|
+
# Group by single column (default behavior):
|
|
47
|
+
ds = load_synapse_report(
|
|
48
|
+
h5_file_path='output/synapse_report.h5',
|
|
49
|
+
config_path='simulation_config.json',
|
|
50
|
+
edge_name='LA_to_LA',
|
|
51
|
+
source_groupby='pop_name',
|
|
52
|
+
target_groupby='pop_name'
|
|
53
|
+
)
|
|
54
|
+
|
|
55
|
+
# Group by multiple columns:
|
|
56
|
+
ds = load_synapse_report(
|
|
57
|
+
h5_file_path='output/synapse_report.h5',
|
|
58
|
+
config_path='simulation_config.json',
|
|
59
|
+
edge_name='LA_to_LA',
|
|
60
|
+
source_groupby=['pop_name', 'model_type'],
|
|
61
|
+
target_groupby=['pop_name', 'model_type']
|
|
62
|
+
)
|
|
63
|
+
# Returns dataset with coordinates:
|
|
64
|
+
# source_pop_name, source_model_type, target_pop_name, target_model_type, connection_label
|
|
65
|
+
"""
|
|
66
|
+
# Normalize groupby parameters to lists
|
|
67
|
+
if isinstance(source_groupby, str):
|
|
68
|
+
source_groupby = [source_groupby]
|
|
69
|
+
if isinstance(target_groupby, str):
|
|
70
|
+
target_groupby = [target_groupby]
|
|
71
|
+
|
|
72
|
+
# Parse edge_name to extract source and target networks
|
|
73
|
+
if "_to_" not in edge_name:
|
|
74
|
+
raise ValueError(
|
|
75
|
+
f"Invalid edge_name format: '{edge_name}'. Expected format: 'source_to_target' "
|
|
76
|
+
"(e.g., 'thalamic_tone_to_LA')"
|
|
77
|
+
)
|
|
78
|
+
|
|
79
|
+
source_network, target_network = edge_name.split("_to_")
|
|
80
|
+
|
|
81
|
+
# Load the h5 file to get synapse mapping data
|
|
82
|
+
with h5py.File(h5_file_path, "r") as file:
|
|
83
|
+
# Get the first (and typically only) network key in the report
|
|
84
|
+
report_networks = list(file["report"].keys())
|
|
85
|
+
if not report_networks:
|
|
86
|
+
raise ValueError(f"No report networks found in {h5_file_path}")
|
|
87
|
+
|
|
88
|
+
# Use the first available network in the h5 file
|
|
89
|
+
report_network = report_networks[0]
|
|
90
|
+
report = file["report"][report_network]
|
|
91
|
+
mapping = report["mapping"]
|
|
92
|
+
|
|
93
|
+
# Get the data - shape is (n_timesteps, n_synapses)
|
|
94
|
+
data = report["data"][:]
|
|
95
|
+
|
|
96
|
+
# Get time information
|
|
97
|
+
time_info = mapping["time"][:] # [start_time, end_time, dt]
|
|
98
|
+
start_time = time_info[0]
|
|
99
|
+
end_time = time_info[1]
|
|
100
|
+
dt = time_info[2]
|
|
101
|
+
|
|
102
|
+
# Create time array
|
|
103
|
+
n_steps = data.shape[0]
|
|
104
|
+
time = np.linspace(start_time, start_time + (n_steps - 1) * dt, n_steps)
|
|
105
|
+
|
|
106
|
+
# Get mapping information
|
|
107
|
+
src_ids = mapping["src_ids"][:]
|
|
108
|
+
trg_ids = mapping["trg_ids"][:]
|
|
109
|
+
sec_id = mapping["element_ids"][:]
|
|
110
|
+
sec_x = mapping["element_pos"][:]
|
|
111
|
+
|
|
112
|
+
# Load node information for both source and target networks
|
|
113
|
+
all_nodes = load_nodes_from_config(config_path)
|
|
114
|
+
|
|
115
|
+
# Get the source and target node dataframes
|
|
116
|
+
if source_network not in all_nodes:
|
|
117
|
+
raise ValueError(
|
|
118
|
+
f"Source network '{source_network}' not found in config. "
|
|
119
|
+
f"Available networks: {list(all_nodes.keys())}"
|
|
120
|
+
)
|
|
121
|
+
if target_network not in all_nodes:
|
|
122
|
+
raise ValueError(
|
|
123
|
+
f"Target network '{target_network}' not found in config. "
|
|
124
|
+
f"Available networks: {list(all_nodes.keys())}"
|
|
125
|
+
)
|
|
126
|
+
|
|
127
|
+
source_nodes = all_nodes[source_network]
|
|
128
|
+
target_nodes = all_nodes[target_network]
|
|
129
|
+
|
|
130
|
+
# Validate that requested groupby columns exist in node dataframes
|
|
131
|
+
missing_src_cols = [col for col in source_groupby if col not in source_nodes.columns]
|
|
132
|
+
if missing_src_cols:
|
|
133
|
+
raise KeyError(
|
|
134
|
+
f"Columns {missing_src_cols} not found in source network '{source_network}'. "
|
|
135
|
+
f"Available columns: {list(source_nodes.columns)}"
|
|
136
|
+
)
|
|
137
|
+
|
|
138
|
+
missing_trg_cols = [col for col in target_groupby if col not in target_nodes.columns]
|
|
139
|
+
if missing_trg_cols:
|
|
140
|
+
raise KeyError(
|
|
141
|
+
f"Columns {missing_trg_cols} not found in target network '{target_network}'. "
|
|
142
|
+
f"Available columns: {list(target_nodes.columns)}"
|
|
143
|
+
)
|
|
144
|
+
|
|
145
|
+
# Create mappings from node IDs to groupby column values
|
|
146
|
+
# source_mappings[col] = {node_id: column_value}
|
|
147
|
+
source_mappings: Dict[str, Dict[int, Any]] = {}
|
|
148
|
+
for col in source_groupby:
|
|
149
|
+
source_mappings[col] = dict(zip(source_nodes.index, source_nodes[col]))
|
|
150
|
+
|
|
151
|
+
target_mappings: Dict[str, Dict[int, Any]] = {}
|
|
152
|
+
for col in target_groupby:
|
|
153
|
+
target_mappings[col] = dict(zip(target_nodes.index, target_nodes[col]))
|
|
154
|
+
|
|
155
|
+
# Determine default values for external inputs (src_id = -1)
|
|
156
|
+
# Get the most common value or "unknown" if heterogeneous
|
|
157
|
+
src_external_values = {}
|
|
158
|
+
for col in source_groupby:
|
|
159
|
+
unique_vals = source_nodes[col].unique()
|
|
160
|
+
if len(unique_vals) == 1:
|
|
161
|
+
src_external_values[col] = unique_vals[0]
|
|
162
|
+
else:
|
|
163
|
+
src_external_values[col] = "unknown"
|
|
164
|
+
|
|
165
|
+
# Get the number of synapses
|
|
166
|
+
n_synapses = data.shape[1]
|
|
167
|
+
|
|
168
|
+
# Create arrays to hold the groupby values for each synapse
|
|
169
|
+
# synapse_values[col] = [val_for_synapse_0, val_for_synapse_1, ...]
|
|
170
|
+
source_values: Dict[str, List[Any]] = {col: [] for col in source_groupby}
|
|
171
|
+
target_values: Dict[str, List[Any]] = {col: [] for col in target_groupby}
|
|
172
|
+
connection_labels = []
|
|
173
|
+
|
|
174
|
+
# Process each synapse
|
|
175
|
+
for i in range(n_synapses):
|
|
176
|
+
src_id = src_ids[i]
|
|
177
|
+
trg_id = trg_ids[i]
|
|
178
|
+
|
|
179
|
+
# Get source groupby values
|
|
180
|
+
src_label_parts = []
|
|
181
|
+
for col in source_groupby:
|
|
182
|
+
if src_id == -1:
|
|
183
|
+
# External input: use default value for this column
|
|
184
|
+
val = src_external_values[col]
|
|
185
|
+
else:
|
|
186
|
+
val = source_mappings[col].get(src_id, f"unknown_{src_id}")
|
|
187
|
+
source_values[col].append(val)
|
|
188
|
+
src_label_parts.append(str(val))
|
|
189
|
+
|
|
190
|
+
# Get target groupby values
|
|
191
|
+
trg_label_parts = []
|
|
192
|
+
for col in target_groupby:
|
|
193
|
+
val = target_mappings[col].get(trg_id, f"unknown_{trg_id}")
|
|
194
|
+
target_values[col].append(val)
|
|
195
|
+
trg_label_parts.append(str(val))
|
|
196
|
+
|
|
197
|
+
# Create connection label with pipe-delimited format
|
|
198
|
+
src_label = "|".join(src_label_parts)
|
|
199
|
+
trg_label = "|".join(trg_label_parts)
|
|
200
|
+
connection_labels.append(f"{src_label}->{trg_label}")
|
|
201
|
+
|
|
202
|
+
# Create coordinates dictionary dynamically based on groupby columns
|
|
203
|
+
coords = {
|
|
204
|
+
"time": time,
|
|
205
|
+
"synapse": np.arange(n_synapses),
|
|
206
|
+
"source_id": ("synapse", src_ids),
|
|
207
|
+
"target_id": ("synapse", trg_ids),
|
|
208
|
+
"sec_id": ("synapse", sec_id),
|
|
209
|
+
"sec_x": ("synapse", sec_x),
|
|
210
|
+
"connection_label": ("synapse", connection_labels),
|
|
211
|
+
}
|
|
212
|
+
|
|
213
|
+
# Add source groupby coordinates
|
|
214
|
+
for col in source_groupby:
|
|
215
|
+
coords[f"source_{col}"] = ("synapse", source_values[col])
|
|
216
|
+
|
|
217
|
+
# Add target groupby coordinates
|
|
218
|
+
for col in target_groupby:
|
|
219
|
+
coords[f"target_{col}"] = ("synapse", target_values[col])
|
|
220
|
+
|
|
221
|
+
# Create xarray dataset
|
|
222
|
+
ds = xr.Dataset(
|
|
223
|
+
data_vars={"synapse_value": (["time", "synapse"], data)},
|
|
224
|
+
coords=coords,
|
|
225
|
+
attrs={"description": "Synapse report data from bmtk simulation"},
|
|
226
|
+
)
|
|
227
|
+
|
|
228
|
+
return ds
|
|
@@ -193,7 +193,11 @@ def percent_connection_matrix(
|
|
|
193
193
|
no_prepend_pop : bool, optional
|
|
194
194
|
If True, population name is not displayed before sid or tid in the plot. Default is False.
|
|
195
195
|
method : str, optional
|
|
196
|
-
Method for calculating percent connectivity. Options:
|
|
196
|
+
Method for calculating percent connectivity. Options:
|
|
197
|
+
- 'total': Pairwise connection probability (connections / (sources * targets) * 100)
|
|
198
|
+
- 'uni': Unidirectional connection probability
|
|
199
|
+
- 'bi': Bidirectional connection probability
|
|
200
|
+
- 'innervation': Percentage of target neurons receiving at least one connection
|
|
197
201
|
Default is 'total'.
|
|
198
202
|
include_gap : bool, optional
|
|
199
203
|
If True, include gap junctions in analysis. If False, only include chemical synapses.
|
|
@@ -249,7 +253,16 @@ def percent_connection_matrix(
|
|
|
249
253
|
include_gap=include_gap,
|
|
250
254
|
)
|
|
251
255
|
if title is None or title == "":
|
|
252
|
-
|
|
256
|
+
if method == "uni":
|
|
257
|
+
title = "Unidirectional Percent Connectivity"
|
|
258
|
+
elif method == "bi":
|
|
259
|
+
title = "Bidirectional Percent Connectivity"
|
|
260
|
+
elif method == "total":
|
|
261
|
+
title = "Total Percent Connectivity"
|
|
262
|
+
elif method == "innervation":
|
|
263
|
+
title = "Percentage of target neurons receiving at least one connection"
|
|
264
|
+
else:
|
|
265
|
+
title = "Percent Connectivity"
|
|
253
266
|
|
|
254
267
|
if return_dict:
|
|
255
268
|
result_dict = plot_connection_info(
|
|
@@ -140,16 +140,129 @@ class StimulusBuilder:
|
|
|
140
140
|
raise ValueError(f"Unknown distribution: {distribution}. Must be 'lognormal' or 'normal'.")
|
|
141
141
|
|
|
142
142
|
return rates
|
|
143
|
+
|
|
144
|
+
def generate_background(self, output_path, network_name, population_params,
|
|
145
|
+
groupby='pop_name', t_start=0.0, t_stop=10.0,
|
|
146
|
+
verbose=False, seed=None):
|
|
147
|
+
"""Generate background (spontaneous) activity for network nodes grouped by property.
|
|
148
|
+
|
|
149
|
+
This function generates baseline spiking activity, grouped by a specified node property.
|
|
150
|
+
Each group can use either a constant firing rate or a distribution-based rate.
|
|
151
|
+
|
|
152
|
+
Args:
|
|
153
|
+
output_path (str): Path to save the resulting .h5 file.
|
|
154
|
+
network_name (str): BMTK network name.
|
|
155
|
+
population_params (dict): Parameters for each population/group.
|
|
156
|
+
Keys should match values in the node property specified by groupby.
|
|
157
|
+
Each value is a dict with:
|
|
158
|
+
- 'mean_firing_rate' (float): Mean firing rate in Hz (required)
|
|
159
|
+
- 'stdev' (float, optional): Standard deviation. If provided, uses lognormal distribution.
|
|
160
|
+
If omitted, uses constant firing rate.
|
|
161
|
+
Example:
|
|
162
|
+
{
|
|
163
|
+
'PN': {'mean_firing_rate': 20.0, 'stdev': 2.0},
|
|
164
|
+
'PV': {'mean_firing_rate': 30.0}, # constant rate
|
|
165
|
+
'SST': {'mean_firing_rate': 15.0, 'stdev': 1.5}
|
|
166
|
+
}
|
|
167
|
+
groupby (str): Node property to group by (default: 'pop_name').
|
|
168
|
+
Will match against keys in population_params.
|
|
169
|
+
t_start, t_stop (float): Time range for activity (seconds).
|
|
170
|
+
verbose (bool): If True, print detailed information (default: False).
|
|
171
|
+
seed (int, optional): Random seed for distribution sampling. Overrides instance psg_seed.
|
|
172
|
+
|
|
173
|
+
Examples:
|
|
174
|
+
# Population-specific rates with mixed distributions
|
|
175
|
+
params = {
|
|
176
|
+
'PN': {'mean_firing_rate': 20.0, 'stdev': 2.0},
|
|
177
|
+
'PV': {'mean_firing_rate': 30.0}, # constant rate
|
|
178
|
+
'SST': {'mean_firing_rate': 15.0, 'stdev': 1.5}
|
|
179
|
+
}
|
|
180
|
+
sb.generate_background(
|
|
181
|
+
output_path='background.h5',
|
|
182
|
+
network_name='input',
|
|
183
|
+
population_params=params,
|
|
184
|
+
t_start=0.0, t_stop=15.0
|
|
185
|
+
)
|
|
186
|
+
|
|
187
|
+
# Group by custom property (e.g., layer)
|
|
188
|
+
layer_params = {
|
|
189
|
+
'L1': {'mean_firing_rate': 10.0, 'stdev': 1.0},
|
|
190
|
+
'L2/3': {'mean_firing_rate': 15.0, 'stdev': 2.0}
|
|
191
|
+
}
|
|
192
|
+
sb.generate_background(
|
|
193
|
+
output_path='layer_background.h5',
|
|
194
|
+
network_name='input',
|
|
195
|
+
population_params=layer_params,
|
|
196
|
+
groupby='layer'
|
|
197
|
+
)
|
|
198
|
+
"""
|
|
199
|
+
if population_params is None or not isinstance(population_params, dict):
|
|
200
|
+
raise ValueError("population_params must be a non-empty dict")
|
|
201
|
+
|
|
202
|
+
nodes_df = self.get_nodes(network_name)
|
|
203
|
+
|
|
204
|
+
# Verify groupby column exists
|
|
205
|
+
if groupby not in nodes_df.columns:
|
|
206
|
+
raise ValueError(f"Node property '{groupby}' not found in network '{network_name}'")
|
|
207
|
+
|
|
208
|
+
# Use provided seed or default to instance psg_seed
|
|
209
|
+
psg_seed = seed if seed is not None else self.psg_seed
|
|
210
|
+
|
|
211
|
+
population = network_name # Default population name in PSG
|
|
212
|
+
psg = PoissonSpikeGenerator(population=population, seed=psg_seed)
|
|
213
|
+
|
|
214
|
+
times = (t_start, t_stop)
|
|
215
|
+
total_nodes = 0
|
|
216
|
+
|
|
217
|
+
for group_key, params in population_params.items():
|
|
218
|
+
# Find nodes matching this group
|
|
219
|
+
nodes_in_group = nodes_df[nodes_df[groupby] == group_key].index.values
|
|
220
|
+
|
|
221
|
+
if len(nodes_in_group) == 0:
|
|
222
|
+
if verbose:
|
|
223
|
+
print(f" Warning: No nodes found with {groupby}='{group_key}'")
|
|
224
|
+
continue
|
|
225
|
+
|
|
226
|
+
total_nodes += len(nodes_in_group)
|
|
227
|
+
|
|
228
|
+
if not isinstance(params, dict) or 'mean_firing_rate' not in params:
|
|
229
|
+
raise ValueError(f"params['{group_key}'] must be a dict with 'mean_firing_rate' key")
|
|
230
|
+
|
|
231
|
+
mean_rate = params['mean_firing_rate']
|
|
232
|
+
stdev = params.get('stdev', None)
|
|
233
|
+
|
|
234
|
+
# Determine: constant vs distribution-based
|
|
235
|
+
if stdev is not None:
|
|
236
|
+
# Use distribution (lognormal)
|
|
237
|
+
firing_rates = self._generate_firing_rates(len(nodes_in_group), mean_rate, stdev, 'lognormal')
|
|
238
|
+
for node_id, rate in zip(nodes_in_group, firing_rates):
|
|
239
|
+
psg.add(node_ids=node_id, firing_rate=rate, times=times)
|
|
240
|
+
if verbose:
|
|
241
|
+
print(f" {group_key}: {len(nodes_in_group)} nodes, {mean_rate:.1f}±{stdev:.1f} Hz (lognormal)")
|
|
242
|
+
else:
|
|
243
|
+
# Use constant firing rate
|
|
244
|
+
psg.add(node_ids=nodes_in_group.tolist(), firing_rate=mean_rate, times=times)
|
|
245
|
+
if verbose:
|
|
246
|
+
print(f" {group_key}: {len(nodes_in_group)} nodes, {mean_rate:.1f} Hz (constant)")
|
|
247
|
+
|
|
248
|
+
# Write to file
|
|
249
|
+
os.makedirs(os.path.dirname(output_path), exist_ok=True)
|
|
250
|
+
psg.to_sonata(output_path)
|
|
251
|
+
if verbose:
|
|
252
|
+
print(f"Generated background activity: {total_nodes} nodes to {output_path}")
|
|
143
253
|
|
|
144
254
|
def generate_stimulus(self, output_path, pattern_type, assembly_name, verbose=False, seed=None, **kwargs):
|
|
145
255
|
"""Generate a BMTK Poisson spike file (SONATA) for a specific assembly group.
|
|
146
256
|
|
|
257
|
+
Use create_assemblies() first to define your stimulus assemblies, then call this
|
|
258
|
+
function to generate time-varying firing patterns for those assemblies.
|
|
259
|
+
|
|
147
260
|
Args:
|
|
148
261
|
output_path (str): Path to save the resulting .h5 file.
|
|
149
262
|
pattern_type (str): Firing rate template ('short', 'long', 'ramp', etc).
|
|
150
263
|
assembly_name (str): Name of the assembly group created via create_assemblies.
|
|
151
264
|
verbose (bool): If True, print detailed information (default: False).
|
|
152
|
-
seed (int, optional): Random seed for Poisson spike generation. Overrides instance psg_seed
|
|
265
|
+
seed (int, optional): Random seed for Poisson spike generation. Overrides instance psg_seed.
|
|
153
266
|
**kwargs: Arguments passed to the generator function and PoissonSpikeGenerator.
|
|
154
267
|
- population (str): Name of the spike population (for BMTK).
|
|
155
268
|
- firing_rate (3-tuple): (off_rate, burst_rate, silent_rate).
|
|
@@ -157,9 +270,20 @@ class StimulusBuilder:
|
|
|
157
270
|
- off_time (float): Duration of silent period.
|
|
158
271
|
- t_start (float): Start time of cycles.
|
|
159
272
|
- t_stop (float): End time of cycles.
|
|
273
|
+
|
|
274
|
+
Example:
|
|
275
|
+
# First create assemblies
|
|
276
|
+
sb.create_assemblies(name='stim_groups', network_name='thalamus',
|
|
277
|
+
method='property', property_name='pulse_group_id')
|
|
278
|
+
|
|
279
|
+
# Then generate stimulus
|
|
280
|
+
sb.generate_stimulus(output_path='stim.h5', pattern_type='long',
|
|
281
|
+
assembly_name='stim_groups', population='thalamus',
|
|
282
|
+
firing_rate=(0.0, 50.0, 0.0), t_start=1.0, t_stop=15.0,
|
|
283
|
+
on_time=1.0, off_time=0.5)
|
|
160
284
|
"""
|
|
161
285
|
if assembly_name not in self.assemblies:
|
|
162
|
-
raise ValueError(f"Assembly '{assembly_name}' not defined.")
|
|
286
|
+
raise ValueError(f"Assembly '{assembly_name}' not defined. Use create_assemblies() first.")
|
|
163
287
|
|
|
164
288
|
assembly_list = self.assemblies[assembly_name]
|
|
165
289
|
n_assemblies = len(assembly_list)
|
|
@@ -192,99 +316,3 @@ class StimulusBuilder:
|
|
|
192
316
|
psg.to_sonata(output_path)
|
|
193
317
|
if verbose:
|
|
194
318
|
print(f"Written stimulus to {output_path}")
|
|
195
|
-
|
|
196
|
-
def generate_baseline(self, output_path, network_name, pop_name=None, distribution='constant',
|
|
197
|
-
mean=None, stdev=None, firing_rate=None, t_start=0.0, t_stop=10.0, verbose=False, seed=None):
|
|
198
|
-
"""Generate baseline activity for a selection of nodes.
|
|
199
|
-
|
|
200
|
-
Args:
|
|
201
|
-
output_path (str): Path to save the resulting .h5 file.
|
|
202
|
-
network_name (str): BMTK network name.
|
|
203
|
-
pop_name (str, optional): Filter nodes by population name.
|
|
204
|
-
distribution (str): 'constant', 'lognormal', or 'normal'.
|
|
205
|
-
mean (float): Mean for lognormal/normal or constant rate (if firing_rate omitted).
|
|
206
|
-
stdev (float): Standard deviation for lognormal/normal.
|
|
207
|
-
firing_rate (float, optional): Constant firing rate.
|
|
208
|
-
t_start, t_stop (float): Time range for activity.
|
|
209
|
-
verbose (bool): If True, print detailed information (default: False).
|
|
210
|
-
seed (int, optional): Random seed for distribution sampling. Overrides instance psg_seed for this call.
|
|
211
|
-
"""
|
|
212
|
-
nodes_df = self.get_nodes(network_name, pop_name)
|
|
213
|
-
node_ids = nodes_df.index.values.tolist()
|
|
214
|
-
|
|
215
|
-
# Use provided seed or default to instance psg_seed
|
|
216
|
-
psg_seed = seed if seed is not None else self.psg_seed
|
|
217
|
-
|
|
218
|
-
population = network_name # Default population name in PSG
|
|
219
|
-
psg = PoissonSpikeGenerator(population=population, seed=psg_seed)
|
|
220
|
-
|
|
221
|
-
times = (t_start, t_stop)
|
|
222
|
-
|
|
223
|
-
if distribution == 'constant':
|
|
224
|
-
if firing_rate is None:
|
|
225
|
-
if mean is not None:
|
|
226
|
-
firing_rate = mean
|
|
227
|
-
else:
|
|
228
|
-
raise ValueError("Must provide firing_rate for constant distribution")
|
|
229
|
-
|
|
230
|
-
psg.add(node_ids=node_ids, firing_rate=firing_rate, times=times)
|
|
231
|
-
|
|
232
|
-
elif distribution in ['lognormal', 'normal']:
|
|
233
|
-
if mean is None or stdev is None:
|
|
234
|
-
raise ValueError(f"Must provide mean and stdev for {distribution} distribution")
|
|
235
|
-
|
|
236
|
-
firing_rates = self._generate_firing_rates(len(node_ids), mean, stdev, distribution)
|
|
237
|
-
|
|
238
|
-
for node_id, fr in zip(node_ids, firing_rates):
|
|
239
|
-
psg.add(node_ids=node_id, firing_rate=fr, times=times)
|
|
240
|
-
|
|
241
|
-
else:
|
|
242
|
-
raise ValueError(f"Unknown distribution: {distribution}")
|
|
243
|
-
|
|
244
|
-
# Write to file
|
|
245
|
-
os.makedirs(os.path.dirname(output_path), exist_ok=True)
|
|
246
|
-
psg.to_sonata(output_path)
|
|
247
|
-
if verbose:
|
|
248
|
-
print(f"Written baseline to {output_path}")
|
|
249
|
-
|
|
250
|
-
def generate_shell_input(self, output_path, network_name, shell_params,
|
|
251
|
-
distribution='lognormal', t_start=0.0, t_stop=15.0, verbose=False, seed=None):
|
|
252
|
-
"""Generate shell (background) stimulus with population-specific rates.
|
|
253
|
-
|
|
254
|
-
Args:
|
|
255
|
-
output_path (str): Path to save the resulting .h5 file.
|
|
256
|
-
network_name (str): BMTK network name.
|
|
257
|
-
shell_params (dict): Population-specific (mean, stdev) tuples.
|
|
258
|
-
Example: {'ET': (1.9, 1.8), 'IT': (1.3, 1.4), 'PV': (7.5, 6.4), 'SST': (5.0, 6.0)}
|
|
259
|
-
distribution (str): 'lognormal' or 'normal' (default: 'lognormal').
|
|
260
|
-
t_start, t_stop (float): Time range for activity.
|
|
261
|
-
verbose (bool): If True, print detailed information (default: False).
|
|
262
|
-
seed (int, optional): Random seed for distribution sampling. Overrides instance psg_seed for this call.
|
|
263
|
-
"""
|
|
264
|
-
nodes_df = self.get_nodes(network_name)
|
|
265
|
-
|
|
266
|
-
# Use provided seed or default to instance psg_seed
|
|
267
|
-
psg_seed = seed if seed is not None else self.psg_seed
|
|
268
|
-
|
|
269
|
-
psg = PoissonSpikeGenerator(population=network_name, seed=psg_seed)
|
|
270
|
-
|
|
271
|
-
total_nodes = 0
|
|
272
|
-
for pop_name, (mean, stdev) in shell_params.items():
|
|
273
|
-
nodes_in_pop = nodes_df[nodes_df['pop_name'] == pop_name].index.values
|
|
274
|
-
if len(nodes_in_pop) == 0:
|
|
275
|
-
continue
|
|
276
|
-
|
|
277
|
-
total_nodes += len(nodes_in_pop)
|
|
278
|
-
|
|
279
|
-
# Generate rates using helper function
|
|
280
|
-
rates = self._generate_firing_rates(len(nodes_in_pop), mean, stdev, distribution)
|
|
281
|
-
|
|
282
|
-
# Add to PSG
|
|
283
|
-
for node_id, rate in zip(nodes_in_pop, rates):
|
|
284
|
-
psg.add(node_ids=node_id, firing_rate=rate, times=(t_start, t_stop))
|
|
285
|
-
|
|
286
|
-
# Write to file
|
|
287
|
-
os.makedirs(os.path.dirname(output_path), exist_ok=True)
|
|
288
|
-
psg.to_sonata(output_path)
|
|
289
|
-
if verbose:
|
|
290
|
-
print(f"Generated shell stimulus ({distribution}): {total_nodes} nodes to {output_path}")
|
|
@@ -1240,12 +1240,18 @@ def percent_connections(
|
|
|
1240
1240
|
uni = round(num_uni / (num_sources * num_targets) * 100, 2)
|
|
1241
1241
|
bi = round(num_bi / (num_sources * num_targets) * 100, 2)
|
|
1242
1242
|
|
|
1243
|
+
# Calculate innervation percentage (% of target neurons receiving at least one connection)
|
|
1244
|
+
unique_target_nodes = cons['target_node_id'].nunique()
|
|
1245
|
+
innervation = round(unique_target_nodes / num_targets * 100, 2)
|
|
1246
|
+
|
|
1243
1247
|
if method == "total":
|
|
1244
1248
|
return total
|
|
1245
1249
|
if method == "uni":
|
|
1246
1250
|
return uni
|
|
1247
1251
|
if method == "bi":
|
|
1248
1252
|
return bi
|
|
1253
|
+
if method == "innervation":
|
|
1254
|
+
return innervation
|
|
1249
1255
|
|
|
1250
1256
|
return relation_matrix(
|
|
1251
1257
|
config, nodes, edges, sources, targets, sids, tids, prepend_pop, relation_func=precent_func
|
|
@@ -1,94 +0,0 @@
|
|
|
1
|
-
import h5py
|
|
2
|
-
import numpy as np
|
|
3
|
-
import xarray as xr
|
|
4
|
-
|
|
5
|
-
from ..util.util import load_nodes_from_config
|
|
6
|
-
|
|
7
|
-
|
|
8
|
-
def load_synapse_report(h5_file_path, config_path, network):
|
|
9
|
-
"""
|
|
10
|
-
Load and process a synapse report from a bmtk simulation into an xarray.
|
|
11
|
-
|
|
12
|
-
Parameters:
|
|
13
|
-
-----------
|
|
14
|
-
h5_file_path : str
|
|
15
|
-
Path to the h5 file containing the synapse report
|
|
16
|
-
config_path : str
|
|
17
|
-
Path to the simulation configuration file
|
|
18
|
-
|
|
19
|
-
Returns:
|
|
20
|
-
--------
|
|
21
|
-
xarray.Dataset
|
|
22
|
-
An xarray containing the synapse report data with proper population labeling
|
|
23
|
-
"""
|
|
24
|
-
# Load the h5 file
|
|
25
|
-
with h5py.File(h5_file_path, "r") as file:
|
|
26
|
-
# Get the report data
|
|
27
|
-
report = file["report"][network]
|
|
28
|
-
mapping = report["mapping"]
|
|
29
|
-
|
|
30
|
-
# Get the data - shape is (n_timesteps, n_synapses)
|
|
31
|
-
data = report["data"][:]
|
|
32
|
-
|
|
33
|
-
# Get time information
|
|
34
|
-
time_info = mapping["time"][:] # [start_time, end_time, dt]
|
|
35
|
-
start_time = time_info[0]
|
|
36
|
-
end_time = time_info[1]
|
|
37
|
-
dt = time_info[2]
|
|
38
|
-
|
|
39
|
-
# Create time array
|
|
40
|
-
n_steps = data.shape[0]
|
|
41
|
-
time = np.linspace(start_time, start_time + (n_steps - 1) * dt, n_steps)
|
|
42
|
-
|
|
43
|
-
# Get mapping information
|
|
44
|
-
src_ids = mapping["src_ids"][:]
|
|
45
|
-
trg_ids = mapping["trg_ids"][:]
|
|
46
|
-
sec_id = mapping["element_ids"][:]
|
|
47
|
-
sec_x = mapping["element_pos"][:]
|
|
48
|
-
|
|
49
|
-
# Load node information
|
|
50
|
-
nodes = load_nodes_from_config(config_path)
|
|
51
|
-
nodes = nodes[network]
|
|
52
|
-
|
|
53
|
-
# Create a mapping from node IDs to population names
|
|
54
|
-
node_to_pop = dict(zip(nodes.index, nodes["pop_name"]))
|
|
55
|
-
|
|
56
|
-
# Get the number of synapses
|
|
57
|
-
n_synapses = data.shape[1]
|
|
58
|
-
|
|
59
|
-
# Create arrays to hold the source and target populations for each synapse
|
|
60
|
-
source_pops = []
|
|
61
|
-
target_pops = []
|
|
62
|
-
connection_labels = []
|
|
63
|
-
|
|
64
|
-
# Process each synapse
|
|
65
|
-
for i in range(n_synapses):
|
|
66
|
-
src_id = src_ids[i]
|
|
67
|
-
trg_id = trg_ids[i]
|
|
68
|
-
|
|
69
|
-
# Get population names (with fallback for unknown IDs)
|
|
70
|
-
src_pop = node_to_pop.get(src_id, f"unknown_{src_id}")
|
|
71
|
-
trg_pop = node_to_pop.get(trg_id, f"unknown_{trg_id}")
|
|
72
|
-
|
|
73
|
-
source_pops.append(src_pop)
|
|
74
|
-
target_pops.append(trg_pop)
|
|
75
|
-
connection_labels.append(f"{src_pop}->{trg_pop}")
|
|
76
|
-
|
|
77
|
-
# Create xarray dataset
|
|
78
|
-
ds = xr.Dataset(
|
|
79
|
-
data_vars={"synapse_value": (["time", "synapse"], data)},
|
|
80
|
-
coords={
|
|
81
|
-
"time": time,
|
|
82
|
-
"synapse": np.arange(n_synapses),
|
|
83
|
-
"source_pop": ("synapse", source_pops),
|
|
84
|
-
"target_pop": ("synapse", target_pops),
|
|
85
|
-
"source_id": ("synapse", src_ids),
|
|
86
|
-
"target_id": ("synapse", trg_ids),
|
|
87
|
-
"sec_id": ("synapse", sec_id),
|
|
88
|
-
"sec_x": ("synapse", sec_x),
|
|
89
|
-
"connection_label": ("synapse", connection_labels),
|
|
90
|
-
},
|
|
91
|
-
attrs={"description": "Synapse report data from bmtk simulation"},
|
|
92
|
-
)
|
|
93
|
-
|
|
94
|
-
return ds
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|