OSCR

Spiking neural network decoders of finger forces from high-density intramuscular microelectrode arrays.

Code ↔ Paper

14 matches between paragraphs of the paper and lines of its authors' code, computed by the harvester (lexical-v1). Click a colored paragraph or line to see its counterpart.

The 14 matches · 1 of them tie a paragraph to a whole file, not to given lines: a weak match, whose lines are not tinted
  1. [1] § Methods › Decoders computational and resource costs ↔ force_regression/models/snn.py, lines 1196–1240 · score 0.80 · multiply accumulate, spike representations, spiking activity, encoding layer, MAC, sparse
  2. [2] § Methods › Baseline regression models › Ordinary least-squares linear regression ↔ notebooks/baselines/BaselineRegressionMUCount.ipynb, lines 63–88 · score 0.79 · linear regression model, low pass filtered, movement direction, force profile, sweep, smooth
  3. [3] § Methods › Spiking neural network models in simulation › Training and inference ↔ force_regression/models/snn.py, lines 911–986 · score 0.64 · Adam optimizer, squared error, gradient, loss, models, trained
  4. [4] § Methods › Spiking neural network models in simulation › Readout configurations ↔ force_regression/models/snn.py, lines 189–328 · score 0.64 · neurons connected, membrane potentials, leaky integrate, weight, synaptic, synapses
  5. [5] § Results › Streamed processing of motor units' spike trains with neurons and synapses ↔ notebooks/baselines/BaselineRegressionMUCount.ipynb, lines 63–88 · score 0.60 · linear regression model, cross validation, movement direction, S1, flexion, baseline
  6. [6] § Results › Networks resilience to real-time decomposition errors ↔ notebooks/analysis/NoiseResults.ipynb, lines 258–340 · score 0.60 · baseline LR, SNN LI, adding spikes, misattribution, 20 %, RMSE
  7. [7] § Methods › Decoders footprint ↔ force_regression/models/snn.py, lines 1748–1791 · score 0.59 · static metrics, connection sparsity, footprint, memory, network, models
  8. [8] § Methods › Baseline regression models › Recurrent neural networks ↔ force_regression/models/rnn.py, lines 24–140 · score 0.58 · exponential kernel, hidden, trainable, bias, vanilla, decay
  9. [9] § Methods › Decoders footprint ↔ force_regression/models/linear_regression.py, lines 267–315 · score 0.56 · static metrics, linear regression, bytes, footprint, memory, models
  10. [10] § Methods › EMG-to-spike encoding › Adaptive leaky integrate-and-fire neurons ↔ configs/constants.py, lines 55–115 · score 0.54 · effective threshold, ALIF neurons, adaptive, adaptation, spiking
  11. [11] § Methods › Spiking neural network on asynchronous neuromorphic hardware › Chip inference latency ↔ force_regression/evaluation/profiling.py, lines 350–398 · score 0.52 · Inference latency, wall clock, metric
  12. [12] § Results › Encoding HD-iEMG into events bypassing motor unit decomposition ↔ configs/constants.py, lines 55–115 · score 0.52 · ALIF neurons, encoding neurons, adaptive, adaptation, membrane, flexion
  13. [13] § Methods › Spiking neural network models in simulation › Readout configurations ↔ force_regression/models/snn.py, lines 1282–1324 · score 0.52 · membrane dynamics, membrane potential, weight, synaptic, modeled, neurons
  14. [14] § Methods › Baseline regression models › Recurrent neural networks ↔ force_regression/config/snnconfig.py, the whole file · a weak match · score 0.50 · Spiking neural network, recurrent, kernel, bias, layer, configured

Paper

Loaded from Europe PMC by your browser, not stored by OSCR: doi.org · Europe PMC

The paper is loaded when this pane is shown.

The authors' code

Python · 2,229 lines · 106 KB · no license · 5 matches

  1. from __future__ import annotations
  2. from typing import Dict, Tuple
  3. import logging
  4. import os
  5. import pickle as pkl
  6. import random
  7. import time
  8. import warnings
  9. from neurobench.benchmarks import Benchmark
  10. from neurobench.metrics.static import (
  11. Footprint,
  12. ConnectionSparsity,
  13. )
  14. from neurobench.metrics.workload import (
  15. ActivationSparsity,
  16. SynapticOperations,
  17. ClassificationAccuracy
  18. )
  19. from neurobench.models import SNNTorchModel
  20. from scipy import stats
  21. from sklearn.metrics import mean_squared_error, r2_score
  22. from snntorch import surrogate
  23. from tabulate import tabulate
  24. from torch.profiler import profile, ProfilerActivity
  25. from torch.utils.data import DataLoader
  26. from torch.utils.data import Dataset
  27. from torch.utils.data import Subset
  28. from torchsummary import summary
  29. import neurobench
  30. import numpy as np
  31. import pandas as pd
  32. import snntorch._neurons as snn
  33. import torch
  34. import torch.nn as nn
  35. import tqdm
  36. import wandb
  37. from configs.constants import *
  38. from force_regression.config.dataconfig import DataConfig
  39. from force_regression.config.snnconfig import SNNConfig
  40. from force_regression.data.preprocessing.emg import apply_butter_lowpass
  41. from force_regression.evaluation.metrics import prepare_snn_metrics_df
  42. from force_regression.training.snn_pipeline import assign_datasets_for_train_test, conc_var_tracker, target_tracker_to_np, create_subsets, reconstruct_rep_from_bins
  43. from force_regression.training.snn_pipeline import map_tau_to_decay
  44. from force_regression.training.snn_pipeline import split_tensors_for_rep
  45. import force_regression.plotting.snn_plot as snnplot
  46. import force_regression.data.preprocessing.spikes as spk
  47. import force_regression.utils.functions as fn
  48. warnings.simplefilter("ignore")
  49. DEVICE = 'cpu'
  50. DTYPE = torch.float
  51. # --- SNN Topologies ---
  52. class SnnTopology(torch.nn.Module):
  53. """
  54. Base class for SNN topologies
  55. """
  56. def __init__(self, timesteps, num_inputs, num_outputs, snn_config, use_wandb, dt):
  57. super(SnnTopology, self).__init__()
  58. self.use_wandb = use_wandb
  59. self.dt = dt
  60. self.timesteps = timesteps
  61. self.num_inputs = num_inputs
  62. self.num_outputs = num_outputs
  63. self.n_electrodes = 3
  64. self.n_channels_per_electrode = 40
  65. self.topology = snn_config["topology"]
  66. self.snn_config = snn_config
  67. self.spike_grad = surrogate.fast_sigmoid()
  68. self.param_std = 0.05 # standard deviation for parameter initialization
  69. self.leaky_integrator_threshold = 300 # spiking threshold for the leaky integrator
  70. self.loss_function = torch.nn.MSELoss()
  71. self.state_vars = self._topology_state_variables()
  72. self.plot_syn_cur, self.plot_neu_cur, self.plot_neu_mem = self._plot_network_variables()
  73. self.empty_record_dict = self.create_record_dict_for_vars()
  74. def _topology_state_variables(self):
  75. """
  76. Returns the state variables of the network
  77. """
  78. state_vars = []
  79. if self.topology == SHALLOW_SPIKING_DOUBLE_FILTER:
  80. state_vars = [SPK_IN, SPK_SYN_CUR, SPK_NEU_CUR,
  81. SPK_NEU_MEM, SPK_OUT, FILT_NEU_CUR,
  82. FILT_NEU_MEM, FILT_SPK_OUT,FILT2_NEU_MEM
  83. ]
  84. if self.topology == SHALLOW_LEAKY:
  85. state_vars = [SPK_IN, LEAK_SYN_CUR, LEAK_NEU_CUR, SPK_OUT,
  86. LEAK_NEU_MEM]
  87. if not self.snn_config["post_process_filt"]:
  88. state_vars.extend([FILT_NEU_MEM, FILT_SPK_OUT])
  89. if self.topology== ENC_SHALLOW_SPIKING:
  90. state_vars = [EMG_IN, ENC_NEU_CUR, ENC_NEU_MEM, ENC_SPK_OUT,
  91. SPK_SYN_CUR, SPK_NEU_CUR, SPK_NEU_MEM, SPK_OUT,
  92. FILT_NEU_CUR, FILT_NEU_MEM, FILT_SPK_OUT,
  93. FILT2_NEU_MEM]
  94. if self.topology == ENC_SHALLOW_LEAKY:
  95. state_vars = [EMG_IN, ENC_NEU_CUR, ENC_NEU_MEM, ENC_SPK_OUT,
  96. LEAK_SYN_CUR, LEAK_NEU_CUR, SPK_OUT,
  97. LEAK_NEU_MEM]
  98. if self.snn_config["first_filter_tau"] > 0:
  99. state_vars.extend([FILT_NEU_MEM, FILT_SPK_OUT])
  100. if self.snn_config["use_aleaky"]:
  101. state_vars.extend([ENC_THRESHOLD_ADAPT, ENC_ALIF_THRESHOLD])
  102. return state_vars
  103. def _set_parameter_per_electrode(self, parameter:float, parameter_name:str):
  104. """
  105. Sets the decay and threshold parameter for the encoding layer
  106. """
  107. if self.set_electrode_specific_parameters:
  108. # assert that self.enc_tau_mem is a list of length num_inputs
  109. assert isinstance(parameter, list), "parameter should be a list"
  110. parameter_per_neuron = np.repeat(parameter , self.n_channels_per_electrode)
  111. assert len(parameter_per_neuron) == self.num_inputs, f"parameter should be a list of length {self.num_inputs}"
  112. if parameter_name == 'tau_mem':
  113. parameter_per_neuron = torch.stack([map_tau_to_decay(parameter_per_neuron[i], self.dt) for i in range(self.num_inputs)],dim=0)
  114. else:
  115. parameter_per_neuron = parameter
  116. if parameter_name == 'tau_mem':
  117. parameter_per_neuron = map_tau_to_decay(parameter,self.dt)
  118. return parameter_per_neuron
  119. def _plot_network_variables(self):
  120. """
  121. Defines which network variables to plot in the network_variables_plot
  122. """
  123. plot_syn_cur = None
  124. plot_neu_cur = None
  125. plot_neu_mem = None
  126. if self.topology == SHALLOW_SPIKING_DOUBLE_FILTER:
  127. plot_syn_cur = SPK_SYN_CUR
  128. plot_neu_cur = SPK_NEU_CUR
  129. plot_neu_mem = SPK_NEU_MEM
  130. if self.topology == SHALLOW_LEAKY:
  131. plot_syn_cur = LEAK_SYN_CUR
  132. plot_neu_cur = LEAK_NEU_CUR
  133. plot_neu_mem = LEAK_NEU_MEM
  134. if self.topology == ENC_SHALLOW_SPIKING:
  135. plot_syn_cur = SPK_SYN_CUR
  136. plot_neu_cur = ENC_NEU_CUR
  137. plot_neu_mem = ENC_NEU_MEM
  138. if self.topology == ENC_SHALLOW_LEAKY:
  139. plot_syn_cur = LEAK_SYN_CUR
  140. plot_neu_cur = ENC_NEU_CUR
  141. plot_neu_mem = ENC_NEU_MEM
  142. return plot_syn_cur, plot_neu_cur, plot_neu_mem
  143. def create_record_dict_for_vars(self):
  144. """
  145. Creates a dictionary to store the variables of the network
  146. """
  147. record_dict = {f'{state_var}':[] for state_var in self.state_vars}
  148. return record_dict
  149. def describe(self):
  150. """
  151. Describes the SNN topology
  152. """
  153. descriptions = {
  154. SHALLOW_LEAKY: 'a single layer with a leaky integrate and fire neuron',
  155. 'shallow_spiking_single_filter': 'a single layer with a spiking neuron and a single filter. Filter is a leaky integrator',
  156. SHALLOW_SPIKING_DOUBLE_FILTER: 'a single layer with a spiking neuron and two filters. Two layers of leaky integrators'
  157. }
  158. return descriptions.get(self.topology, 'Unknown topology')
  159. def forward(self):
  160. """
  161. Forward pass for the network
  162. """
  163. print("Forward is implemented in the child class")
  164. class ShallowSpikingDoubleFilter(SnnTopology):
  165. def __init__(self, timesteps, num_inputs, num_outputs, snn_config, use_wandb, dt):
  166. SnnTopology.__init__(self, timesteps, num_inputs, num_outputs, snn_config, use_wandb, dt)
  167. self._init_network_parameters()
  168. self._create_network_layers()
  169. self.prediction_var = FILT2_NEU_MEM if self.snn_config["second_filter_tau"] > 0 else FILT_NEU_MEM
  170. self.input_var = SPK_IN
  171. self.leaky_integrators_mem = [FILT2_NEU_MEM, FILT2_NEU_MEM]
  172. self.leaky_integrators_spk= [FILT_SPK_OUT]
  173. self.filter_layers = [self.first_filter, self.second_filter] #if not self.post_process_filt else [self.first_filter]
  174. def _create_network_layers(self):
  175. self.fc1 = torch.nn.Linear(in_features=self.num_inputs,
  176. out_features=self.num_outputs)
  177. self.lif1 = snn.Synaptic(beta=self.neurons_betas,
  178. alpha=self.synapses_alphas,
  179. spike_grad=self.spike_grad,
  180. reset_mechanism="zero",
  181. threshold=self.neurons_thresholds,
  182. learn_threshold=self.learn_threshold,
  183. learn_beta=self.learn_beta,
  184. learn_alpha=self.learn_alpha)
  185. self.fc2 = torch.nn.Linear(in_features=self.num_outputs,
  186. out_features=self.num_outputs,
  187. bias=False)
  188. self.first_filter = snn.Leaky(beta=self.first_filter_beta,
  189. threshold=self.leaky_integrator_threshold,
  190. learn_beta=False)
  191. # if not self.post_process_filt:
  192. self.second_filter = snn.Leaky(beta=self.second_filter_beta,
  193. threshold=self.leaky_integrator_threshold,
  194. learn_beta=False)
  195. if self.snn_config["w_init_dist"] == 'uniform':
  196. nn.init.uniform_(self.fc1.weight,a=self.snn_config["w_init"][0],b=self.snn_config["w_init"][1])
  197. if self.snn_config["w_init_dist"] == 'normal':
  198. nn.init.normal_(self.fc1.weight, mean=self.snn_config["w_init_mean"],
  199. std=self.snn_config["w_init_std"])
  200. nn.init.uniform(self.fc1.bias,a=self.snn_config["bias_init"][0],b=self.snn_config["bias_init"][1])
  201. # enforce a one-to-one mapping between the spiking layer and the filtering layer
  202. self.fc2.weight = nn.Parameter(torch.eye(self.num_outputs) * self.snn_config["w_fixed_filt"])
  203. self.fc2.bias = nn.Parameter(torch.zeros(self.num_outputs))
  204. self.fc2.requires_grad_(False)
  205. def _init_network_parameters(self):
  206. # Input to spiking neuron connected via alpha synapse with parameters tau_syn
  207. # Spiking neuron has parameters tau_mem and threshold
  208. self.tau_syn = self.snn_config["tau_syn"]
  209. self.tau_mem = self.snn_config["tau_mem"]
  210. self.spk_threshold = self.snn_config["spk_threshold"]
  211. self.alpha = map_tau_to_decay(self.tau_syn, self.dt)
  212. self.beta = map_tau_to_decay(self.tau_mem, self.dt)
  213. self.learn_alpha = self.snn_config["learn_tau_syn"]
  214. self.learn_beta = self.snn_config["learn_tau_mem"]
  215. self.learn_threshold = self.snn_config["learn_threshold"]
  216. self.post_process_filt = self.snn_config["post_process_filt"]
  217. self.post_process_filt_cutoff = self.snn_config["post_filt_cutoff"]
  218. # In case parameters are learned, initialize them from a distribution
  219. if self.learn_beta:
  220. neurons_betas = torch.rand(self.num_outputs)
  221. nn.init.normal_(neurons_betas, mean=self.beta,
  222. std=self.param_std * self.beta)
  223. print(f"Init beta:{neurons_betas}\n")
  224. else:
  225. neurons_betas = self.beta
  226. if self.learn_alpha:
  227. synapses_alphas = torch.rand(self.num_outputs)
  228. nn.init.normal_(synapses_alphas, mean=self.alpha,
  229. std=self.param_std * self.alpha)
  230. else:
  231. synapses_alphas = self.alpha
  232. if self.learn_threshold:
  233. neurons_thresholds = torch.rand(self.num_outputs)
  234. nn.init.normal_(neurons_thresholds, mean=self.spk_threshold,
  235. std=self.param_std * self.spk_threshold)
  236. else:
  237. neurons_thresholds = self.spk_threshold
  238. self.neurons_betas = neurons_betas
  239. self.synapses_alphas = synapses_alphas
  240. self.neurons_thresholds = neurons_thresholds
  241. self.first_filter_beta = map_tau_to_decay(self.snn_config["first_filter_tau"], self.dt)
  242. self.second_filter_beta = map_tau_to_decay(self.snn_config["second_filter_tau"], self.dt)
  243. def forward(self, spk_in, state_dict=None):
  244. """Forward pass for each time step"""
  245. empty_record_dict = self.create_record_dict_for_vars()
  246. if state_dict is None:
  247. # Initalize membrane potential as empty tensor
  248. spk_syn_cur, spk_neu_mem = self.lif1.init_synaptic()
  249. filt_neu_mem = self.first_filter.init_leaky()
  250. # if not self.post_process_filt:
  251. filt2_neu_mem = self.second_filter.init_leaky()
  252. else: # if not initializing as empty rely on last timestep
  253. spk_syn_cur = state_dict[SPK_SYN_CUR]
  254. spk_neu_cur = state_dict[SPK_NEU_CUR]
  255. spk_neu_mem = state_dict[SPK_NEU_MEM]
  256. spk_out = state_dict[SPK_OUT]
  257. filt_neu_mem = state_dict[FILT_NEU_MEM]
  258. filt_neu_cur = state_dict[FILT_NEU_CUR]
  259. filt_spk_out = state_dict[FILT_SPK_OUT]
  260. # if not self.post_process_filt:
  261. filt2_neu_mem = state_dict[FILT2_NEU_MEM]
  262. for step in range(self.timesteps):
  263. if state_dict is None or step > 0:
  264. spk_neu_cur = self.fc1(spk_in[:, step])
  265. spk_out, spk_syn_cur, spk_neu_mem = self.lif1(spk_neu_cur, spk_syn_cur, spk_neu_mem)
  266. filt_neu_cur = self.fc2(spk_out)
  267. filt_spk_out, filt_neu_mem = self.first_filter(filt_neu_cur, filt_neu_mem)
  268. # if not self.post_process_filt:
  269. _, filt2_neu_mem = self.second_filter(filt_neu_mem, filt2_neu_mem)
  270. # record the variables for layer 1
  271. empty_record_dict[SPK_IN].append(spk_in[:, step])
  272. empty_record_dict[SPK_SYN_CUR].append(spk_syn_cur)
  273. empty_record_dict[SPK_NEU_CUR].append(spk_neu_cur)
  274. empty_record_dict[SPK_NEU_MEM].append(spk_neu_mem)
  275. empty_record_dict[SPK_OUT].append(spk_out)
  276. empty_record_dict[FILT_NEU_CUR].append(filt_neu_cur)
  277. empty_record_dict[FILT_NEU_MEM].append(filt_neu_mem)
  278. empty_record_dict[FILT_SPK_OUT].append(filt_spk_out)
  279. # if not self.post_process_filt:
  280. empty_record_dict[FILT2_NEU_MEM].append(filt2_neu_mem)
  281. # stack the recorded variables
  282. for var in empty_record_dict:
  283. empty_record_dict[var] = torch.stack(empty_record_dict[var], dim=1)
  284. filled_record_dict = empty_record_dict
  285. return filled_record_dict
  286. class ShallowLeaky(SnnTopology):
  287. """
  288. A shallow network with lekay integrator neurons followed by a filtering neuron
  289. """
  290. def __init__(self, timesteps, num_inputs, num_outputs, snn_config, use_wandb, dt):
  291. SnnTopology.__init__(self, timesteps, num_inputs, num_outputs, snn_config, use_wandb, dt)
  292. self._init_network_parameters()
  293. self._create_network_layers()
  294. self.prediction_var = FILT_NEU_MEM if not self.post_process_filt else LEAK_NEU_MEM
  295. self.input_var = SPK_IN
  296. self.leaky_integrators_mem = [FILT_NEU_MEM]
  297. self.leaky_integrators_spk= [SPK_OUT, FILT_SPK_OUT] if not self.post_process_filt else [SPK_OUT]
  298. self.filter_layers = [self.first_filter] if not self.post_process_filt else []
  299. def _init_network_parameters(self):
  300. # Input to leaky neuron connected via alpha synapse with parameters tau_syn
  301. # Leaky neuron has parameters tau_mem
  302. self.tau_syn = self.snn_config["tau_syn"]
  303. self.tau_mem = self.snn_config["tau_mem"]
  304. self.alpha = map_tau_to_decay(self.tau_syn, self.dt)
  305. self.beta = map_tau_to_decay(self.tau_mem, self.dt)
  306. self.learn_alpha = self.snn_config["learn_tau_syn"]
  307. self.learn_beta = self.snn_config["learn_tau_mem"]
  308. self.post_process_filt = self.snn_config["post_process_filt"]
  309. self.post_process_filt_cutoff = self.snn_config["post_filt_cutoff"]
  310. # In case parameters are learned, initialize them from a distribution
  311. if self.learn_beta:
  312. neurons_betas = torch.rand(self.num_outputs)
  313. nn.init.normal_(neurons_betas, mean=self.beta,
  314. std=self.param_std * self.beta)
  315. print(f"Init beta:{neurons_betas}\n")
  316. else:
  317. neurons_betas = self.beta
  318. if self.learn_alpha:
  319. synapses_alphas = torch.rand(self.num_outputs)
  320. nn.init.normal_(synapses_alphas, mean=self.alpha,
  321. std=self.param_std * self.alpha)
  322. else:
  323. synapses_alphas = self.alpha
  324. self.neurons_betas = neurons_betas
  325. self.synapses_alphas = synapses_alphas
  326. self.first_filter_beta = map_tau_to_decay(self.snn_config["first_filter_tau"], self.dt)
  327. def _create_network_layers(self):
  328. self.fc1 = torch.nn.Linear(in_features=self.num_inputs,
  329. out_features=self.num_outputs)
  330. self.lif1 = snn.Synaptic(beta=self.neurons_betas,
  331. alpha=self.synapses_alphas,
  332. spike_grad=self.spike_grad,
  333. reset_mechanism="zero",
  334. threshold=self.leaky_integrator_threshold,
  335. learn_threshold=False,
  336. learn_beta=self.learn_beta,
  337. learn_alpha=self.learn_alpha)
  338. if not self.post_process_filt:
  339. self.first_filter = snn.Leaky(beta=self.first_filter_beta,
  340. threshold=self.leaky_integrator_threshold,
  341. learn_beta=False)
  342. if self.snn_config["w_init_dist"] == 'uniform':
  343. nn.init.uniform_(self.fc1.weight,a=self.snn_config["w_init"][0],b=self.snn_config["w_init"][1])
  344. if self.snn_config["w_init_dist"] == 'normal':
  345. nn.init.normal_(self.fc1.weight, mean=self.snn_config["w_init_mean"],
  346. std=self.snn_config["w_init_std"])
  347. nn.init.uniform(self.fc1.bias,a=self.snn_config["bias_init"][0],b=self.snn_config["bias_init"][1])
  348. def forward(self, spk_in, state_dict=None):
  349. empty_record_dict = self.create_record_dict_for_vars()
  350. if state_dict is None:
  351. # Initalize membrane potential as empty tensor
  352. leaky_syn_cur, leaky_neu_mem = self.lif1.init_synaptic()
  353. if not self.post_process_filt:
  354. filt_neu_mem = self.first_filter.init_leaky()
  355. else: # if not initializing as empty rely on last timestep values, restroring the state
  356. leaky_syn_cur = state_dict[LEAK_SYN_CUR]
  357. leaky_neu_cur = state_dict[LEAK_NEU_CUR]
  358. leaky_neu_mem = state_dict[LEAK_NEU_MEM]
  359. spk_out = state_dict[SPK_OUT]
  360. if not self.post_process_filt:
  361. filt_neu_mem = state_dict[FILT_NEU_MEM]
  362. filt_spk_out = state_dict[FILT_SPK_OUT]
  363. for step in range(self.timesteps):
  364. if state_dict is None or step > 0:
  365. leaky_neu_cur = self.fc1(spk_in[:, step])
  366. spk_out, leaky_syn_cur, leaky_neu_mem = self.lif1(leaky_neu_cur, leaky_syn_cur, leaky_neu_mem)
  367. # clip membrane potential to 0
  368. leaky_neu_mem = torch.clamp(leaky_neu_mem, min=0)
  369. if not self.post_process_filt:
  370. filt_spk_out, filt_neu_mem = self.first_filter(leaky_neu_mem, filt_neu_mem)
  371. # record the variables for layer 1
  372. empty_record_dict[SPK_IN].append(spk_in[:, step])
  373. empty_record_dict[LEAK_SYN_CUR].append(leaky_syn_cur)
  374. empty_record_dict[LEAK_NEU_CUR].append(leaky_neu_cur)
  375. empty_record_dict[LEAK_NEU_MEM].append(leaky_neu_mem)
  376. empty_record_dict[SPK_OUT].append(spk_out)
  377. if not self.post_process_filt:
  378. empty_record_dict[FILT_NEU_MEM].append(filt_neu_mem)
  379. empty_record_dict[FILT_SPK_OUT].append(filt_spk_out)
  380. # stack the recorded variables
  381. for var in empty_record_dict:
  382. empty_record_dict[var] = torch.stack(empty_record_dict[var], dim=1)
  383. filled_record_dict = empty_record_dict
  384. return filled_record_dict
  385. class EncodingShallowSpikingDoubleFilter(SnnTopology):
  386. """
  387. A shallow SNN encoding iEMG into spikes then decoding forces
  388. """
  389. def __init__(self, timesteps, num_inputs, num_outputs, snn_config, use_wandb, dt):
  390. SnnTopology.__init__(self, timesteps, num_inputs, num_outputs, snn_config, use_wandb, dt)
  391. self._init_network_parameters()
  392. self._create_network_layers()
  393. self.prediction_var = FILT2_NEU_MEM if self.snn_config["second_filter_tau"] > 0 else FILT_NEU_MEM
  394. self.input_var = EMG_IN
  395. self.leaky_integrators_mem = [FILT_NEU_MEM, FILT2_NEU_MEM]
  396. self.leaky_integrators_spk= [FILT_SPK_OUT]
  397. self.filter_layers = [self.first_filter, self.second_filter]
  398. def _create_network_layers(self):
  399. self.enc_fc = torch.nn.Linear(in_features=self.num_inputs,
  400. out_features=self.num_inputs,
  401. bias=False)
  402. # set weights of encoding layer to identity
  403. self.enc_fc.weight = nn.Parameter(torch.eye(self.num_inputs) * self.snn_config["w_fixed_enc"])
  404. # self.enc_fc.bias = nn.Parameter(torch.zeros(self.num_inputs))
  405. self.enc_fc.requires_grad_(False)
  406. self.encoding_layer = snn.Leaky(beta=self.enc_beta,
  407. threshold=self.enc_neurons_threshold,
  408. learn_beta=False,
  409. learn_threshold=self.enc_learn_threshold,
  410. reset_mechanism="zero")
  411. self.fc1 = torch.nn.Linear(in_features=self.num_inputs,
  412. out_features=self.num_outputs)
  413. self.lif1 = snn.Synaptic(beta=self.neurons_betas,
  414. alpha=self.synapses_alphas,
  415. spike_grad=self.spike_grad,
  416. reset_mechanism="subtract",
  417. threshold=self.neurons_thresholds,
  418. learn_threshold=self.learn_threshold,
  419. learn_beta=self.learn_beta,
  420. learn_alpha=self.learn_alpha)
  421. self.fc2 = torch.nn.Linear(in_features=self.num_outputs,
  422. out_features=self.num_outputs,
  423. bias=False)
  424. self.first_filter = snn.Leaky(beta=self.first_filter_beta,
  425. threshold=self.leaky_integrator_threshold,
  426. learn_beta=False)
  427. self.second_filter = snn.Leaky(beta=self.second_filter_beta,
  428. threshold=self.leaky_integrator_threshold,
  429. learn_beta=False)
  430. if self.snn_config["w_init_dist"] == 'uniform':
  431. nn.init.uniform_(self.fc1.weight,a=self.snn_config["w_init"][0],b=self.snn_config["w_init"][1])
  432. if self.snn_config["w_init_dist"] == 'normal':
  433. nn.init.normal_(self.fc1.weight, mean=self.snn_config["w_init_mean"],
  434. std=self.snn_config["w_init_std"])
  435. nn.init.uniform(self.fc1.bias,a=self.snn_config["bias_init"][0],b=self.snn_config["bias_init"][1])
  436. # enforce a one-to-one mapping between the spiking layer and the filtering layer
  437. self.fc2.weight = nn.Parameter(torch.eye(self.num_outputs) * self.snn_config["w_fixed_filt"])
  438. self.fc2.bias = nn.Parameter(torch.zeros(self.num_outputs))
  439. self.fc2.requires_grad_(False)
  440. def _init_network_parameters(self):
  441. # Input to spiking neuron connected via alpha synapse with parameters tau_syn
  442. # Spiking neuron has parameters tau_mem and threshold
  443. self.set_electrode_specific_parameters = self.snn_config["set_electrode_specific_parameters"]
  444. self.enc_tau_mem = self.snn_config["enc_tau_mem"]
  445. self.enc_beta = self._set_parameter_per_electrode(self.enc_tau_mem, 'tau_mem') #map_tau_to_decay(self.enc_tau_mem, self.dt)
  446. self.enc_learn_threshold = self.snn_config["learn_enc_threshold"]
  447. self.tau_syn = self.snn_config["tau_syn"]
  448. self.tau_mem = self.snn_config["tau_mem"]
  449. self.spk_threshold = self.snn_config["spk_threshold"]
  450. self.enc_threshold= self._set_parameter_per_electrode(self.snn_config["enc_spk_threshold"], 'threshold') #self.snn_config.enc_spk_threshold
  451. self.alpha = map_tau_to_decay(self.tau_syn, self.dt)
  452. self.beta = map_tau_to_decay(self.tau_mem, self.dt)
  453. self.learn_alpha = self.snn_config["learn_tau_syn"]
  454. self.learn_beta = self.snn_config["learn_tau_mem"]
  455. self.learn_threshold = self.snn_config["learn_threshold"]
  456. # In case parameters are learned, initialize them from a distribution
  457. if self.enc_learn_threshold:
  458. enc_neurons_thresholds = torch.rand(self.num_inputs)
  459. nn.init.normal_(enc_neurons_thresholds, mean=self.enc_threshold,
  460. std=self.param_std * self.enc_threshold)
  461. else:
  462. enc_neurons_thresholds = self.enc_threshold
  463. if self.learn_beta:
  464. neurons_betas = torch.rand(self.num_outputs)
  465. nn.init.normal_(neurons_betas, mean=self.beta,
  466. std=self.param_std * self.beta)
  467. print(f"Init beta:{neurons_betas}\n")
  468. else:
  469. neurons_betas = self.beta
  470. if self.learn_alpha:
  471. synapses_alphas = torch.rand(self.num_outputs)
  472. nn.init.normal_(synapses_alphas, mean=self.alpha,
  473. std=self.param_std * self.alpha)
  474. else:
  475. synapses_alphas = self.alpha
  476. if self.learn_threshold:
  477. neurons_thresholds = torch.rand(self.num_outputs)
  478. nn.init.normal_(neurons_thresholds, mean=self.spk_threshold,
  479. std=self.param_std * self.spk_threshold)
  480. else:
  481. neurons_thresholds = self.spk_threshold
  482. self.neurons_betas = neurons_betas
  483. self.synapses_alphas = synapses_alphas
  484. self.neurons_thresholds = neurons_thresholds
  485. self.enc_neurons_threshold = enc_neurons_thresholds
  486. self.first_filter_beta = map_tau_to_decay(self.snn_config["first_filter_tau"], self.dt)
  487. self.second_filter_beta = map_tau_to_decay(self.snn_config["second_filter_tau"], self.dt)
  488. def forward(self, emg_in, state_dict):
  489. """Forward pass for each time step"""
  490. empty_record_dict = self.create_record_dict_for_vars()
  491. if state_dict is None:
  492. # Initalize membrane potential as empty tensor
  493. enc_neu_mem = self.encoding_layer.init_leaky()
  494. spk_syn_cur, spk_neu_mem = self.lif1.init_synaptic()
  495. filt_neu_mem = self.first_filter.init_leaky()
  496. filt2_neu_mem = self.second_filter.init_leaky()
  497. else: # if not initializing as empty rely on last timestep
  498. enc_neu_mem = state_dict[ENC_NEU_MEM]
  499. enc_neu_cur = state_dict[ENC_NEU_CUR]
  500. enc_spk_out = state_dict[ENC_SPK_OUT]
  501. spk_syn_cur = state_dict[SPK_SYN_CUR]
  502. spk_neu_cur = state_dict[SPK_NEU_CUR]
  503. spk_neu_mem = state_dict[SPK_NEU_MEM]
  504. spk_out = state_dict[SPK_OUT]
  505. filt_neu_mem = state_dict[FILT_NEU_MEM]
  506. filt_neu_cur = state_dict[FILT_NEU_CUR]
  507. filt_spk_out = state_dict[FILT_SPK_OUT]
  508. filt2_neu_mem = state_dict[FILT2_NEU_MEM]
  509. for step in range(self.timesteps):
  510. if state_dict is None or step > 0:
  511. enc_neu_cur = self.enc_fc(emg_in[:, step])
  512. enc_spk_out, enc_neu_mem = self.encoding_layer(enc_neu_cur, enc_neu_mem)
  513. spk_neu_cur = self.fc1(enc_spk_out)
  514. spk_out, spk_syn_cur, spk_neu_mem = self.lif1(spk_neu_cur, spk_syn_cur, spk_neu_mem)
  515. filt_neu_cur = self.fc2(spk_out)
  516. filt_spk_out, filt_neu_mem = self.first_filter(filt_neu_cur, filt_neu_mem)
  517. _, filt2_neu_mem = self.second_filter(filt_neu_mem, filt2_neu_mem)
  518. # record the variables for layer 1
  519. empty_record_dict[EMG_IN].append(emg_in[:, step])
  520. empty_record_dict[ENC_NEU_CUR].append(enc_neu_cur)
  521. empty_record_dict[ENC_NEU_MEM].append(enc_neu_mem)
  522. empty_record_dict[ENC_SPK_OUT].append(enc_spk_out)
  523. empty_record_dict[SPK_SYN_CUR].append(spk_syn_cur)
  524. empty_record_dict[SPK_NEU_CUR].append(spk_neu_cur)
  525. empty_record_dict[SPK_NEU_MEM].append(spk_neu_mem)
  526. empty_record_dict[SPK_OUT].append(spk_out)
  527. empty_record_dict[FILT_NEU_CUR].append(filt_neu_cur)
  528. empty_record_dict[FILT_NEU_MEM].append(filt_neu_mem)
  529. empty_record_dict[FILT_SPK_OUT].append(filt_spk_out)
  530. empty_record_dict[FILT2_NEU_MEM].append(filt2_neu_mem)
  531. # stack the recorded variables
  532. for var in empty_record_dict:
  533. empty_record_dict[var] = torch.stack(empty_record_dict[var], dim=1)
  534. filled_record_dict = empty_record_dict
  535. return filled_record_dict
  536. class EncodingShallowLeaky(SnnTopology):
  537. """
  538. A shallow network with lekay integrator neurons followed by a filtering neuron
  539. """
  540. def __init__(self, timesteps, num_inputs, num_outputs, snn_config, use_wandb, dt):
  541. SnnTopology.__init__(self, timesteps, num_inputs, num_outputs, snn_config, use_wandb, dt)
  542. self._init_network_parameters()
  543. self._create_network_layers()
  544. self.prediction_var = FILT_NEU_MEM if self.snn_config["first_filter_tau"] >0 else LEAK_NEU_MEM
  545. self.input_var = EMG_IN
  546. self.leaky_integrators_mem = [FILT_NEU_MEM]
  547. self.leaky_integrators_spk= [SPK_OUT, FILT_SPK_OUT] if self.snn_config["first_filter_tau"] >0 else [SPK_OUT]
  548. self.filter_layers = [self.first_filter] if self.snn_config["first_filter_tau"] >0 else []
  549. def _init_network_parameters(self):
  550. # Input to leaky neuron connected via alpha synapse with parameters tau_syn
  551. # Leaky neuron has parameters tau_mem
  552. self.set_electrode_specific_parameters = self.snn_config["set_electrode_specific_parameters"]
  553. self.enc_tau_mem = self.snn_config["enc_tau_mem"]
  554. self.enc_beta = self._set_parameter_per_electrode(self.enc_tau_mem, 'tau_mem') #map_tau_to_decay(self.enc_tau_mem, self.dt)
  555. self.enc_learn_threshold = self.snn_config["learn_enc_threshold"]
  556. self.tau_syn = self.snn_config["tau_syn"]
  557. self.tau_mem = self.snn_config["tau_mem"]
  558. self.w_recurrent = self.snn_config["w_recurrent"]
  559. self.alpha = map_tau_to_decay(self.tau_syn, self.dt)
  560. self.beta = map_tau_to_decay(self.tau_mem, self.dt)
  561. self.enc_threshold= self._set_parameter_per_electrode(self.snn_config["enc_spk_threshold"], 'threshold') #self.snn_config.enc_spk_threshold
  562. self.learn_alpha = self.snn_config["learn_tau_syn"]
  563. self.learn_beta = self.snn_config["learn_tau_mem"]
  564. self.post_process_filt = self.snn_config["post_process_filt"]
  565. self.post_process_filt_cutoff = self.snn_config["post_filt_cutoff"]
  566. enc_neurons_thresholds = self.enc_threshold
  567. if self.learn_beta:
  568. neurons_betas = torch.rand(self.num_outputs)
  569. nn.init.normal_(neurons_betas, mean=self.beta,
  570. std=self.param_std * self.beta)
  571. print(f"Init beta:{neurons_betas}\n")
  572. else:
  573. neurons_betas = self.beta
  574. if self.learn_alpha:
  575. synapses_alphas = torch.rand(self.num_outputs)
  576. nn.init.normal_(synapses_alphas, mean=self.alpha,
  577. std=self.param_std * self.alpha)
  578. else:
  579. synapses_alphas = self.alpha
  580. self.neurons_betas = neurons_betas
  581. self.synapses_alphas = synapses_alphas
  582. self.enc_neurons_threshold = enc_neurons_thresholds
  583. if self.snn_config["first_filter_tau"] > 0:
  584. self.first_filter_beta = map_tau_to_decay(self.snn_config["first_filter_tau"], self.dt)
  585. def _create_network_layers(self):
  586. self.enc_fc = torch.nn.Linear(in_features=self.num_inputs,
  587. out_features=self.num_inputs,
  588. bias=False)
  589. # set weights of encoding layer to identity
  590. self.enc_fc.weight = nn.Parameter(torch.eye(self.num_inputs) * self.snn_config["w_fixed_enc"])
  591. self.enc_fc.requires_grad_(False)
  592. if self.snn_config["add_recurrent"]:
  593. self.encoding_layer = snn.RLeaky(beta=self.enc_beta,
  594. threshold=self.enc_neurons_threshold,
  595. learn_beta=False,
  596. learn_threshold=self.enc_learn_threshold,
  597. reset_mechanism="subtract",
  598. all_to_all=False,
  599. V=self.w_recurrent,
  600. learn_recurrent=False)
  601. elif self.snn_config["use_aleaky"]:
  602. self.encoding_layer = snn.ALeaky(beta=self.enc_beta,
  603. threshold=self.enc_neurons_threshold,
  604. learn_beta=False,
  605. learn_threshold=self.enc_learn_threshold,
  606. reset_mechanism="zero",
  607. tau_adapt = self.snn_config["threshold_tau_adapt"],
  608. threshold_weight = self.snn_config["threshold_scale_adapt"],
  609. dt = self.dt)
  610. else:
  611. self.encoding_layer = snn.Leaky(beta=self.enc_beta,
  612. threshold=self.enc_neurons_threshold,
  613. learn_beta=False,
  614. learn_threshold=self.enc_learn_threshold,
  615. reset_mechanism="zero")
  616. self.fc1 = torch.nn.Linear(in_features=self.num_inputs,
  617. out_features=self.num_outputs,
  618. )
  619. self.lif1 = snn.Synaptic(beta=self.neurons_betas,
  620. alpha=self.synapses_alphas,
  621. spike_grad=self.spike_grad,
  622. reset_mechanism="zero",
  623. threshold=self.leaky_integrator_threshold,
  624. learn_threshold=False,
  625. learn_beta=self.learn_beta,
  626. learn_alpha=self.learn_alpha)
  627. if self.snn_config["first_filter_tau"] > 0:
  628. self.first_filter = snn.Leaky(beta=self.first_filter_beta,
  629. threshold=self.leaky_integrator_threshold,
  630. learn_beta=False)
  631. if self.snn_config["w_init_dist"] == 'uniform':
  632. nn.init.uniform_(self.fc1.weight,a=self.snn_config["w_init"][0],b=self.snn_config["w_init"][1])
  633. if self.snn_config["w_init_dist"] == 'normal':
  634. nn.init.normal_(self.fc1.weight, mean=self.snn_config["w_init_mean"],
  635. std=self.snn_config["w_init_std"])
  636. # keep bias fixed
  637. # self.fc1.bias = nn.Parameter(torch.zeros(self.num_outputs))
  638. # self.fc1.bias.requires_grad = False
  639. nn.init.uniform(self.fc1.bias,a=self.snn_config["bias_init"][0],b=self.snn_config["bias_init"][1])
  640. def forward(self, emg_in, state_dict=None):
  641. empty_record_dict = self.create_record_dict_for_vars()
  642. if state_dict is None:
  643. # Initalize membrane potential as empty tensor
  644. if self.snn_config["add_recurrent"]:
  645. enc_spk_out, enc_neu_mem = self.encoding_layer.init_rleaky()
  646. elif self.snn_config["use_aleaky"]:
  647. enc_neu_mem, threshold_adapt = self.encoding_layer.init_aleaky()
  648. else:
  649. enc_neu_mem = self.encoding_layer.init_leaky()
  650. leaky_syn_cur, leaky_neu_mem = self.lif1.init_synaptic()
  651. if self.snn_config["first_filter_tau"] > 0:
  652. filt_neu_mem = self.first_filter.init_leaky()
  653. else: # if not initializing as empty rely on last timestep values, restroring the state
  654. enc_neu_mem = state_dict[ENC_NEU_MEM]
  655. enc_neu_cur = state_dict[ENC_NEU_CUR]
  656. enc_spk_out = state_dict[ENC_SPK_OUT]
  657. leaky_syn_cur = state_dict[LEAK_SYN_CUR]
  658. leaky_neu_cur = state_dict[LEAK_NEU_CUR]
  659. leaky_neu_mem = state_dict[LEAK_NEU_MEM]
  660. spk_out = state_dict[SPK_OUT]
  661. if self.snn_config["first_filter_tau"] > 0:
  662. filt_neu_mem = state_dict[FILT_NEU_MEM]
  663. filt_spk_out = state_dict[FILT_SPK_OUT]
  664. if self.snn_config["use_aleaky"]:
  665. threshold_adapt = state_dict[ENC_THRESHOLD_ADAPT]
  666. threshold = state_dict[ENC_ALIF_THRESHOLD]
  667. for step in range(self.timesteps):
  668. if state_dict is None or step > 0:
  669. enc_neu_cur = self.enc_fc(emg_in[:, step])
  670. if self.snn_config["add_recurrent"]:
  671. enc_spk_out, enc_neu_mem = self.encoding_layer(enc_neu_cur,enc_spk_out, enc_neu_mem)
  672. elif self.snn_config["use_aleaky"]:
  673. enc_spk_out, enc_neu_mem, threshold_adapt, threshold = self.encoding_layer(enc_neu_cur, enc_neu_mem, threshold_adapt)
  674. else:
  675. enc_spk_out, enc_neu_mem = self.encoding_layer(enc_neu_cur, enc_neu_mem)
  676. leaky_neu_cur = self.fc1(enc_spk_out)
  677. spk_out, leaky_syn_cur, leaky_neu_mem = self.lif1(leaky_neu_cur, leaky_syn_cur, leaky_neu_mem)
  678. # clip membrane potential to 0
  679. leaky_neu_mem = torch.clamp(leaky_neu_mem, min=0)
  680. if self.snn_config["first_filter_tau"] > 0:
  681. filt_spk_out, filt_neu_mem = self.first_filter(leaky_neu_mem, filt_neu_mem)
  682. # record the variables for layer 1
  683. empty_record_dict[EMG_IN].append(emg_in[:, step])
  684. empty_record_dict[ENC_NEU_CUR].append(enc_neu_cur)
  685. empty_record_dict[ENC_NEU_MEM].append(enc_neu_mem)
  686. empty_record_dict[ENC_SPK_OUT].append(enc_spk_out)
  687. empty_record_dict[LEAK_SYN_CUR].append(leaky_syn_cur)
  688. empty_record_dict[LEAK_NEU_CUR].append(leaky_neu_cur)
  689. empty_record_dict[LEAK_NEU_MEM].append(leaky_neu_mem)
  690. empty_record_dict[SPK_OUT].append(spk_out)
  691. if self.snn_config["first_filter_tau"] > 0:
  692. empty_record_dict[FILT_NEU_MEM].append(filt_neu_mem)
  693. empty_record_dict[FILT_SPK_OUT].append(filt_spk_out)
  694. if self.snn_config["use_aleaky"]:
  695. empty_record_dict[ENC_THRESHOLD_ADAPT].append(threshold_adapt)
  696. empty_record_dict[ENC_ALIF_THRESHOLD].append(threshold)
  697. # stack the recorded variables
  698. for var in empty_record_dict:
  699. empty_record_dict[var] = torch.stack(empty_record_dict[var], dim=1)
  700. filled_record_dict = empty_record_dict
  701. return filled_record_dict
  702. # --- SNN Regression ---
  703. class SnnReg():
  704. """
  705. Single-layer spiking neural network in snntorch to regress finger forces
  706. """
  707. def __init__(self, timesteps:int, num_inputs:int, num_outputs:int, snn_config:SNNConfig, train_dataset, test_dataset,
  708. ):
  709. self.snn_config = snn_config
  710. self.num_inputs = num_inputs
  711. self.num_outputs = num_outputs
  712. self.train_dataset = train_dataset
  713. self.test_dataset = test_dataset
  714. torch_seed = snn_config.training["torch_seed"]
  715. use_wandb = snn_config.logging["wandb"]["use_wandb"]
  716. dt = snn_config.training["dt"]
  717. self._set_random_seeds(torch_seed)
  718. passed_net_config = snn_config.decoder_type
  719. if passed_net_config["topology"] == SHALLOW_SPIKING_DOUBLE_FILTER:
  720. self.network = ShallowSpikingDoubleFilter(timesteps,num_inputs,num_outputs, passed_net_config,
  721. use_wandb=use_wandb, dt=dt)
  722. if passed_net_config["topology"] == SHALLOW_LEAKY:
  723. self.network = ShallowLeaky(timesteps,num_inputs,num_outputs, passed_net_config,
  724. use_wandb=use_wandb, dt=dt)
  725. if passed_net_config["topology"] == ENC_SHALLOW_SPIKING:
  726. self.network = EncodingShallowSpikingDoubleFilter(timesteps,num_inputs,num_outputs, passed_net_config,
  727. use_wandb=use_wandb, dt=dt)
  728. if passed_net_config["topology"] == ENC_SHALLOW_LEAKY:
  729. self.network = EncodingShallowLeaky(timesteps,num_inputs,num_outputs, passed_net_config,
  730. use_wandb=use_wandb, dt=dt)
  731. def _set_random_seeds(self,seed):
  732. """
  733. Sets seed for random shuffle of the fingers
  734. """
  735. torch.manual_seed(seed)
  736. torch.cuda.manual_seed(seed)
  737. np.random.seed(seed)
  738. random.seed(seed)
  739. print(f"\nnetwork initialization seed:{torch.random.initial_seed()}\n")
  740. def _check_leaky_integrators_not_firing(self, record_dict:dict):
  741. """
  742. Checks that the filtering layer is not emitting any spikes.
  743. """
  744. for var in self.network.leaky_integrators_spk:
  745. if torch.sum(record_dict[var]) > 0:
  746. print(f"Careful!! Spikes in the filtering layer: {torch.sum(record_dict[var])}...\n")
  747. def _save_network_state_dict(self, previous_record_dict:dict):
  748. """
  749. Saves the state of the network.
  750. Variables saved are the membrane potential, synaptic current, synaptic conductance,
  751. and the input spikes.
  752. """
  753. split_type = self.snn_config.task["split_type"]
  754. overlap_perc = self.snn_config.task["overlap_perc"]
  755. if split_type == 'without_overlap':
  756. t_id = -1 # saved time step should be the last of the previous sample
  757. else:
  758. t_id = int(overlap_perc * self.network.timesteps)
  759. # in case the consecutive samples are non-overlapping, use the last time step
  760. state_dict = {}
  761. for var in self.network.state_vars:
  762. state_dict[var] = previous_record_dict[var].detach()[:, t_id, :]
  763. return state_dict
  764. def _initialize_state_dict(self):
  765. """
  766. Initializes the state dictionary of the network.
  767. """
  768. state_dict = {}
  769. for var in self.network.state_vars:
  770. state_dict[var] = torch.zeros((1, self.num_outputs)) if 'enc' not in var else torch.zeros((1, self.num_inputs))
  771. return state_dict
  772. def _train(self):
  773. """
  774. Trains the network using Adam optimizer and logs losses to wandb.
  775. """
  776. lr = self.snn_config.decoder_type["lr"]
  777. num_iter = self.snn_config.training["num_iter"]
  778. batch_size = self.snn_config.training["train_batch_size"]
  779. use_wandb = self.snn_config.logging["wandb"]["use_wandb"]
  780. log_ts_freq = self.snn_config.logging["log_ts_freq"]
  781. snn_network = self.network
  782. optimizer = torch.optim.Adam(params=snn_network.parameters(), lr=lr)
  783. train_loader = DataLoader(self.train_dataset, batch_size=batch_size, shuffle=False)
  784. record_dict_cache = [] # records variables from all samples
  785. targets_cache = []
  786. tr_loss_per_epoch = []
  787. ts_loss_every_nepochs = []
  788. tr_batch_loss_hist = [] # record loss for all iterations and each batch
  789. # test print torch summary
  790. # summary(snn_network, (1000, 120), None)
  791. with tqdm.trange(num_iter) as pbar:
  792. for i in pbar:
  793. train_batch = iter(train_loader)
  794. ep_running_loss = 0 # the current epoch loss averaged over the batches and samples
  795. mse_epoch = 0
  796. r2_epoch = 0
  797. state_dict = None
  798. for sample_id, (data, label, _) in enumerate(train_batch):
  799. spk_or_emg_in = data.to(DEVICE)
  800. targets = label.to(DEVICE)
  801. if sample_id > 0 and batch_size==1: # save the state of vmem to use with the next sample
  802. state_dict = self._save_network_state_dict(record_dict)
  803. # print(f"Sample {sample_id} of {len(train_loader)}")
  804. record_dict = snn_network.forward(spk_or_emg_in, state_dict)
  805. self._check_leaky_integrators_not_firing(record_dict)
  806. if i == num_iter - 1:
  807. record_dict_cache.append(record_dict)
  808. targets_cache.append(targets)
  809. prediction_var = snn_network.prediction_var
  810. sample_loss_value = snn_network.loss_function(record_dict[prediction_var], targets)
  811. predicted_value = record_dict[prediction_var].cpu().detach()
  812. predicted_value = predicted_value.numpy().reshape(-1, predicted_value.shape[-1])
  813. targets_reshaped = targets.cpu().detach().numpy().reshape(-1, targets.shape[-1])
  814. r2_epoch += r2_score(predicted_value, targets_reshaped)
  815. mse_epoch += mean_squared_error(predicted_value, targets_reshaped)
  816. optimizer.zero_grad()
  817. sample_loss_value.backward() # calculate gradients
  818. optimizer.step()
  819. # store losses: loss_hist is for all iterations,
  820. # hist_loss_epoch is for the current epoch
  821. tr_batch_loss_hist.append(sample_loss_value.item())
  822. ep_running_loss += sample_loss_value.item()
  823. n_samples = (sample_id+1)
  824. mean_batch_loss = ep_running_loss / n_samples
  825. pbar.set_postfix(loss="%.3e" % mean_batch_loss)
  826. if use_wandb:
  827. wandb.log({'mse_tr_per_epoch': mse_epoch / n_samples,'epoch': i})
  828. if i % log_ts_freq == 0:
  829. test_eloss,_, _, _= self.evaluate(self.test_dataset)
  830. ts_loss_every_nepochs.append(test_eloss)
  831. tr_loss_per_epoch.append(mean_batch_loss)
  832. print("Training complete!")
  833. return tr_loss_per_epoch, tr_batch_loss_hist, ts_loss_every_nepochs, record_dict_cache, targets_cache
  834. def post_process_predictions(self,y_pred:torch.Tensor):
  835. """
  836. Post-processes the predictions by applying a low-pass filter.
  837. """
  838. post_filt_cutoff = self.snn_config.decoder_type["post_filt_cutoff"]
  839. post_filt_order = self.snn_config.decoder_type["post_filt_order"]
  840. y_pred_post = torch.zeros_like(y_pred)
  841. for batch in range(y_pred.shape[0]):
  842. for output in range(y_pred.shape[-1]):
  843. filtered_y_pred = apply_butter_lowpass(y_pred[batch,:, output], cutoff=post_filt_cutoff,
  844. fs=1/self.network.dt, order=post_filt_order)
  845. y_pred_post[batch,:, output] = torch.tensor(filtered_y_pred.copy(), dtype=DTYPE, device=DEVICE)
  846. return y_pred_post
  847. def evaluate(self, eval_dataset, enable_profiling: bool = False,
  848. num_profile_samples: int = 10):
  849. """
  850. Gets the predictions for the test set.
  851. Concatenates the predictions made over the time segments of the test repetition.
  852. The assumption that there is a single test repetition, and it is segmented into smaller windows.
  853. Args:
  854. eval_dataset: Dataset to evaluate on
  855. enable_profiling: If True, profile CPU time and memory for first N samples
  856. num_profile_samples: Number of samples to profile (default: 10)
  857. Returns:
  858. If enable_profiling=False:
  859. (average_loss, record_dict_cache, predicted_values_cache, true_values_cache)
  860. If enable_profiling=True:
  861. (average_loss, record_dict_cache, predicted_values_cache, true_values_cache, profiling_results)
  862. """
  863. batch_size = self.snn_config.training["test_batch_size"]
  864. percent_omission = self.snn_config.task["percent_omission"]
  865. percent_addition = self.snn_config.task["percent_addition"]
  866. percent_misattribution = self.snn_config.task["percent_misattribution"]
  867. jitter_std = self.snn_config.task["jitter_std"]
  868. noise_mode = self.snn_config.task["noise_mode"]
  869. inf_rep_binwidth = self.snn_config.training["inf_rep_binwidth"]
  870. post_process_filt = self.snn_config.decoder_type["post_process_filt"]
  871. if noise_mode=='omission':
  872. percent_noise = percent_omission
  873. elif noise_mode=='addition':
  874. percent_noise = percent_addition
  875. elif noise_mode=='location':
  876. percent_noise = jitter_std
  877. elif noise_mode=='misattribution':
  878. percent_noise = percent_misattribution
  879. else:
  880. percent_noise = 0
  881. snn_network = self.network
  882. test_loader = DataLoader(eval_dataset, batch_size=batch_size, shuffle=False)
  883. # Profiling storage
  884. profiling_samples = []
  885. with torch.no_grad():
  886. snn_network.eval()
  887. loss_per_batch = []
  888. record_dict_cache = []
  889. predicted_values_cache = []
  890. true_values_cache = []
  891. if batch_size ==1:
  892. state_dict = self._initialize_state_dict()
  893. else:
  894. state_dict = None
  895. no_input_counter = 0 # counter for the number of samples with no input spikes
  896. no_input_id = None
  897. for i, (data, label, noisy_data) in enumerate(test_loader):
  898. if noise_mode is None or percent_noise == 0 :
  899. spk_or_emg_in = data.to(DEVICE)
  900. else:
  901. spk_or_emg_in = noisy_data.to(DEVICE)
  902. target_values = label.to(DEVICE)
  903. if torch.sum(spk_or_emg_in) == 0 :
  904. if no_input_id is None or no_input_id==i-1: # check on first all zeros sample
  905. no_input_counter += 1
  906. no_input_id = i
  907. else:
  908. no_input_counter = 0
  909. no_input_id = None
  910. # forward pass
  911. if i > 0 and batch_size==1: # save the last time step
  912. state_dict = self._save_network_state_dict(record_dict)
  913. # if it stays silent for more than 0.5 seconds, reset the state dict
  914. if no_input_counter > int(40/inf_rep_binwidth):
  915. print("No input spikes for more than 10 samples..Resetting the state dict and counter")
  916. state_dict = None
  917. no_input_counter = 0
  918. no_input_id = None
  919. # Profile this sample if enabled and we haven't reached the limit
  920. if enable_profiling and len(profiling_samples) < num_profile_samples:
  921. profile_result = self.profile_snn_inference(spk_or_emg_in, state_dict)
  922. profiling_samples.append(profile_result)
  923. record_dict = profile_result['record_dict']
  924. else:
  925. record_dict = snn_network.forward(spk_or_emg_in, state_dict)
  926. prediction_var = snn_network.prediction_var
  927. predicted_values = record_dict[prediction_var].cpu().detach()
  928. if post_process_filt:
  929. predicted_values = self.post_process_predictions(predicted_values)
  930. loss_value = snn_network.loss_function(predicted_values, target_values)
  931. loss_per_batch.append(loss_value.item())
  932. predicted_values_cache.append(predicted_values)
  933. true_values_cache.append(target_values)
  934. record_dict_cache.append(record_dict)
  935. average_loss = np.sum(loss_per_batch) / len(loss_per_batch)
  936. # Aggregate and return profiling results
  937. if enable_profiling and profiling_samples:
  938. profiling_results = aggregate_snn_profiling_results(
  939. profiling_samples, effective_ops=None, model=snn_network, verbose=True
  940. )
  941. return average_loss, record_dict_cache, predicted_values_cache, true_values_cache, profiling_results
  942. return average_loss, record_dict_cache, predicted_values_cache, true_values_cache
  943. def train_and_evaluate_on_training_set(self):
  944. """
  945. Trains the network on the training dataset. The train function also returns
  946. some test losses these are the loss on the test every n epochs
  947. During offline mode, the batch size is greater than 1.
  948. """
  949. # return tr_loss_per_epoch, tr_batch_loss_hist, ts_loss_every_nepochs, record_dict_cache, targets_cache
  950. tr_loss_per_epoch, tr_batch_loss_hist, ts_loss_every_nepochs, _, _ = self._train()
  951. e_loss_tr, rec_tr, y_pred_tr, y_true_tr = self.evaluate(self.train_dataset)
  952. print(f"Train MSE loss: {e_loss_tr:.4f}")
  953. rec_tr, _ = conc_var_tracker(tracker=rec_tr, is_target_tracker=False)
  954. y_pred_tr = target_tracker_to_np(y_pred_tr)
  955. y_true_tr = target_tracker_to_np(y_true_tr)
  956. return y_pred_tr,y_true_tr,tr_loss_per_epoch,tr_batch_loss_hist,ts_loss_every_nepochs,rec_tr
  957. def get_initial_network_params(self):
  958. """
  959. Gets the initial values of the network parameters to be used for plotting."""
  960. snn_network = self.network
  961. beta_init = snn_network.lif1.beta.detach().clone().cpu().numpy()
  962. alpha_init = snn_network.lif1.alpha.detach().clone().cpu().numpy()
  963. initial_weights = snn_network.fc1.weight.detach().clone().cpu().numpy()
  964. threshold_init = snn_network.lif1.threshold.detach().clone().cpu().numpy()
  965. return beta_init,alpha_init,threshold_init,initial_weights
  966. def evaluate_and_log_on_test_set(self, data_config, num_segments_per_rep:int,
  967. repetition_dur:float,
  968. fold_i:int, mode:str, binwidth,
  969. enable_profiling: bool = False,
  970. num_profile_samples: int = 10):
  971. """
  972. Evaluates the model after training is complete using the test set.
  973. This test set can have a batch size =1 to mimic online inference or
  974. a larger batch size to evaluate the model.
  975. Args:
  976. data_config: Data configuration
  977. num_segments_per_rep: Number of segments per repetition
  978. repetition_dur: Duration of repetition
  979. fold_i: Fold index
  980. mode: 'training' or 'inference'
  981. binwidth: Bin width for reconstruction
  982. enable_profiling: If True, profile CPU time and memory
  983. num_profile_samples: Number of samples to profile
  984. Returns:
  985. If enable_profiling=False:
  986. (y_pred_ts, y_true_ts, rec_ts)
  987. If enable_profiling=True:
  988. (y_pred_ts, y_true_ts, rec_ts, profiling_results)
  989. """
  990. profiling_results = None
  991. if enable_profiling:
  992. _, rec_ts, y_pred_ts, y_true_ts, profiling_results = self.evaluate(
  993. self.test_dataset, enable_profiling=True, num_profile_samples=num_profile_samples
  994. )
  995. else:
  996. _, rec_ts, y_pred_ts, y_true_ts = self.evaluate(self.test_dataset)
  997. rec_ts, _ = conc_var_tracker(tracker=rec_ts, is_target_tracker=False)
  998. #Log the output spikes: here I assume that given the overlap percentage,ol, there is ol% of the spikes from the previous sample
  999. overlap_perc = self.snn_config.task["overlap_perc"]
  1000. output_spike_count_df = sum_output_spikes_per_active_finger(overlap_perc, data_config,
  1001. num_segments_per_rep,
  1002. repetition_dur,
  1003. fold_i, binwidth, self.network,
  1004. mode, rec_ts)
  1005. print(f"Output spike count df\n{tabulate(output_spike_count_df, headers='keys', tablefmt='psql')}")
  1006. use_wandb = self.snn_config.logging["wandb"]["use_wandb"]
  1007. if use_wandb:
  1008. log_spike_count_df(output_spike_count_df, fold_i, binwidth)
  1009. # Compute metrics and create y_df
  1010. y_pred_ts = target_tracker_to_np(y_pred_ts)
  1011. y_true_ts = target_tracker_to_np(y_true_ts)
  1012. if enable_profiling and profiling_results is not None:
  1013. return y_pred_ts, y_true_ts, rec_ts, profiling_results
  1014. return y_pred_ts, y_true_ts, rec_ts
  1015. def calculate_effective_ops(self, record_dict: dict, verbose: bool = False) -> Dict[str, float]:
  1016. """
  1017. Calculate effective operations for SNN based on actual spike activity.
  1018. This accounts for the sparse nature of spiking neural networks.
  1019. Operation Types:
  1020. ----------------
  1021. 1. **SynOps (Synaptic Operations)**: Spike-dependent operations at synapses
  1022. - For encoding layer: Treated as MAC (weight × spike_value) since encoding
  1023. may produce non-binary spike representations
  1024. - For non-encoding layer: Treated as AC (accumulate) only since spikes are
  1025. binary (just add the weight when spike=1, no multiplication needed)
  1026. 2. **MAC Operations**: Multiply-accumulate for membrane/synapse dynamics
  1027. - Membrane decay: beta × mem
  1028. - Synaptic decay: alpha × syn_cur
  1029. - These run every timestep (dense operations)
  1030. 3. **Add Operations**: Simple additions
  1031. - Adding synaptic current to membrane potential
  1032. - Dense operations that run every timestep
  1033. Total Effective FLOPs = SynOps_FLOPs + MAC_ops × 2 + Add_ops
  1034. Where:
  1035. - SynOps_FLOPs = SynOps × 2 (for encoding) or SynOps × 1 (for non-encoding)
  1036. - MAC × 2 because each MAC = 1 multiply + 1 add
  1037. Args:
  1038. record_dict: Dictionary containing recorded network variables from forward pass
  1039. verbose: If True, print detailed breakdown of operations
  1040. Returns:
  1041. Dictionary containing various operation counts and metrics
  1042. """
  1043. ops = {
  1044. 'mac_ops': 0, # Multiply-accumulate operations (dense, membrane/synapse dynamics)
  1045. 'add_ops': 0, # Addition operations (dense)
  1046. 'comparison_ops': 0, # Threshold comparisons
  1047. 'synops_mac': 0, # SynOps counted as MAC (encoding layer) - fc1
  1048. 'synops_ac': 0, # SynOps counted as AC (non-encoding layer) - fc1
  1049. 'synops_ac_fc2': 0, # SynOps for fc2 layer (output spikes -> filter neurons)
  1050. 'total_input_spikes': 0,
  1051. 'total_output_spikes': 0,
  1052. 'total_encoding_spikes': 0,
  1053. }
  1054. # Get batch size and timesteps from recorded data
  1055. input_var = self.network.input_var
  1056. input_data = record_dict[input_var]
  1057. batch_size = input_data.shape[0]
  1058. timesteps = input_data.shape[1]
  1059. topology = self.network.topology
  1060. # Count spikes based on topology
  1061. if topology in [ENC_SHALLOW_SPIKING, ENC_SHALLOW_LEAKY]:
  1062. # For encoding topologies, count encoding layer spikes
  1063. enc_spikes = record_dict[ENC_SPK_OUT]
  1064. ops['total_encoding_spikes'] = torch.sum(enc_spikes).item()
  1065. # Encoding layer operations (EMG to spikes) - DENSE operations
  1066. # Linear layer: enc_fc (identity matrix, but still processes each input)
  1067. ops['mac_ops'] += batch_size * timesteps * self.num_inputs
  1068. # Leaky integrator update for encoding layer
  1069. # beta * mem + current (1 MAC + 1 add per neuron per timestep)
  1070. ops['mac_ops'] += batch_size * timesteps * self.num_inputs
  1071. ops['add_ops'] += batch_size * timesteps * self.num_inputs
  1072. # Threshold comparison for encoding layer
  1073. ops['comparison_ops'] += batch_size * timesteps * self.num_inputs
  1074. # Decoding layer: fc1 operations (only for active spikes due to sparsity)
  1075. # For encoding layer output, spikes may have continuous values, so treat as MAC
  1076. ops['synops_mac'] = ops['total_encoding_spikes'] * self.num_outputs
  1077. else:
  1078. # For non-encoding topologies, count input spikes
  1079. input_spikes = record_dict[SPK_IN]
  1080. ops['total_input_spikes'] = torch.sum(input_spikes).item()
  1081. # Linear layer operations (only for active input spikes)
  1082. # Since spikes are binary (0 or 1), weight × spike = weight when spike=1
  1083. # So this is just an accumulate operation (AC), not multiply-accumulate (MAC)
  1084. ops['synops_ac'] = ops['total_input_spikes'] * self.num_outputs
  1085. # Output layer spikes
  1086. output_spikes = record_dict[SPK_OUT]
  1087. ops['total_output_spikes'] = torch.sum(output_spikes).item()
  1088. # fc2 synaptic operations: output spikes -> filter neurons
  1089. # In SHALLOW_SPIKING_DOUBLE_FILTER, fc2 processes binary output spikes (spk_out)
  1090. # to produce input current for the filter neurons
  1091. # Since spk_out contains binary spikes, this is an AC operation (just accumulate weight)
  1092. if topology == SHALLOW_SPIKING_DOUBLE_FILTER:
  1093. # fc2: output_spikes × num_outputs (filter neurons have same size as output)
  1094. ops['synops_ac_fc2'] = ops['total_output_spikes'] * self.num_outputs
  1095. elif topology == ENC_SHALLOW_SPIKING:
  1096. # For encoding topology, output spikes also feed fc2 to filter neurons
  1097. ops['synops_ac_fc2'] = ops['total_output_spikes'] * self.num_outputs
  1098. # Synaptic and membrane dynamics for output neurons - DENSE operations
  1099. if topology in [SHALLOW_SPIKING_DOUBLE_FILTER, ENC_SHALLOW_SPIKING]:
  1100. # Synaptic current update: alpha * syn_cur (1 MAC per neuron per timestep)
  1101. ops['mac_ops'] += batch_size * timesteps * self.num_outputs
  1102. # Membrane potential update: beta * mem (1 MAC per neuron per timestep)
  1103. ops['mac_ops'] += batch_size * timesteps * self.num_outputs
  1104. # Adding synaptic current to membrane (2 adds per neuron per timestep)
  1105. ops['add_ops'] += batch_size * timesteps * self.num_outputs * 2
  1106. elif topology in [SHALLOW_LEAKY, ENC_SHALLOW_LEAKY]:
  1107. # Leaky neuron: alpha * syn_cur + beta * mem (2 MACs + 2 adds)
  1108. ops['mac_ops'] += batch_size * timesteps * self.num_outputs * 2
  1109. ops['add_ops'] += batch_size * timesteps * self.num_outputs * 2
  1110. # Threshold comparison for output neurons
  1111. ops['comparison_ops'] += batch_size * timesteps * self.num_outputs
  1112. # Filter layers (always dense, not spike-based)
  1113. if hasattr(self.network, 'filter_layers'):
  1114. for _ in self.network.filter_layers:
  1115. # Each filter: beta * mem + input (1 MAC + 1 add per output per timestep)
  1116. ops['mac_ops'] += batch_size * timesteps * self.num_outputs
  1117. ops['add_ops'] += batch_size * timesteps * self.num_outputs
  1118. # Calculate total synaptic operations (SynOps)
  1119. # This is the spike-dependent operation count reported in SNN literature
  1120. # Includes fc1 SynOps (input/encoding -> output) + fc2 SynOps (output -> filter)
  1121. ops['synaptic_ops'] = ops['synops_mac'] + ops['synops_ac'] + ops['synops_ac_fc2']
  1122. # Calculate total FLOPs with proper accounting:
  1123. # - MAC operations = 2 FLOPs each (1 multiply + 1 accumulate)
  1124. # - Add operations = 1 FLOP each
  1125. # - SynOps as MAC (encoding) = 2 FLOPs each
  1126. # - SynOps as AC (non-encoding) = 1 FLOP each (just accumulate)
  1127. # - SynOps fc2 as AC = 1 FLOP each (binary output spikes)
  1128. ops['synops_flops'] = ops['synops_mac'] * 2 + ops['synops_ac'] * 1 + ops['synops_ac_fc2'] * 1
  1129. ops['membrane_synapse_flops'] = ops['mac_ops'] * 2 + ops['add_ops']
  1130. ops['total_flops'] = ops['synops_flops'] + ops['membrane_synapse_flops']
  1131. ops['total_ops'] = ops['total_flops'] + ops['comparison_ops']
  1132. # Calculate sparsity metrics
  1133. total_possible_input_spikes = batch_size * timesteps * self.num_inputs
  1134. total_possible_output_spikes = batch_size * timesteps * self.num_outputs
  1135. if topology in [ENC_SHALLOW_SPIKING, ENC_SHALLOW_LEAKY]:
  1136. ops['encoding_spike_rate'] = ops['total_encoding_spikes'] / total_possible_input_spikes
  1137. else:
  1138. ops['input_spike_rate'] = ops['total_input_spikes'] / total_possible_input_spikes
  1139. ops['output_spike_rate'] = ops['total_output_spikes'] / total_possible_output_spikes
  1140. # Calculate what percentage of total FLOPs are synaptic operations
  1141. ops['synaptic_percentage'] = (ops['synops_flops']) / ops['total_flops'] * 100 if ops['total_flops'] > 0 else 0
  1142. # Store dimensions for reference
  1143. ops['batch_size'] = batch_size
  1144. ops['timesteps'] = timesteps
  1145. ops['num_inputs'] = self.num_inputs
  1146. ops['num_outputs'] = self.num_outputs
  1147. # Calculate per-sample averages (useful for reporting and comparison)
  1148. ops['total_flops_per_sample'] = ops['total_flops'] / batch_size
  1149. ops['synaptic_ops_per_sample'] = ops['synaptic_ops'] / batch_size
  1150. ops['synops_flops_per_sample'] = ops['synops_flops'] / batch_size
  1151. ops['total_ops_per_sample'] = ops['total_ops'] / batch_size
  1152. if verbose:
  1153. print("\n" + "="*70)
  1154. print("SNN Effective Operations Analysis")
  1155. print("="*70)
  1156. print(f"Topology: {topology}")
  1157. print(f"Batch size: {batch_size}, Timesteps: {timesteps}")
  1158. print(f"Inputs: {self.num_inputs}, Outputs: {self.num_outputs}")
  1159. print("-"*70)
  1160. if topology in [ENC_SHALLOW_SPIKING, ENC_SHALLOW_LEAKY]:
  1161. print(f"Encoding spikes: {ops['total_encoding_spikes']:,} "
  1162. f"(rate: {ops['encoding_spike_rate']:.4f})")
  1163. else:
  1164. print(f"Input spikes: {ops['total_input_spikes']:,} "
  1165. f"(rate: {ops['input_spike_rate']:.4f})")
  1166. print(f"Output spikes: {ops['total_output_spikes']:,} "
  1167. f"(rate: {ops['output_spike_rate']:.4f})")
  1168. print("-"*70)
  1169. print("SPIKE-DEPENDENT (Sparse) Operations:")
  1170. print(f" Synaptic operations (SynOps): {ops['synaptic_ops']:,}")
  1171. if ops['synops_mac'] > 0:
  1172. print(f" - fc1 SynOps as MAC (encoding): {ops['synops_mac']:,} -> {ops['synops_mac']*2:,} FLOPs")
  1173. if ops['synops_ac'] > 0:
  1174. print(f" - fc1 SynOps as AC (non-enc): {ops['synops_ac']:,} -> {ops['synops_ac']:,} FLOPs")
  1175. if ops['synops_ac_fc2'] > 0:
  1176. print(f" - fc2 SynOps as AC (out->filt): {ops['synops_ac_fc2']:,} -> {ops['synops_ac_fc2']:,} FLOPs")
  1177. print(f" SynOps FLOPs: {ops['synops_flops']:,}")
  1178. print(f" Per sample: {ops['synaptic_ops_per_sample']:,.0f} SynOps")
  1179. print("-"*70)
  1180. print("DENSE Operations (membrane/synapse dynamics):")
  1181. print(f" MAC operations: {ops['mac_ops']:,} -> {ops['mac_ops']*2:,} FLOPs")
  1182. print(f" Add operations: {ops['add_ops']:,} -> {ops['add_ops']:,} FLOPs")
  1183. print(f" Dense FLOPs: {ops['membrane_synapse_flops']:,}")
  1184. print("-"*70)
  1185. print(f"Comparison operations: {ops['comparison_ops']:,}")
  1186. print("-"*70)
  1187. print(f"TOTAL FLOPs: {ops['total_flops']:,}")
  1188. print(f" = SynOps FLOPs ({ops['synops_flops']:,}) + Dense FLOPs ({ops['membrane_synapse_flops']:,})")
  1189. print(f" Per sample: {ops['total_flops_per_sample']:,.0f} FLOPs")
  1190. print(f" SynOps as % of total FLOPs: {ops['synaptic_percentage']:.1f}%")
  1191. print(f"Total operations (incl. comparisons): {ops['total_ops']:,}")
  1192. print("="*70 + "\n")
  1193. return ops
  1194. def profile_snn_inference(self, input_tensor: torch.Tensor,
  1195. state_dict: dict = None) -> Dict:
  1196. """
  1197. Profile a single forward pass of the SNN model.
  1198. Args:
  1199. input_tensor: Input tensor for forward pass (batch, timesteps, inputs)
  1200. state_dict: Optional state dictionary for maintaining membrane potentials
  1201. Returns:
  1202. Dictionary with profiling metrics for this sample:
  1203. - cpu_memory: CPU memory usage in bytes
  1204. - cpu_time: CPU time in microseconds
  1205. - flops: FLOPs from torch.profiler
  1206. - record_dict: Network output dictionary with spikes and membrane potentials
  1207. """
  1208. result = {
  1209. 'cpu_memory': 0,
  1210. 'cpu_time': 0,
  1211. 'flops': 0,
  1212. 'record_dict': None
  1213. }
  1214. # Profile using torch.profiler
  1215. with profile(
  1216. activities=[ProfilerActivity.CPU],
  1217. record_shapes=True,
  1218. profile_memory=True,
  1219. with_flops=True,
  1220. with_modules=True
  1221. ) as prof:
  1222. record_dict = self.network.forward(input_tensor, state_dict)
  1223. result['record_dict'] = record_dict
  1224. # Extract metrics from profiler
  1225. for event in prof.key_averages():
  1226. if event.flops is not None:
  1227. result['flops'] += event.flops
  1228. result['cpu_time'] += event.cpu_time_total
  1229. result['cpu_memory'] += event.cpu_memory_usage
  1230. return result
  1231. def calculate_theoretical_ops(self, batch_size: int, timesteps: int) -> Dict[str, float]:
  1232. """
  1233. Calculate theoretical maximum operations assuming dense (non-sparse) computation.
  1234. This represents the worst-case scenario where all neurons spike at every timestep (100% spike rate).
  1235. Args:
  1236. batch_size: Batch size for the calculation
  1237. timesteps: Number of timesteps
  1238. Returns:
  1239. Dictionary containing theoretical operation counts
  1240. """
  1241. ops = {
  1242. 'mac_ops': 0,
  1243. 'add_ops': 0,
  1244. 'comparison_ops': 0,
  1245. }
  1246. topology = self.network.topology
  1247. # Input layer operations
  1248. if topology in [ENC_SHALLOW_SPIKING, ENC_SHALLOW_LEAKY]:
  1249. # Encoding layer: EMG to spikes
  1250. # Linear layer (identity matrix): num_inputs multiplications
  1251. ops['mac_ops'] += batch_size * timesteps * self.num_inputs
  1252. # Leaky integrator: beta * mem + current
  1253. ops['mac_ops'] += batch_size * timesteps * self.num_inputs
  1254. ops['add_ops'] += batch_size * timesteps * self.num_inputs
  1255. ops['comparison_ops'] += batch_size * timesteps * self.num_inputs
  1256. # Decoding layer: assuming every neuron spikes at every timestep (worst case)
  1257. ops['mac_ops'] += batch_size * timesteps * self.num_inputs * self.num_outputs
  1258. else:
  1259. # Standard input to fc1: assuming all inputs are active at every timestep
  1260. ops['mac_ops'] += batch_size * timesteps * self.num_inputs * self.num_outputs
  1261. # Output neuron dynamics
  1262. if topology in [SHALLOW_SPIKING_DOUBLE_FILTER, ENC_SHALLOW_SPIKING]:
  1263. # Synaptic current: alpha * syn_cur
  1264. ops['mac_ops'] += batch_size * timesteps * self.num_outputs
  1265. # Membrane potential: beta * mem
  1266. ops['mac_ops'] += batch_size * timesteps * self.num_outputs
  1267. # Additions (syn_cur + input, mem + syn)
  1268. ops['add_ops'] += batch_size * timesteps * self.num_outputs * 2
  1269. elif topology in [SHALLOW_LEAKY, ENC_SHALLOW_LEAKY]:
  1270. # Leaky dynamics: alpha * syn_cur + beta * mem
  1271. ops['mac_ops'] += batch_size * timesteps * self.num_outputs * 2
  1272. ops['add_ops'] += batch_size * timesteps * self.num_outputs * 2
  1273. # Threshold comparisons
  1274. ops['comparison_ops'] += batch_size * timesteps * self.num_outputs
  1275. # Filter layers
  1276. if hasattr(self.network, 'filter_layers'):
  1277. num_filters = len(self.network.filter_layers)
  1278. for _ in range(num_filters):
  1279. ops['mac_ops'] += batch_size * timesteps * self.num_outputs
  1280. ops['add_ops'] += batch_size * timesteps * self.num_outputs
  1281. # Total FLOPs
  1282. ops['total_flops'] = ops['mac_ops'] * 2 + ops['add_ops']
  1283. ops['total_ops'] = ops['total_flops'] + ops['comparison_ops']
  1284. ops['batch_size'] = batch_size
  1285. ops['timesteps'] = timesteps
  1286. ops['assumption'] = 'full_density'
  1287. return ops
  1288. def compare_ops_with_baseline(self, effective_ops,
  1289. baseline_type: str = 'dense_ann') -> Dict[str, float]:
  1290. """
  1291. Compare SNN operations with baseline architectures.
  1292. Args:
  1293. record_dict: Dictionary containing recorded network variables
  1294. baseline_type: Type of baseline ('dense_ann', 'theoretical_snn')
  1295. Returns:
  1296. Dictionary with comparison metrics
  1297. """
  1298. # effective_ops = self.calculate_effective_ops(record_dict, verbose=False)
  1299. batch_size = effective_ops['batch_size']
  1300. timesteps = effective_ops['timesteps']
  1301. comparison = {
  1302. 'effective_ops': effective_ops,
  1303. }
  1304. if baseline_type == 'dense_ann':
  1305. # Dense ANN: all neurons active at every timestep
  1306. dense_ops = batch_size * timesteps * self.num_inputs * self.num_outputs * 2 # fc1
  1307. dense_ops += batch_size * timesteps * self.num_outputs * 2 # bias and activation
  1308. comparison['dense_ann_flops'] = dense_ops
  1309. comparison['reduction_vs_dense'] = dense_ops / effective_ops['total_flops']
  1310. elif baseline_type == 'theoretical_snn':
  1311. theoretical_ops = self.calculate_theoretical_ops(batch_size, timesteps)
  1312. comparison['theoretical_snn_ops'] = theoretical_ops
  1313. comparison['reduction_vs_theoretical'] = (theoretical_ops['total_flops'] /
  1314. effective_ops['total_flops'])
  1315. return comparison
  1316. def create_inference_net(snn_config:SNNConfig, trained_params_file:str, timesteps:int,
  1317. num_inputs:int, num_outputs:int,test_dataset:EMGDataset):
  1318. """
  1319. Creates a new network object using the trained parameters.
  1320. """
  1321. trained_params = load_params_from_pickle(trained_params_file)
  1322. snnreg_model = SnnReg(timesteps,num_inputs,num_outputs, snn_config, None, test_dataset)
  1323. net = snnreg_model.network
  1324. net.fc1.weight = nn.Parameter(trained_params['fc1.weight'])
  1325. net.fc1.bias = nn.Parameter(trained_params['fc1.bias'])
  1326. net.lif1.alpha = nn.Parameter(trained_params['lif1.alpha'])
  1327. net.lif1.beta = nn.Parameter(trained_params['lif1.beta'])
  1328. # net.filter.beta = nn.Parameter(trained_params['filter.beta'])
  1329. return snnreg_model
  1330. def load_params_from_pickle(trained_params_file):
  1331. """
  1332. Loads the trained parameters from a pickle file.
  1333. """
  1334. with open(trained_params_file, 'rb') as f:
  1335. trained_params = pkl.load(f)
  1336. return trained_params
  1337. def log_spike_count_df(output_spike_count_df:pd.DataFrame,fold_i:int, binwidth:float)-> None:
  1338. """
  1339. Logs table of output spike counts to wandb.
  1340. """
  1341. wandb.log({f'Spike Count table Fold {fold_i} {binwidth}': wandb.Table(data=output_spike_count_df,
  1342. columns=output_spike_count_df.columns.tolist())})
  1343. def sum_output_spikes_per_active_finger(overlap_perc:float, data_config:DataConfig,
  1344. num_segments_per_rep:int,
  1345. repetition_dur:float,
  1346. fold_i:int, binwidth:float, snn_network,
  1347. mode:str,
  1348. rec_ts:dict) -> pd.DataFrame:
  1349. """
  1350. Sums the output spikes for each active finger by first reconstructing the whole repetition using the overlapping windows.
  1351. """
  1352. reconst_spikes = reconstruct_rep_from_bins(rec_ts[SPK_OUT], num_segments_per_rep,
  1353. binwidth, overlap_perc,
  1354. snn_network.dt,repetition_dur)
  1355. output_spike_count = torch.sum(reconst_spikes,dim=0)
  1356. output_spike_count_df = pd.DataFrame(output_spike_count, columns=[f'neu_{i}' for i in range(output_spike_count.shape[1])])
  1357. output_spike_count_df[MODE_COL] = mode
  1358. output_spike_count_df[REP_DUR_COL] = reconst_spikes.shape[0]* snn_network.dt
  1359. output_spike_count_df[FOLD_COL] = fold_i
  1360. output_spike_count_df['active_fing_name'] = list(data_config.finger_label_map.keys())
  1361. return output_spike_count_df
  1362. def run_net_single_split(snn_config:SNNConfig, data_config:DataConfig, dataset:EMGDataset,
  1363. rep_on_rep:str, fold_i:int, use_inference_network:bool=False,
  1364. num_profile_samples:int=10,
  1365. enable_profiling:bool=False):
  1366. """
  1367. Runs the network for a single split of the dataset.
  1368. This function is used for both training (or offline) and inference modes (mimicks online inference with batch size=1).
  1369. In training mode, the dataset is split into train and test sets, and the network is trained on the train set then evaluated
  1370. on both train and test sets.
  1371. In inference mode, a new network is created using the trained weights and and other parameters saved in a pickle file is used to make predictions on the test set.
  1372. This new network is created to run on the test set using a smaller binwidth than the one used for training to mimic online inference.
  1373. """
  1374. dataset_rep_1, dataset_rep_2 = create_subsets(data_config, snn_config, dataset)
  1375. test_rep = int(rep_on_rep.split('on')[1])
  1376. neuron_beta_before_training = None # inital tau_mem for the network prior training
  1377. synapse_alpha_before_training = None # inital tau_syn for the network prior training
  1378. neuron_thr_before_training = None
  1379. weights_before_training = None
  1380. y_pred_tr, y_true_tr = None, None
  1381. trained_params_file = generate_trained_params_filename(snn_config, data_config, fold_i)
  1382. trained_params_df = None
  1383. if use_inference_network:
  1384. mode = 'inference'
  1385. test_dataset, _ = assign_datasets_for_train_test(dataset_rep_1, dataset_rep_2, rep_on_rep)
  1386. snnreg_model = create_inference_net(snn_config, trained_params_file,
  1387. dataset.num_steps,dataset.num_inputs,dataset.num_outputs,
  1388. test_dataset)
  1389. snn_network = snnreg_model.network
  1390. binwidth = snn_config.training["inf_rep_binwidth"]
  1391. else:
  1392. mode = 'training'
  1393. test_dataset, train_dataset = assign_datasets_for_train_test(dataset_rep_1, dataset_rep_2, rep_on_rep)
  1394. snnreg_model = SnnReg(dataset.num_steps,dataset.num_inputs,dataset.num_outputs,
  1395. snn_config, train_dataset, test_dataset)
  1396. snn_network = snnreg_model.network
  1397. save_network_statedict_to_file(snn_config, data_config, snn_network)
  1398. neuron_beta_before_training, synapse_alpha_before_training, neuron_thr_before_training, weights_before_training =snnreg_model.get_initial_network_params()
  1399. binwidth = snn_config.training["train_rep_binwidth"]
  1400. y_pred_tr, y_true_tr, tr_loss_hist, loss_per_epoch, ts_loss_hist, rec_tr = snnreg_model.train_and_evaluate_on_training_set()
  1401. trained_params_df = prepare_and_save_trained_params(fold_i, trained_params_file, snn_network,
  1402. neuron_beta_before_training, synapse_alpha_before_training,
  1403. neuron_thr_before_training, weights_before_training,
  1404. binwidth)
  1405. trained_weights = snn_network.fc1.weight.detach().clone().numpy()
  1406. plot_network_summary(snn_config, data_config, dataset, fold_i, neuron_beta_before_training, synapse_alpha_before_training,
  1407. neuron_thr_before_training, weights_before_training, y_true_tr, snn_network, tr_loss_hist,
  1408. loss_per_epoch, ts_loss_hist, rec_tr, trained_weights)
  1409. profiling_results = None
  1410. if enable_profiling:
  1411. eval_result = snnreg_model.evaluate_and_log_on_test_set(data_config, dataset.num_segments,
  1412. dataset.rep_dur, fold_i, mode, binwidth,
  1413. enable_profiling=True,
  1414. num_profile_samples=num_profile_samples)
  1415. if len(eval_result) == 4:
  1416. y_pred_ts, y_true_ts, rec_ts, profiling_results = eval_result
  1417. else:
  1418. y_pred_ts, y_true_ts, rec_ts = eval_result
  1419. else:
  1420. y_pred_ts, y_true_ts, rec_ts = snnreg_model.evaluate_and_log_on_test_set(data_config, dataset.num_segments,
  1421. dataset.rep_dur, fold_i, mode, binwidth)
  1422. # Note: Effective SynOps calculation has been moved to profile_snn_on_unsegmented_data()
  1423. # to only run when profiling is enabled, not on every inference run.
  1424. metrics_df = prepare_snn_metrics_df(y_pred_tr, y_true_tr, y_pred_ts, y_true_ts, test_rep, fold_i, mode, binwidth,
  1425. fing_list=dataset.active_fingers_order[test_rep])
  1426. y_df = prepare_predictions_df(y_pred_ts, y_true_ts, test_rep, fold_i, mode, binwidth,
  1427. dataset.active_fingers_order[test_rep])
  1428. snnplot.plot_pred_vs_true(snn_network, snn_config, data_config, dataset,
  1429. y_pred_ts, y_true_ts, test_rep,
  1430. fold_i, mode=mode)
  1431. return metrics_df, y_df, trained_params_df, rec_ts, profiling_results
  1432. def profile_snn_on_unsegmented_data(snn_config: SNNConfig, data_config: DataConfig,
  1433. unsegmented_dataset: EMGDataset, rep_on_rep: str,
  1434. fold_i: int, num_profile_samples: int = 5,
  1435. ) -> Dict:
  1436. """
  1437. Profile SNN inference on unsegmented data for accurate operation counting.
  1438. This function profiles the SNN on full-length repetitions (unsegmented data)
  1439. to get accurate FLOP counts that match the RNN profiling approach. This avoids
  1440. the overlap-induced redundancy from windowed segmentation.
  1441. Args:
  1442. snn_config: SNN configuration
  1443. data_config: Data configuration
  1444. unsegmented_dataset: Dataset with unsegmented inputs (full repetition)
  1445. rep_on_rep: Training/test split string (e.g., '2on1')
  1446. fold_i: Fold index
  1447. num_profile_samples: Number of samples to profile (default: 5)
  1448. Returns:
  1449. Dictionary with profiling results including:
  1450. - CPU memory and time statistics
  1451. - Effective operations (SynOps, MACs, etc.)
  1452. - Spike counts and rates
  1453. - Comparison with dense ANN baseline
  1454. """
  1455. # Load trained parameters
  1456. trained_params_file = generate_trained_params_filename(snn_config, data_config, fold_i)
  1457. # For unsegmented data, samples are arranged as:
  1458. # [finger0_rep1, finger0_rep2, finger1_rep1, finger1_rep2, ...]
  1459. # So even indices (0,2,4,6,8) are rep1, odd indices (1,3,5,7,9) are rep2
  1460. test_rep = int(rep_on_rep.split('on')[1])
  1461. # Create test subset based on rep_on_rep
  1462. # For unsegmented data with 10 samples (5 fingers x 2 reps):
  1463. # rep1 indices: 0, 2, 4, 6, 8 (even)
  1464. # rep2 indices: 1, 3, 5, 7, 9 (odd)
  1465. n_fingers = data_config.n_ind_fingers
  1466. if test_rep == 1:
  1467. test_indices = np.arange(0, n_fingers * 2, 2) # [0, 2, 4, 6, 8]
  1468. else:
  1469. test_indices = np.arange(1, n_fingers * 2, 2) # [1, 3, 5, 7, 9]
  1470. test_dataset = Subset(unsegmented_dataset, test_indices)
  1471. snnreg_model = create_inference_net(snn_config, trained_params_file,
  1472. unsegmented_dataset.num_steps,
  1473. unsegmented_dataset.num_inputs,
  1474. unsegmented_dataset.num_outputs,
  1475. test_dataset)
  1476. snn_network = snnreg_model.network
  1477. # Profile on unsegmented samples
  1478. test_loader = DataLoader(test_dataset, batch_size=1, shuffle=False)
  1479. profiling_samples = []
  1480. record_dicts = []
  1481. print(f"\n[Unsegmented Profiling] Test dataset size: {len(test_dataset)}")
  1482. print(f"[Unsegmented Profiling] Timesteps per sample: {unsegmented_dataset.num_steps}")
  1483. # import the NeuroBench wrapper to wrap the snnTorch model for usage in the NeuroBench framework
  1484. # import the benchmark class
  1485. neuro_model = SNNTorchModel(snn_network, custom_forward=True, prediction_var=snn_network.prediction_var)
  1486. static_metrics = [Footprint, ConnectionSparsity]
  1487. workload_metrics = [ActivationSparsity, SynapticOperations]
  1488. benchmark = Benchmark(neuro_model, test_loader, [], [], [static_metrics, workload_metrics])
  1489. results = benchmark.run(device=DEVICE)
  1490. # results = {}
  1491. # Add parameter and buffer breakdown to NeuroBench results
  1492. results['parameters'] = [
  1493. (name, list(p.shape), p.requires_grad, p.numel(), p.dtype)
  1494. for name, p in snn_network.named_parameters()
  1495. ]
  1496. results['buffers'] = [
  1497. (name, list(b.shape), b.numel(), b.dtype)
  1498. for name, b in snn_network.named_buffers()
  1499. ]
  1500. with torch.no_grad():
  1501. snn_network.eval()
  1502. for i, (data, label, noisy_data) in enumerate(test_loader):
  1503. if i >= num_profile_samples:
  1504. break
  1505. spk_or_emg_in = data.to(DEVICE)
  1506. # Profile this sample
  1507. profile_result = {
  1508. 'cpu_memory': 0,
  1509. 'cpu_time': 0,
  1510. 'flops': 0,
  1511. 'wall_clock_time': 0,
  1512. 'record_dict': None
  1513. }
  1514. # Measure wall-clock time separately (without profiler overhead)
  1515. wall_start = time.perf_counter()
  1516. _ = snn_network.forward(spk_or_emg_in, None)
  1517. wall_end = time.perf_counter()
  1518. profile_result['wall_clock_time'] = (wall_end - wall_start) * 1e6 # Convert to microseconds
  1519. # Now run with profiler for FLOPs and memory (separate pass)
  1520. with profile(
  1521. activities=[ProfilerActivity.CPU],
  1522. record_shapes=True,
  1523. profile_memory=True,
  1524. with_flops=True,
  1525. with_modules=True
  1526. ) as prof:
  1527. record_dict = snn_network.forward(spk_or_emg_in, None)
  1528. profile_result['record_dict'] = record_dict
  1529. record_dicts.append(record_dict)
  1530. for event in prof.key_averages():
  1531. if event.flops is not None:
  1532. profile_result['flops'] += event.flops
  1533. profile_result['cpu_time'] += event.cpu_time_total
  1534. profile_result['cpu_memory'] += event.cpu_memory_usage
  1535. profiling_samples.append(profile_result)
  1536. print(f" Profiled sample {i+1}/{min(num_profile_samples, len(test_dataset))}")
  1537. # Calculate effective ops on one sample (full repetition)
  1538. if record_dicts:
  1539. # Use first sample for effective ops calculation
  1540. sample_record_dict = record_dicts[0]
  1541. test_ops = snnreg_model.calculate_effective_ops(sample_record_dict, verbose=True)
  1542. comparison = snnreg_model.compare_ops_with_baseline(test_ops, baseline_type='dense_ann')
  1543. else:
  1544. print("[Warning] No samples were profiled - record_dicts is empty")
  1545. test_ops = {}
  1546. comparison = {'dense_ann_flops': 0, 'reduction_vs_dense': 0}
  1547. # Check if we have profiling samples before aggregation
  1548. if not profiling_samples:
  1549. print("[Warning] No profiling samples collected - returning empty results")
  1550. return {
  1551. 'peak_cpu_memory_bytes': 0,
  1552. 'avg_cpu_memory_bytes': 0,
  1553. 'std_cpu_memory_bytes': 0,
  1554. 'peak_cpu_time_us': 0,
  1555. 'avg_cpu_time_us': 0,
  1556. 'std_cpu_time_us': 0,
  1557. 'avg_flops': 0,
  1558. 'avg_flops_torch_profiler': 0,
  1559. 'num_samples_profiled': 0,
  1560. 'data_type': 'unsegmented',
  1561. 'timesteps_per_sample': unsegmented_dataset.num_steps,
  1562. 'dense_ann_flops': 0,
  1563. 'reduction_vs_dense': 0
  1564. }
  1565. # Aggregate profiling results
  1566. profiling_results = aggregate_snn_profiling_results(
  1567. profiling_samples, effective_ops=test_ops, model=snn_network, verbose=True
  1568. )
  1569. # Add comparison metrics
  1570. profiling_results['dense_ann_flops'] = comparison['dense_ann_flops']
  1571. profiling_results['reduction_vs_dense'] = comparison['reduction_vs_dense']
  1572. # Add metadata
  1573. profiling_results['data_type'] = 'unsegmented'
  1574. profiling_results['timesteps_per_sample'] = unsegmented_dataset.num_steps
  1575. # Add NeuroBench benchmark results
  1576. profiling_results['neurobench_results'] = results
  1577. print(f"[Unsegmented Profiling] FLOP reduction vs dense ANN: {comparison['reduction_vs_dense']:.2f}x")
  1578. return profiling_results
  1579. def plot_network_summary(snn_config, data_config, dataset, fold_i, beta_init, alpha_init, threshold_init,
  1580. initial_weights, y_true_tr, snn_network, tr_loss_hist, loss_per_epoch,
  1581. ts_loss_hist, rec_tr, trained_weights):
  1582. """
  1583. Plots the training an test losses per epoch, learnt vs initial weights, input spikes,
  1584. network intermediary variables, and the learnt vs initial weights.
  1585. """
  1586. snnplot.plot_epochs_loss(snn_config, data_config, tr_loss_hist, ts_loss_hist, loss_per_epoch, fold_i)
  1587. snnplot.plot_learnt_wdist(snn_config, data_config, trained_weights, initial_weights)
  1588. # Plot the fold input
  1589. input_var = rec_tr[snn_network.input_var].cpu().detach().numpy()
  1590. snnplot.plot_network_variable(input_var, dataset, snn_network.dt, snn_config,data_config, fold_i, var_name=snn_network.input_var)
  1591. snnplot.network_variables_plot(snn_network, snn_config, data_config, rec_tr, y_true_tr,
  1592. dataset, fold_i, beta_init, alpha_init, threshold_init)
  1593. snnplot.plot_weight_heatmap(snn_network.fc1.weight.detach().cpu().numpy(),snn_config, data_config, fold_i)
  1594. def prepare_and_save_trained_params(fold_i, trained_params_file, snn_network,
  1595. neuron_beta_before_training, synapse_alpha_before_training,
  1596. neuron_thr_before_training, weights_before_training,
  1597. binwidth):
  1598. """
  1599. Prepares the trained parameters dataframe and saves it to a pickle file."""
  1600. trained_params_dict = {
  1601. 'fc1.weight.before': weights_before_training,
  1602. 'lif1.alpha.before': synapse_alpha_before_training,
  1603. 'lif1.beta.before': neuron_beta_before_training,
  1604. 'lif1.threshold.before': neuron_thr_before_training,
  1605. 'fc1.weight': snn_network.fc1.weight.detach().clone(),
  1606. 'fc1.bias': snn_network.fc1.bias.detach().clone(),
  1607. 'lif1.alpha': snn_network.lif1.alpha.detach().clone(),
  1608. 'lif1.beta': snn_network.lif1.beta.detach().clone(),
  1609. 'lif1.threshold': snn_network.lif1.threshold.detach().clone(),
  1610. }
  1611. for name, param in snn_network.named_parameters():
  1612. if param.requires_grad and name not in trained_params_dict.keys():
  1613. trained_params_dict[name] = param.detach().clone()
  1614. print(f"Adding parameter to trained params dict:{name} with values:\n{param.detach().clone()}")
  1615. trained_params_df = prepare_trained_params_df(trained_params_dict, fold_i, binwidth)
  1616. pkl.dump(trained_params_dict, open(trained_params_file, 'wb'))
  1617. return trained_params_df
  1618. def generate_trained_params_filename(snn_config, data_config, fold_i):
  1619. """
  1620. Generates the file path for the trained parameters."""
  1621. batch_size = snn_config.training["train_batch_size"]
  1622. num_iter = snn_config.training["num_iter"]
  1623. learn_tau_mem = snn_config.decoder_type["learn_tau_mem"]
  1624. first_filter_tau = snn_config.decoder_type["first_filter_tau"]
  1625. second_filter_tau = snn_config.decoder_type["second_filter_tau"]
  1626. tau_syn = snn_config.decoder_type["tau_syn"]
  1627. exp_dir = snn_config.task["exp_dir"]
  1628. train_rep_binwidth = snn_config.training["train_rep_binwidth"]
  1629. temp_filename = f"{data_config.mvc}_{exp_dir}_ep_{num_iter}_trained_params_fold_{fold_i}_binwidth_{train_rep_binwidth}_batchsize_{batch_size}_taufilt1_{first_filter_tau}_taufilt2_{second_filter_tau}_tau_syn_{tau_syn}_learn_beta_{learn_tau_mem}.pkl"
  1630. trained_params_file = os.path.join(data_config.tr_weights_data_path, temp_filename)
  1631. return trained_params_file
  1632. def save_network_statedict_to_file(snn_config, data_config, snn_network):
  1633. """
  1634. Saves the model parameters to a file.
  1635. """
  1636. exp_dir = snn_config.task["exp_dir"]
  1637. topology = snn_config.decoder_type["topology"]
  1638. subject = snn_config.task["subject"]
  1639. filename = f"{subject}_mvc_{data_config.mvc}_{exp_dir}_networkparam_topology_{topology}.pt"
  1640. common_path_across_subjects = os.path.dirname(data_config.snn_temp_data_path)
  1641. torch.save(snn_network.state_dict(), os.path.join(common_path_across_subjects, filename))
  1642. def prepare_predictions_df(y_pred_test:np.ndarray, y_true_test:np.ndarray, test_rep:int,
  1643. fold_index:int, mode:str,
  1644. input_dur:float,fing_list:list[str]) -> pd.DataFrame:
  1645. """
  1646. Prepares a dataframe with the predictions and the true values.
  1647. """
  1648. df_columns = [Y_PRED_TEST, Y_TRUE_TEST,TEST_ON_REP, FOLD_COL,
  1649. INPUT_DUR_COL, FING_ORDER_COL, MODE_COL]
  1650. y_pred_df = pd.DataFrame(columns=df_columns)
  1651. y_pred_df.loc[0, Y_PRED_TEST] = y_pred_test
  1652. y_pred_df.loc[0, Y_TRUE_TEST] = y_true_test
  1653. y_pred_df.loc[0, TEST_ON_REP] = test_rep
  1654. y_pred_df.loc[0, FOLD_COL] = fold_index
  1655. y_pred_df.loc[0, MODE_COL] = mode
  1656. y_pred_df.loc[0, INPUT_DUR_COL] = input_dur
  1657. y_pred_df.loc[0, FING_ORDER_COL] = fing_list
  1658. return y_pred_df
  1659. def prepare_trained_params_df(trained_params:dict, fold_i:int, input_dur:float) -> pd.DataFrame:
  1660. """
  1661. Prepares a dataframe with the trained parameters.
  1662. """
  1663. trained_params_df = pd.DataFrame(columns=list(trained_params.keys()))
  1664. for key in trained_params.keys():
  1665. trained_params_df.loc[0, key] = trained_params[key].cpu().numpy() if isinstance(trained_params[key], torch.Tensor) else trained_params[key]
  1666. trained_params_df.loc[0, 'fold'] = fold_i
  1667. trained_params_df.loc[0, 'input_dur'] = input_dur
  1668. return trained_params_df
  1669. def calculate_ops_for_dataset(snn_model: SnnReg, dataset: EMGDataset,
  1670. verbose: bool = True) -> Tuple[Dict, pd.DataFrame]:
  1671. """
  1672. Calculate operations for an entire dataset by running inference and aggregating results.
  1673. Args:
  1674. snn_model: Trained SNN regression model
  1675. dataset: Dataset to evaluate
  1676. verbose: If True, print summary statistics
  1677. Returns:
  1678. Tuple of (aggregated_ops_dict, ops_dataframe)
  1679. """
  1680. batch_size = 1 # Process one sample at a time for accurate spike counting
  1681. data_loader = DataLoader(dataset, batch_size=batch_size, shuffle=False)
  1682. all_ops = []
  1683. with torch.no_grad():
  1684. snn_model.network.eval()
  1685. state_dict = None
  1686. for i, (data, label, _) in enumerate(data_loader):
  1687. spk_or_emg_in = data.to(DEVICE)
  1688. # Forward pass
  1689. if i > 0 and batch_size == 1:
  1690. state_dict = snn_model._save_network_state_dict(record_dict)
  1691. record_dict = snn_model.network.forward(spk_or_emg_in, state_dict)
  1692. # Calculate ops for this sample
  1693. ops = snn_model.calculate_effective_ops(record_dict, verbose=False)
  1694. ops['sample_id'] = i
  1695. all_ops.append(ops)
  1696. # Convert to DataFrame
  1697. ops_df = pd.DataFrame(all_ops)
  1698. # Calculate aggregated statistics
  1699. aggregated = {
  1700. 'total_samples': len(all_ops),
  1701. 'mean_flops_per_sample': ops_df['total_flops'].mean(),
  1702. 'std_flops_per_sample': ops_df['total_flops'].std(),
  1703. 'total_flops': ops_df['total_flops'].sum(),
  1704. 'mean_spike_rate': ops_df.get('encoding_spike_rate', ops_df.get('input_spike_rate', 0)).mean(),
  1705. 'mean_output_spike_rate': ops_df['output_spike_rate'].mean(),
  1706. }
  1707. if verbose:
  1708. print("\n" + "="*70)
  1709. print("Dataset-wide Operations Summary")
  1710. print("="*70)
  1711. print(f"Total samples processed: {aggregated['total_samples']}")
  1712. print(f"Total FLOPs: {aggregated['total_flops']:,.0f}")
  1713. print(f"Mean FLOPs per sample: {aggregated['mean_flops_per_sample']:,.0f} ± {aggregated['std_flops_per_sample']:,.0f}")
  1714. print(f"Mean spike rate: {aggregated['mean_spike_rate']:.4f}")
  1715. print(f"Mean output spike rate: {aggregated['mean_output_spike_rate']:.4f}")
  1716. print("="*70 + "\n")
  1717. return aggregated, ops_df
  1718. def log_ops_to_wandb(ops_dict: Dict, prefix: str = 'snn', fold_i: int = None):
  1719. """
  1720. Log operation counts to Weights & Biases.
  1721. Args:
  1722. ops_dict: Dictionary of operation counts from calculate_effective_ops
  1723. prefix: Prefix for wandb logging keys (e.g., 'snn_train', 'snn_test')
  1724. fold_i: Optional fold index for cross-validation
  1725. """
  1726. log_dict = {
  1727. # Total operations
  1728. f'{prefix}/total_flops': ops_dict['total_flops'],
  1729. f'{prefix}/synaptic_ops': ops_dict['synaptic_ops'],
  1730. f'{prefix}/synaptic_percentage': ops_dict['synaptic_percentage'],
  1731. f'{prefix}/mac_ops': ops_dict['mac_ops'],
  1732. f'{prefix}/add_ops': ops_dict['add_ops'],
  1733. f'{prefix}/comparison_ops': ops_dict['comparison_ops'],
  1734. # Per-sample averages (important for fair comparison)
  1735. f'{prefix}/total_flops_per_sample': ops_dict['total_flops_per_sample'],
  1736. f'{prefix}/synaptic_ops_per_sample': ops_dict['synaptic_ops_per_sample'],
  1737. # Spike rates
  1738. f'{prefix}/output_spike_rate': ops_dict['output_spike_rate'],
  1739. }
  1740. # Add spike rate (encoding or input depending on topology)
  1741. if 'encoding_spike_rate' in ops_dict:
  1742. log_dict[f'{prefix}/encoding_spike_rate'] = ops_dict['encoding_spike_rate']
  1743. elif 'input_spike_rate' in ops_dict:
  1744. log_dict[f'{prefix}/input_spike_rate'] = ops_dict['input_spike_rate']
  1745. if fold_i is not None:
  1746. log_dict['fold'] = fold_i
  1747. wandb.log(log_dict)
  1748. def aggregate_snn_profiling_results(profiling_samples: list, effective_ops: Dict = None,
  1749. model: torch.nn.Module = None, verbose: bool = True) -> Dict:
  1750. """
  1751. Aggregate profiling results across multiple SNN samples.
  1752. Args:
  1753. profiling_samples: List of profiling results from profile_snn_inference
  1754. effective_ops: Dictionary from calculate_effective_ops (optional, for SynOps metrics)
  1755. model: The SNN model for parameter counting (optional)
  1756. verbose: If True, print summary
  1757. Returns:
  1758. Dictionary with aggregated statistics:
  1759. - peak/avg/std for cpu_memory, cpu_time
  1760. - total_params, trainable_params, counted_params (if model provided)
  1761. - effective_ops metrics if provided
  1762. """
  1763. cpu_mem = np.array([s['cpu_memory'] for s in profiling_samples])
  1764. cpu_time = np.array([s['cpu_time'] for s in profiling_samples])
  1765. flops = np.array([s['flops'] for s in profiling_samples])
  1766. wall_clock_time = np.array([s.get('wall_clock_time', 0) for s in profiling_samples])
  1767. results = {
  1768. # Per-sample arrays
  1769. 'cpu_memory_per_sample': cpu_mem.tolist(),
  1770. 'cpu_time_per_sample': cpu_time.tolist(),
  1771. 'flops_per_sample': flops.tolist(),
  1772. 'wall_clock_time_per_sample': wall_clock_time.tolist(),
  1773. # Aggregated CPU memory stats
  1774. 'peak_cpu_memory_bytes': int(np.max(cpu_mem)),
  1775. 'avg_cpu_memory_bytes': float(np.mean(cpu_mem)),
  1776. 'std_cpu_memory_bytes': float(np.std(cpu_mem)),
  1777. # Aggregated CPU time stats (from profiler - includes overhead)
  1778. 'peak_cpu_time_us': int(np.max(cpu_time)),
  1779. 'avg_cpu_time_us': float(np.mean(cpu_time)),
  1780. 'std_cpu_time_us': float(np.std(cpu_time)),
  1781. # Wall-clock time stats (actual inference time without profiler overhead)
  1782. 'peak_wall_clock_time_us': float(np.max(wall_clock_time)),
  1783. 'avg_wall_clock_time_us': float(np.mean(wall_clock_time)),
  1784. 'std_wall_clock_time_us': float(np.std(wall_clock_time)),
  1785. # Aggregated FLOPs stats (from torch.profiler)
  1786. 'peak_flops': int(np.max(flops)),
  1787. 'avg_flops': float(np.mean(flops)),
  1788. 'avg_flops_torch_profiler': float(np.mean(flops)),
  1789. 'num_samples_profiled': len(profiling_samples)
  1790. }
  1791. # Add model parameter counts if model is provided
  1792. if model is not None:
  1793. results['total_params'] = sum(p.numel() for p in model.parameters())
  1794. results['trainable_params'] = sum(p.numel() for p in model.parameters() if p.requires_grad)
  1795. results['counted_params'] = [
  1796. (name, list(param.data.size()), param.requires_grad, param.numel(), param.dtype)
  1797. for name, param in model.named_parameters()
  1798. ]
  1799. # Add effective ops metrics if provided
  1800. if effective_ops is not None:
  1801. # Core operation counts
  1802. results['synaptic_ops'] = effective_ops.get('synaptic_ops', 0)
  1803. results['synaptic_ops_per_sample'] = effective_ops.get('synaptic_ops_per_sample', 0)
  1804. results['synops_mac'] = effective_ops.get('synops_mac', 0) # SynOps as MAC (encoding) - fc1
  1805. results['synops_ac'] = effective_ops.get('synops_ac', 0) # SynOps as AC (non-encoding) - fc1
  1806. results['synops_ac_fc2'] = effective_ops.get('synops_ac_fc2', 0) # SynOps as AC for fc2 (out->filt)
  1807. results['synops_flops'] = effective_ops.get('synops_flops', 0)
  1808. results['synops_flops_per_sample'] = effective_ops.get('synops_flops_per_sample', 0)
  1809. results['mac_ops'] = effective_ops.get('mac_ops', 0) # Dense MAC ops
  1810. results['add_ops'] = effective_ops.get('add_ops', 0) # Dense add ops
  1811. results['membrane_synapse_flops'] = effective_ops.get('membrane_synapse_flops', 0)
  1812. results['comparison_ops'] = effective_ops.get('comparison_ops', 0)
  1813. results['total_effective_flops'] = effective_ops.get('total_flops', 0)
  1814. results['total_effective_flops_per_sample'] = effective_ops.get('total_flops_per_sample', 0)
  1815. results['total_ops'] = effective_ops.get('total_ops', 0)
  1816. results['total_ops_per_sample'] = effective_ops.get('total_ops_per_sample', 0)
  1817. results['synaptic_percentage'] = effective_ops.get('synaptic_percentage', 0)
  1818. # Spike counts and rates
  1819. results['total_input_spikes'] = effective_ops.get('total_input_spikes', 0)
  1820. results['total_output_spikes'] = effective_ops.get('total_output_spikes', 0)
  1821. results['total_encoding_spikes'] = effective_ops.get('total_encoding_spikes', 0)
  1822. if 'encoding_spike_rate' in effective_ops:
  1823. results['encoding_spike_rate'] = effective_ops['encoding_spike_rate']
  1824. if 'input_spike_rate' in effective_ops:
  1825. results['input_spike_rate'] = effective_ops['input_spike_rate']
  1826. results['output_spike_rate'] = effective_ops.get('output_spike_rate', 0)
  1827. # Dimensions
  1828. results['timesteps'] = effective_ops.get('timesteps', 0)
  1829. results['batch_size'] = effective_ops.get('batch_size', 0)
  1830. results['num_inputs'] = effective_ops.get('num_inputs', 0)
  1831. results['num_outputs'] = effective_ops.get('num_outputs', 0)
  1832. # Comparison metrics (if available)
  1833. if 'dense_ann_flops' in effective_ops:
  1834. results['dense_ann_flops'] = effective_ops['dense_ann_flops']
  1835. if 'reduction_vs_dense' in effective_ops:
  1836. results['reduction_vs_dense'] = effective_ops['reduction_vs_dense']
  1837. if verbose:
  1838. print("\n" + "=" * 70)
  1839. print("SNN Profiling Results (Aggregated)")
  1840. print("=" * 70)
  1841. print(f"Samples profiled: {results['num_samples_profiled']}")
  1842. if model is not None:
  1843. print(f"\n[Model Parameters]")
  1844. print(f" Total: {results['total_params']:,}")
  1845. print(f" Trainable: {results['trainable_params']:,}")
  1846. print(f"\n[CPU Memory]")
  1847. print(f" Peak: {results['peak_cpu_memory_bytes'] / 1024**2:.2f} MiB")
  1848. print(f" Average: {results['avg_cpu_memory_bytes'] / 1024**2:.2f} MiB")
  1849. print(f" Std: {results['std_cpu_memory_bytes'] / 1024**2:.2f} MiB")
  1850. print(f"\n[Wall-Clock Time (actual inference)]")
  1851. print(f" Peak: {results['peak_wall_clock_time_us'] / 1000:.3f} ms")
  1852. print(f" Average: {results['avg_wall_clock_time_us'] / 1000:.3f} ms")
  1853. print(f" Std: {results['std_wall_clock_time_us'] / 1000:.3f} ms")
  1854. print(f"\n[FLOPs - torch.profiler (dense ops only)]")
  1855. print(f" Average: {results['avg_flops_torch_profiler']:,.0f}")
  1856. if effective_ops is not None:
  1857. print(f"\n[Effective Ops - SNN-aware FLOP Breakdown]")
  1858. print(f" --- Spike-Dependent (Sparse) ---")
  1859. print(f" SynOps: {results['synaptic_ops']:,}")
  1860. if results.get('synops_mac', 0) > 0:
  1861. print(f" - As MAC (encoding): {results['synops_mac']:,} -> {results['synops_mac']*2:,} FLOPs")
  1862. if results.get('synops_ac', 0) > 0:
  1863. print(f" - As AC (non-enc): {results['synops_ac']:,} -> {results['synops_ac']:,} FLOPs")
  1864. print(f" SynOps FLOPs: {results.get('synops_flops', 0):,}")
  1865. print(f" --- Dense (membrane/synapse) ---")
  1866. print(f" MAC ops: {results['mac_ops']:,} -> {results['mac_ops']*2:,} FLOPs")
  1867. print(f" Add ops: {results['add_ops']:,} -> {results['add_ops']:,} FLOPs")
  1868. print(f" Dense FLOPs: {results.get('membrane_synapse_flops', 0):,}")
  1869. print(f" --- Totals ---")
  1870. print(f" Total Effective FLOPs: {results['total_effective_flops']:,}")
  1871. print(f" Per sample: {results['total_effective_flops_per_sample']:,.0f}")
  1872. print(f" SynOps as % of FLOPs: {results['synaptic_percentage']:.1f}%")
  1873. print(f"\n[Spike Rates]")
  1874. if 'encoding_spike_rate' in results:
  1875. print(f" Encoding spike rate: {results['encoding_spike_rate']:.4f}")
  1876. elif 'input_spike_rate' in results:
  1877. print(f" Input spike rate: {results['input_spike_rate']:.4f}")
  1878. print(f" Output spike rate: {results['output_spike_rate']:.4f}")
  1879. print("=" * 70 + "\n")
  1880. return results

snn.py at commit 00dadd7, no license · at the source

Overview

  1. Institute of Neuroinformatics, University of Zürich and ETH Zürich,Zürich, Switzerland
  2. Department of Bioengineering, Imperial College London,London, UK
Institutions: University of Zurich (Switzerland); Imperial College London (United Kingdom)
Journal: Nature communications, volume 17, issue 1, article 7937
Dates: received 1 September 2025; accepted 1 June 2026; published online 24 June 2026
Type: Research article · Language: English
License: CC BY
Identifiers: DOI 10.1038/s41467-026-74243-1 · PMID 42342667 · PMCID PMC13447840 · OpenAlex W7165739264
Open access: gold, a free copy (OpenAlex)
Status: code verified
Categories: other (modality), extracellular electrophysiology (units, LFP) (modality), human (organism), computational (subfield)
Methods: Connectivity, Machine learning, Statistics, Preprocessing, Graphs, Single-unit activity, calcium imaging, Physiology & signal measures
Keywords: Biomedical engineering, Motor neuron, Electrical and electronic engineering, Computer science
MeSH: Action Potentials*, Fingers*, Muscle, Skeletal*, Neural Networks, Computer*, Electromyography, Humans, Isometric Contraction, Microelectrodes, Motor Neurons (* major topic)
Topic: Muscle activation and electromyography studies (Biomedical Engineering, Engineering), according to OpenAlex
Citations: not cited yet (Europe PMC); 64 references in the paper

Abstract

Assistive technologies for restoring naturalistic finger control require continuous and robust decoding of motor intent, with high accuracy and low latency. Here, we present a spike-based decoding framework that exploits the dynamics of spiking neural networks (SNNs) to efficiently process motor unit activity extracted from high-density intramuscular microelectrode arrays. Using this framework, we demonstrate simultaneous and proportional decoding of individual finger forces from motor unit spike trains during isometric contractions at 15% of maximum voluntary contraction. We systematically evaluated the properties of different SNN decoder configurations, comparing two possible input modalities: physiologically grounded motor unit spike trains and spike-encoded intramuscular EMG signals. Through this comparison, we determined the trade-offs between decoding accuracy, memory footprint, and robustness to input errors. Our results show that lean shallow SNNs are sufficient to decode finger-level motor intent with competitive accuracy, while operating, with minimal memory requirements and without the need for external pre-processing modules. This work provides a practical blueprint for integrating compact, low-power and low-latency SNNs into finger-level force decoding systems, demonstrating how the choice of input representation can be strategically tailored to meet application-specific requirements for accuracy, robustness, and memory efficiency.

Reproduced under the paper's license (CC BY), from the paper cited above.

Repositories

Its files are read in the Code ↔ Paper reader above, with 14 matches between paragraphs and lines of code.

FarahBaracat/snn-fingerforce-decoding

License: none: the authors keep all their rights
State: the link answers, verified on 27 September 2026
Evidence: files inventoried
Commit: 00dadd75490c7d3c32bcf5d1e9a89cc35f2750ff, 23 April 2026
Languages: Python (43), Jupyter (13)
Size: 78 files, 56 scripts
Software Heritage: not archived
Found in: “Code availability”
Holds: README, environment (environment.yml, pyproject.toml), 13 notebooks
Not found: license file, CITATION.cff, tests, continuous integration, documentation
Tools: pandas (30 files), NumPy (28 files), PyTorch (20 files), Matplotlib (14 files), SciPy (11 files), seaborn (11 files), scikit-learn (4 files)
Availability: 1 check, the latest on 27 September 2026: the link answers
  • 27 September 2026: the link answers
57 files

Zenodo 19724261

License: CC-BY-4.0
State: the link answers, verified on 27 September 2026
Evidence: files inventoried
Size: 1 file
Software Heritage: not checked
Found in: the references
Not found: README, license file, CITATION.cff, environment file, tests, continuous integration, documentation
Tools: pandas (30 files), NumPy (28 files), PyTorch (20 files), Matplotlib (14 files), SciPy (11 files), seaborn (11 files), scikit-learn (4 files)
Availability: 1 check, the latest on 27 September 2026: the link answers (HTTP 200)
  • 27 September 2026: the link answers (HTTP 200)
57 files
At the source:

Code availability

The code used to analyze the data and reproduce the main findings is available at https://github.com/FarahBaracat/snn-fingerforce-decoding64.

Reproduced under the paper's license (CC BY), from the paper cited above.

Tracing map

Proposed by the machine: these links were found in the paper and verified at the source, without human review. The map will receive a Zenodo DOI once one of the paper's authors has validated it with their ORCID.

What the map holds:

  • 2 repositories of the authors' code, each at its verified commit, with its license and how the link was found in the paper;
  • 112 scripts, each with its path and the digest of its content;
  • 14 matches between paragraphs of the paper and lines of the code (method lexical-v1);
  • neither the text of the paper nor the code itself.

Its JSON (tracing-map.json) is deposited on Zenodo with its DOI once the map is validated.

Data

Datasets cited

Data availability

The motor unit decomposition data that support the findings of this study are publicly available on Figshare63. The original HD-iEMG recordings are not publicly available due to ethical, privacy considerations and ongoing research, but are available from the corresponding author (FB) upon reasonable request and subject to institutional approval. Source data are provided with this paper.

Reproduced under the paper's license (CC BY), from the paper cited above.

Versions

The history of this record: each version stored by the harvester or made by a correction of its authors or of the maintainers of its code, and what changed in its facts. The texts of the paper (its abstract, its availability statements) are not part of it; versions that changed only those are not listed.

Version 1, 27 September 2026: the first record

Recorded: type, language, journal, volume, issue, pages, dates, 5 authors, 4 keywords, 9 MeSH terms, 3 funders, 55 references.

Cite

This paper

Baracat, F., Grison, A., Farina, D., Indiveri, G., & Donati, E. (2026). Spiking neural network decoders of finger forces from high-density intramuscular microelectrode arrays. Nature communications, 17(1), 7937. https://doi.org/10.1038/s41467-026-74243-1

BibTeX

@article{baracat2026spiking,
author = {Baracat, Farah and Grison, Agnese and Farina, Dario and Indiveri, Giacomo and Donati, Elisa},
title = {{Spiking neural network decoders of finger forces from high-density intramuscular microelectrode arrays}},
journal = {Nature communications},
year = {2026},
month = jun,
volume = {17},
number = {1},
pages = {7937},
publisher = {Nature Publishing Group},
issn = {2041-1723},
doi = {10.1038/s41467-026-74243-1},
url = {https://doi.org/10.1038/s41467-026-74243-1},
pmid = {42342667},
pmcid = {PMC13447840}
}

RIS

TY - JOUR
AU - Baracat, Farah
AU - Grison, Agnese
AU - Farina, Dario
AU - Indiveri, Giacomo
AU - Donati, Elisa
TI - Spiking neural network decoders of finger forces from high-density intramuscular microelectrode arrays
T2 - Nature communications
J2 - Nat Commun
PY - 2026
DA - 2026/06/24
VL - 17
IS - 1
SP - 7937
SN - 2041-1723
PB - Nature Publishing Group
DO - 10.1038/s41467-026-74243-1
UR - https://doi.org/10.1038/s41467-026-74243-1
LA - en
ER -

CSL-JSON

{
"id": "10.1038/s41467-026-74243-1",
"type": "article-journal",
"title": "Spiking neural network decoders of finger forces from high-density intramuscular microelectrode arrays",
"container-title": "Nature communications",
"author": [
{
"family": "Baracat",
"given": "Farah"
},
{
"family": "Grison",
"given": "Agnese"
},
{
"family": "Farina",
"given": "Dario"
},
{
"family": "Indiveri",
"given": "Giacomo"
},
{
"family": "Donati",
"given": "Elisa"
}
],
"container-title-short": "Nat Commun",
"volume": "17",
"issue": "1",
"page": "7937",
"DOI": "10.1038/s41467-026-74243-1",
"PMID": "42342667",
"PMCID": "PMC13447840",
"ISSN": "2041-1723",
"publisher": "Nature Publishing Group",
"URL": "https://doi.org/10.1038/s41467-026-74243-1",
"language": "en",
"issued": {
"date-parts": [
[
2026,
6,
24
]
]
}
}

The tracing map gets a citation of its own once an author has validated it and it has a DOI.

Similar papers

The papers with a page that share the most with this one: the tools found in their code, their categories, datasets, cited references and authors, the rarest counting most.

[1] doi:10.1126/sciadv.aee9425 [code]
Probabilistic inference of homonymous and heteronymous recurrent inhibition in human muscles from large-scale motor neuron recordings.
Journal: Science advances
In common: PyTorch, seaborn, scikit-learn, 4 other tools, 2 references, author Dario Farina
[2] doi:10.1113/jp290395
Spinal motor neuron pools may be partly driven by impulsive common inputs.
Journal: The Journal of physiology
In common: 3 references, author Dario Farina
[3] doi:10.1186/s12984-026-02041-3 [code]
Mental tasks induce common modulations of oscillations in cortex and spinal cord.
Journal: Journal of neuroengineering and rehabilitation
In common: Matplotlib, other, 2 references, author Dario Farina
[4] doi:10.1038/s41467-026-74460-8 [code]
Spike-based alignment learning solves the weight transport problem.
Journal: Nature communications
In common: PyTorch, seaborn, scikit-learn, 4 other tools, 2 references
[5] doi:10.1038/s42003-026-10032-2 [code]
Data-driven mouse motor thalamus model reveals topography and spatial weight scaling govern spindle dynamics.
Journal: Communications biology
In common: seaborn, scikit-learn, pandas, 3 other tools, computational, 1 reference
[6] doi:10.7554/elife.110588 [code]
Opening the black box toward a modular approach to spike sorting.
Journal: eLife
In common: PyTorch, seaborn, scikit-learn, 4 other tools, extracellular electrophysiology (units, LFP)
[7] doi:10.1371/journal.pcbi.1014615 [code]
Toward reliable machine learning models for neural circuit inference: A diagnostic study of CNNs on spike trains.
Journal: PLoS computational biology
In common: PyTorch, seaborn, scikit-learn, 4 other tools, extracellular electrophysiology (units, LFP)
[8] doi:10.1038/s41467-026-75704-3 [code]
A minimal model of working memory in neural systems and neuromorphic circuits.
Journal: Nature communications
In common: SciPy, Matplotlib, NumPy, 3 references
[9] doi:10.1371/journal.pcbi.1014337 [code]
Fast reconstruction of degenerate populations of conductance-based neuron models from spike times.
Journal: PLoS computational biology
In common: PyTorch, seaborn, scikit-learn, 4 other tools, computational
[10] doi:10.1016/j.patter.2026.101590 [code]
Density-based longitudinal neuron tracking in high-density electrophysiological recordings.
Journal: Patterns (New York, N.Y.)
In common: PyTorch, seaborn, scikit-learn, 4 other tools, extracellular electrophysiology (units, LFP)

Contribute

The authors of this paper can claim it, correct its record and validate its tracing map, and the maintainers of its code (its owner, or a public member of its organization) correct what it says of their repository; anyone signed in can ask for its removal. Every request goes to OSCR's own machine, which answers it; your account page follows them.

Sign in with ORCID to claim this paper as one of its authors, correct its record or validate its tracing map: when the paper's metadata lists your ORCID iD, you are recognized at once. Maintainers of its code: sign in with GitHub, then claim the repository on your account page.

Request its removal

To ask OSCR to remove this record, the copies of its authors' scripts or its tracing map, use the removal request page: signed in, you say who you are, what to remove and why, then review and confirm the request. Published rules decide every request (how).

Discussion, reproductions, activity

Discussion: questions and error reports about this paper and its code, from signed-in readers and its authors. It opens with sign-in.

Reproductions: reports from readers who ran the authors' code: what they reproduced, with which environment, commit and data. It opens with sign-in.

Activity: what happens around this paper: new versions of its record, its map's validation, discussions and reproductions. It opens with sign-in.