1use crate::analysis;
12use numpy::{
13 IntoPyArray, PyArray1, PyArray2, PyReadonlyArray1, PyReadonlyArray2, PyUntypedArrayMethods,
14};
15use pyo3::prelude::*;
16use pyo3::types::PyDict;
17
18pub(crate) fn register(m: &Bound<'_, PyModule>) -> PyResult<()> {
20 m.add_function(wrap_pyfunction!(py_spike_times, m)?)?;
22 m.add_function(wrap_pyfunction!(py_isi, m)?)?;
23 m.add_function(wrap_pyfunction!(py_firing_rate, m)?)?;
24 m.add_function(wrap_pyfunction!(py_spike_count, m)?)?;
25 m.add_function(wrap_pyfunction!(py_bin_spike_train, m)?)?;
26 m.add_function(wrap_pyfunction!(py_instantaneous_rate, m)?)?;
27 m.add_function(wrap_pyfunction!(py_psth, m)?)?;
28 m.add_function(wrap_pyfunction!(py_cv_isi, m)?)?;
29 m.add_function(wrap_pyfunction!(py_cv2, m)?)?;
30 m.add_function(wrap_pyfunction!(py_local_variation, m)?)?;
31 m.add_function(wrap_pyfunction!(py_fano_factor, m)?)?;
32 m.add_function(wrap_pyfunction!(py_lempel_ziv_complexity, m)?)?;
33 m.add_function(wrap_pyfunction!(py_permutation_entropy, m)?)?;
34 m.add_function(wrap_pyfunction!(py_hurst_exponent, m)?)?;
35 m.add_function(wrap_pyfunction!(py_approximate_entropy, m)?)?;
36 m.add_function(wrap_pyfunction!(py_sample_entropy, m)?)?;
37 m.add_function(wrap_pyfunction!(py_cross_correlation, m)?)?;
39 m.add_function(wrap_pyfunction!(py_pairwise_correlation, m)?)?;
40 m.add_function(wrap_pyfunction!(py_event_synchronization, m)?)?;
41 m.add_function(wrap_pyfunction!(py_spike_train_coherence, m)?)?;
42 m.add_function(wrap_pyfunction!(py_spike_time_tiling_coefficient, m)?)?;
43 m.add_function(wrap_pyfunction!(py_covariance_matrix, m)?)?;
44 m.add_function(wrap_pyfunction!(py_autocorrelation_time, m)?)?;
45 m.add_function(wrap_pyfunction!(py_noise_correlation, m)?)?;
46 m.add_function(wrap_pyfunction!(py_signal_correlation, m)?)?;
47 m.add_function(wrap_pyfunction!(py_spike_count_covariance, m)?)?;
48 m.add_function(wrap_pyfunction!(py_joint_psth, m)?)?;
49 m.add_function(wrap_pyfunction!(py_coincidence_index, m)?)?;
50 m.add_function(wrap_pyfunction!(py_van_rossum_distance, m)?)?;
52 m.add_function(wrap_pyfunction!(py_victor_purpura_distance, m)?)?;
53 m.add_function(wrap_pyfunction!(py_isi_distance, m)?)?;
54 m.add_function(wrap_pyfunction!(py_spike_distance, m)?)?;
55 m.add_function(wrap_pyfunction!(py_spike_sync, m)?)?;
56 m.add_function(wrap_pyfunction!(py_spike_sync_profile, m)?)?;
57 m.add_function(wrap_pyfunction!(py_spike_profile, m)?)?;
58 m.add_function(wrap_pyfunction!(py_isi_profile, m)?)?;
59 m.add_function(wrap_pyfunction!(py_adaptive_spike_distance, m)?)?;
60 m.add_function(wrap_pyfunction!(py_schreiber_similarity, m)?)?;
61 m.add_function(wrap_pyfunction!(py_hunter_milton_similarity, m)?)?;
62 m.add_function(wrap_pyfunction!(py_earth_movers_distance, m)?)?;
63 m.add_function(wrap_pyfunction!(py_multi_neuron_victor_purpura, m)?)?;
64 m.add_function(wrap_pyfunction!(py_spike_distance_matrix, m)?)?;
65 m.add_function(wrap_pyfunction!(py_mutual_information, m)?)?;
67 m.add_function(wrap_pyfunction!(py_transfer_entropy, m)?)?;
68 m.add_function(wrap_pyfunction!(py_spike_train_entropy, m)?)?;
69 m.add_function(wrap_pyfunction!(py_noise_entropy, m)?)?;
70 m.add_function(wrap_pyfunction!(py_stimulus_specific_information, m)?)?;
71 m.add_function(wrap_pyfunction!(py_kozachenko_leonenko_mi, m)?)?;
72 m.add_function(wrap_pyfunction!(py_pairwise_granger_causality, m)?)?;
74 m.add_function(wrap_pyfunction!(py_conditional_granger_causality, m)?)?;
75 m.add_function(wrap_pyfunction!(py_spectral_granger_causality, m)?)?;
76 m.add_function(wrap_pyfunction!(py_partial_directed_coherence, m)?)?;
77 m.add_function(wrap_pyfunction!(py_directed_transfer_function, m)?)?;
78 m.add_function(wrap_pyfunction!(py_population_vector_decode, m)?)?;
80 m.add_function(wrap_pyfunction!(py_bayesian_decode, m)?)?;
81 m.add_function(wrap_pyfunction!(py_maximum_likelihood_decode, m)?)?;
82 m.add_function(wrap_pyfunction!(py_linear_discriminant_decode, m)?)?;
83 m.add_function(wrap_pyfunction!(py_naive_bayes_decode, m)?)?;
84 m.add_function(wrap_pyfunction!(py_tokenise_spikes, m)?)?;
86 m.add_function(wrap_pyfunction!(py_sinusoidal_position_encode, m)?)?;
87 m.add_function(wrap_pyfunction!(py_scaled_dot_product_attention, m)?)?;
88 m.add_function(wrap_pyfunction!(py_gaussian_attention, m)?)?;
89 m.add_function(wrap_pyfunction!(py_infonce_loss, m)?)?;
90 m.add_function(wrap_pyfunction!(py_functional_connectivity, m)?)?;
92 m.add_function(wrap_pyfunction!(py_unitary_events, m)?)?;
93 m.add_function(wrap_pyfunction!(py_cell_assembly_detection, m)?)?;
94 m.add_function(wrap_pyfunction!(py_synfire_chain_detection, m)?)?;
95 m.add_function(wrap_pyfunction!(py_surrogate_isi_shuffle, m)?)?;
97 m.add_function(wrap_pyfunction!(py_surrogate_dither, m)?)?;
98 m.add_function(wrap_pyfunction!(py_homogeneous_poisson, m)?)?;
99 m.add_function(wrap_pyfunction!(py_gamma_process, m)?)?;
100 m.add_function(wrap_pyfunction!(py_compound_poisson_process, m)?)?;
101 m.add_function(wrap_pyfunction!(py_surrogate_joint_isi, m)?)?;
102 m.add_function(wrap_pyfunction!(py_surrogate_bin_shuffling, m)?)?;
103 m.add_function(wrap_pyfunction!(py_surrogate_spike_train_shifting, m)?)?;
104 m.add_function(wrap_pyfunction!(py_burst_detection, m)?)?;
106 m.add_function(wrap_pyfunction!(py_first_spike_latency, m)?)?;
107 m.add_function(wrap_pyfunction!(py_response_onset, m)?)?;
108 m.add_function(wrap_pyfunction!(py_change_point_detection, m)?)?;
109 m.add_function(wrap_pyfunction!(py_spike_directionality, m)?)?;
111 m.add_function(wrap_pyfunction!(py_spike_train_order, m)?)?;
112 m.add_function(wrap_pyfunction!(py_cubic_higher_order, m)?)?;
113 m.add_function(wrap_pyfunction!(py_power_spectrum, m)?)?;
115 m.add_function(wrap_pyfunction!(py_waveform_width, m)?)?;
117 m.add_function(wrap_pyfunction!(py_waveform_amplitude, m)?)?;
118 m.add_function(wrap_pyfunction!(py_waveform_repolarization_slope, m)?)?;
119 m.add_function(wrap_pyfunction!(py_waveform_recovery_slope, m)?)?;
120 m.add_function(wrap_pyfunction!(py_waveform_halfwidth, m)?)?;
121 m.add_function(wrap_pyfunction!(py_waveform_pt_ratio, m)?)?;
122 m.add_function(wrap_pyfunction!(py_conditional_intensity, m)?)?;
124 m.add_function(wrap_pyfunction!(py_isi_hazard_function, m)?)?;
125 m.add_function(wrap_pyfunction!(py_isi_survivor_function, m)?)?;
126 m.add_function(wrap_pyfunction!(py_renewal_density, m)?)?;
127 m.add_function(wrap_pyfunction!(py_spike_triggered_average, m)?)?;
129 m.add_function(wrap_pyfunction!(py_spike_triggered_covariance, m)?)?;
130 m.add_function(wrap_pyfunction!(py_spatial_information, m)?)?;
131 m.add_function(wrap_pyfunction!(py_place_field_detection, m)?)?;
132 m.add_function(wrap_pyfunction!(py_tuning_curve, m)?)?;
133 m.add_function(wrap_pyfunction!(py_phase_locking_value, m)?)?;
135 m.add_function(wrap_pyfunction!(py_spike_field_coherence, m)?)?;
136 m.add_function(wrap_pyfunction!(py_spike_phase_histogram, m)?)?;
137 m.add_function(wrap_pyfunction!(py_isolation_distance, m)?)?;
139 m.add_function(wrap_pyfunction!(py_l_ratio, m)?)?;
140 m.add_function(wrap_pyfunction!(py_silhouette_score, m)?)?;
141 m.add_function(wrap_pyfunction!(py_d_prime, m)?)?;
142 m.add_function(wrap_pyfunction!(py_isi_violation_rate, m)?)?;
143 m.add_function(wrap_pyfunction!(py_presence_ratio, m)?)?;
144 m.add_function(wrap_pyfunction!(py_amplitude_cutoff, m)?)?;
145 m.add_function(wrap_pyfunction!(py_snr, m)?)?;
146 m.add_function(wrap_pyfunction!(py_nn_hit_rate, m)?)?;
147 m.add_function(wrap_pyfunction!(py_drift_metric, m)?)?;
148 m.add_function(wrap_pyfunction!(py_spike_train_pca, m)?)?;
150 m.add_function(wrap_pyfunction!(py_demixed_pca, m)?)?;
151 m.add_function(wrap_pyfunction!(py_factor_analysis, m)?)?;
152 m.add_function(wrap_pyfunction!(py_pca_components, m)?)?;
153 m.add_function(wrap_pyfunction!(py_demixed_components, m)?)?;
154 m.add_function(wrap_pyfunction!(py_factor_loadings, m)?)?;
155 m.add_function(wrap_pyfunction!(py_gpfa, m)?)?;
157 m.add_function(wrap_pyfunction!(py_gpfa_em, m)?)?;
158 m.add_function(wrap_pyfunction!(py_gpfa_transform, m)?)?;
159 m.add_function(wrap_pyfunction!(py_spade_detect, m)?)?;
161 Ok(())
162}
163
164#[pyfunction]
167#[pyo3(signature = (binary_train, dt=0.001))]
168fn py_spike_times(
169 py: Python<'_>,
170 binary_train: PyReadonlyArray1<'_, i32>,
171 dt: f64,
172) -> Py<PyArray1<f64>> {
173 let data = binary_train.as_slice().unwrap();
174 analysis::basic::spike_times(data, dt)
175 .into_pyarray(py)
176 .into()
177}
178
179#[pyfunction]
180#[pyo3(signature = (binary_train, dt=0.001))]
181fn py_isi(py: Python<'_>, binary_train: PyReadonlyArray1<'_, i32>, dt: f64) -> Py<PyArray1<f64>> {
182 let data = binary_train.as_slice().unwrap();
183 analysis::basic::isi(data, dt).into_pyarray(py).into()
184}
185
186#[pyfunction]
187#[pyo3(signature = (binary_train, dt=0.001))]
188fn py_firing_rate(binary_train: PyReadonlyArray1<'_, i32>, dt: f64) -> f64 {
189 let data = binary_train.as_slice().unwrap();
190 analysis::basic::firing_rate(data, dt)
191}
192
193#[pyfunction]
194fn py_spike_count(binary_train: PyReadonlyArray1<'_, i32>) -> i64 {
195 let data = binary_train.as_slice().unwrap();
196 analysis::basic::spike_count(data)
197}
198
199#[pyfunction]
200#[pyo3(signature = (binary_train, bin_size=10))]
201fn py_bin_spike_train(
202 py: Python<'_>,
203 binary_train: PyReadonlyArray1<'_, i32>,
204 bin_size: usize,
205) -> Py<PyArray1<i64>> {
206 let data = binary_train.as_slice().unwrap();
207 analysis::basic::bin_spike_train(data, bin_size)
208 .into_pyarray(py)
209 .into()
210}
211
212#[pyfunction]
213#[pyo3(signature = (binary_train, dt=0.001, kernel="gaussian", sigma_ms=10.0))]
214fn py_instantaneous_rate(
215 py: Python<'_>,
216 binary_train: PyReadonlyArray1<'_, f64>,
217 dt: f64,
218 kernel: &str,
219 sigma_ms: f64,
220) -> Py<PyArray1<f64>> {
221 let data = binary_train.as_slice().unwrap();
222 analysis::rate::instantaneous_rate(data, dt, kernel, sigma_ms)
223 .into_pyarray(py)
224 .into()
225}
226
227#[pyfunction]
228#[pyo3(signature = (trials, bin_ms=10.0, dt=0.001))]
229fn py_psth(
230 py: Python<'_>,
231 trials: Vec<PyReadonlyArray1<'_, f64>>,
232 bin_ms: f64,
233 dt: f64,
234) -> (Py<PyArray1<f64>>, Py<PyArray1<f64>>) {
235 let vecs: Vec<Vec<f64>> = trials
236 .iter()
237 .map(|t| t.as_slice().unwrap().to_vec())
238 .collect();
239 let (rates, centers) = analysis::rate::psth(&vecs, bin_ms, dt);
240 (
241 rates.into_pyarray(py).into(),
242 centers.into_pyarray(py).into(),
243 )
244}
245
246#[pyfunction]
247#[pyo3(signature = (binary_train, dt=0.001))]
248fn py_cv_isi(binary_train: PyReadonlyArray1<'_, i32>, dt: f64) -> f64 {
249 let data = binary_train.as_slice().unwrap();
250 analysis::variability::cv_isi(data, dt)
251}
252
253#[pyfunction]
254#[pyo3(signature = (binary_train, dt=0.001))]
255fn py_cv2(binary_train: PyReadonlyArray1<'_, i32>, dt: f64) -> f64 {
256 let data = binary_train.as_slice().unwrap();
257 analysis::variability::cv2(data, dt)
258}
259
260#[pyfunction]
261#[pyo3(signature = (binary_train, dt=0.001))]
262fn py_local_variation(binary_train: PyReadonlyArray1<'_, i32>, dt: f64) -> f64 {
263 let data = binary_train.as_slice().unwrap();
264 analysis::variability::local_variation(data, dt)
265}
266
267#[pyfunction]
268#[pyo3(signature = (binary_train, window_ms=50.0, dt=0.001))]
269fn py_fano_factor(binary_train: PyReadonlyArray1<'_, i32>, window_ms: f64, dt: f64) -> f64 {
270 let data = binary_train.as_slice().unwrap();
271 analysis::variability::fano_factor(data, window_ms, dt)
272}
273
274#[pyfunction]
275fn py_lempel_ziv_complexity(binary_train: PyReadonlyArray1<'_, i32>) -> f64 {
276 let data = binary_train.as_slice().unwrap();
277 analysis::variability::lempel_ziv_complexity(data)
278}
279
280#[pyfunction]
281#[pyo3(signature = (binary_train, order=3, delay=1))]
282fn py_permutation_entropy(
283 binary_train: PyReadonlyArray1<'_, i32>,
284 order: usize,
285 delay: usize,
286) -> f64 {
287 let data = binary_train.as_slice().unwrap();
288 analysis::variability::permutation_entropy(data, order, delay)
289}
290
291#[pyfunction]
292#[pyo3(signature = (binary_train, min_window=10))]
293fn py_hurst_exponent(binary_train: PyReadonlyArray1<'_, i32>, min_window: usize) -> f64 {
294 let data = binary_train.as_slice().unwrap();
295 analysis::variability::hurst_exponent(data, min_window)
296}
297
298#[pyfunction]
299#[pyo3(signature = (binary_train, m=2, r_factor=0.2))]
300fn py_approximate_entropy(binary_train: PyReadonlyArray1<'_, i32>, m: usize, r_factor: f64) -> f64 {
301 let data = binary_train.as_slice().unwrap();
302 analysis::variability::approximate_entropy(data, m, r_factor)
303}
304
305#[pyfunction]
306#[pyo3(signature = (binary_train, m=2, r_factor=0.2))]
307fn py_sample_entropy(binary_train: PyReadonlyArray1<'_, i32>, m: usize, r_factor: f64) -> f64 {
308 let data = binary_train.as_slice().unwrap();
309 analysis::variability::sample_entropy(data, m, r_factor)
310}
311
312#[pyfunction]
315#[pyo3(signature = (train_a, train_b, max_lag_ms=50.0, dt=0.001))]
316fn py_cross_correlation(
317 py: Python<'_>,
318 train_a: PyReadonlyArray1<'_, i32>,
319 train_b: PyReadonlyArray1<'_, i32>,
320 max_lag_ms: f64,
321 dt: f64,
322) -> (Py<PyArray1<f64>>, Py<PyArray1<f64>>) {
323 let a = train_a.as_slice().unwrap();
324 let b = train_b.as_slice().unwrap();
325 let (cc, lags) = analysis::correlation::cross_correlation(a, b, max_lag_ms, dt);
326 (cc.into_pyarray(py).into(), lags.into_pyarray(py).into())
327}
328
329#[pyfunction]
330#[pyo3(signature = (trains, dt=0.001))]
331fn py_pairwise_correlation(
332 py: Python<'_>,
333 trains: Vec<PyReadonlyArray1<'_, i32>>,
334 dt: f64,
335) -> Py<PyArray2<f64>> {
336 let vecs: Vec<Vec<i32>> = trains
337 .iter()
338 .map(|t| t.as_slice().unwrap().to_vec())
339 .collect();
340 let refs: Vec<&[i32]> = vecs.iter().map(|v| v.as_slice()).collect();
341 let mat = analysis::correlation::pairwise_correlation(&refs, dt);
342 let n = mat.len();
343 let flat: Vec<f64> = mat.into_iter().flatten().collect();
344 numpy::PyArray2::from_vec2(py, &flat.chunks(n).map(|c| c.to_vec()).collect::<Vec<_>>())
345 .unwrap()
346 .into()
347}
348
349#[pyfunction]
350#[pyo3(signature = (train_a, train_b, dt=0.001, tau_ms=5.0))]
351fn py_event_synchronization(
352 train_a: PyReadonlyArray1<'_, i32>,
353 train_b: PyReadonlyArray1<'_, i32>,
354 dt: f64,
355 tau_ms: f64,
356) -> f64 {
357 let a = train_a.as_slice().unwrap();
358 let b = train_b.as_slice().unwrap();
359 analysis::correlation::event_synchronization(a, b, dt, tau_ms)
360}
361
362#[pyfunction]
363#[pyo3(signature = (train_a, train_b, dt=0.001))]
364fn py_spike_train_coherence(
365 py: Python<'_>,
366 train_a: PyReadonlyArray1<'_, i32>,
367 train_b: PyReadonlyArray1<'_, i32>,
368 dt: f64,
369) -> (Py<PyArray1<f64>>, Py<PyArray1<f64>>) {
370 let a = train_a.as_slice().unwrap();
371 let b = train_b.as_slice().unwrap();
372 let (coh, freqs) = analysis::correlation::spike_train_coherence(a, b, dt);
373 (coh.into_pyarray(py).into(), freqs.into_pyarray(py).into())
374}
375
376#[pyfunction]
377#[pyo3(signature = (train_a, train_b, dt=0.001, delta_ms=5.0))]
378fn py_spike_time_tiling_coefficient(
379 train_a: PyReadonlyArray1<'_, i32>,
380 train_b: PyReadonlyArray1<'_, i32>,
381 dt: f64,
382 delta_ms: f64,
383) -> f64 {
384 let a = train_a.as_slice().unwrap();
385 let b = train_b.as_slice().unwrap();
386 analysis::correlation::spike_time_tiling_coefficient(a, b, dt, delta_ms)
387}
388
389#[pyfunction]
390#[pyo3(signature = (trains, bin_size=10))]
391fn py_covariance_matrix(
392 py: Python<'_>,
393 trains: Vec<PyReadonlyArray1<'_, i32>>,
394 bin_size: usize,
395) -> Py<PyArray2<f64>> {
396 let vecs: Vec<Vec<i32>> = trains
397 .iter()
398 .map(|t| t.as_slice().unwrap().to_vec())
399 .collect();
400 let refs: Vec<&[i32]> = vecs.iter().map(|v| v.as_slice()).collect();
401 let mat = analysis::correlation::covariance_matrix(&refs, bin_size);
402 let n = mat.len();
403 let rows: Vec<Vec<f64>> = mat.into_iter().collect();
404 numpy::PyArray2::from_vec2(py, &rows)
405 .unwrap_or_else(|_| numpy::PyArray2::zeros(py, [n, n], false))
406 .into()
407}
408
409#[pyfunction]
410#[pyo3(signature = (binary_train, dt=0.001, max_lag_ms=100.0))]
411fn py_autocorrelation_time(
412 binary_train: PyReadonlyArray1<'_, i32>,
413 dt: f64,
414 max_lag_ms: f64,
415) -> f64 {
416 let data = binary_train.as_slice().unwrap();
417 analysis::correlation::autocorrelation_time(data, dt, max_lag_ms)
418}
419
420#[pyfunction]
421#[pyo3(signature = (trains, bin_size=50))]
422fn py_noise_correlation(
423 py: Python<'_>,
424 trains: Vec<PyReadonlyArray1<'_, i32>>,
425 bin_size: usize,
426) -> Py<PyArray2<f64>> {
427 let vecs: Vec<Vec<i32>> = trains
428 .iter()
429 .map(|t| t.as_slice().unwrap().to_vec())
430 .collect();
431 let refs: Vec<&[i32]> = vecs.iter().map(|v| v.as_slice()).collect();
432 let mat = analysis::correlation::noise_correlation(&refs, bin_size);
433 let rows: Vec<Vec<f64>> = mat.into_iter().collect();
434 let n = rows.len();
435 numpy::PyArray2::from_vec2(py, &rows)
436 .unwrap_or_else(|_| numpy::PyArray2::zeros(py, [n, n], false))
437 .into()
438}
439
440#[pyfunction]
441#[pyo3(signature = (trains, bin_size=50))]
442fn py_signal_correlation(
443 py: Python<'_>,
444 trains: Vec<PyReadonlyArray1<'_, i32>>,
445 bin_size: usize,
446) -> Py<PyArray2<f64>> {
447 let vecs: Vec<Vec<i32>> = trains
448 .iter()
449 .map(|t| t.as_slice().unwrap().to_vec())
450 .collect();
451 let refs: Vec<&[i32]> = vecs.iter().map(|v| v.as_slice()).collect();
452 let mat = analysis::correlation::signal_correlation(&refs, bin_size);
453 let rows: Vec<Vec<f64>> = mat.into_iter().collect();
454 let n = rows.len();
455 numpy::PyArray2::from_vec2(py, &rows)
456 .unwrap_or_else(|_| numpy::PyArray2::zeros(py, [n, n], false))
457 .into()
458}
459
460#[pyfunction]
461#[pyo3(signature = (trains, window=50))]
462fn py_spike_count_covariance(
463 py: Python<'_>,
464 trains: Vec<PyReadonlyArray1<'_, i32>>,
465 window: usize,
466) -> Py<PyArray2<f64>> {
467 let vecs: Vec<Vec<i32>> = trains
468 .iter()
469 .map(|t| t.as_slice().unwrap().to_vec())
470 .collect();
471 let refs: Vec<&[i32]> = vecs.iter().map(|v| v.as_slice()).collect();
472 let mat = analysis::correlation::spike_count_covariance(&refs, window);
473 let rows: Vec<Vec<f64>> = mat.into_iter().collect();
474 let n = rows.len();
475 numpy::PyArray2::from_vec2(py, &rows)
476 .unwrap_or_else(|_| numpy::PyArray2::zeros(py, [n, n], false))
477 .into()
478}
479
480#[pyfunction]
481#[pyo3(signature = (train_a, train_b, bin_size=10))]
482fn py_joint_psth(
483 py: Python<'_>,
484 train_a: PyReadonlyArray1<'_, i32>,
485 train_b: PyReadonlyArray1<'_, i32>,
486 bin_size: usize,
487) -> Py<PyArray2<f64>> {
488 let a = train_a.as_slice().unwrap();
489 let b = train_b.as_slice().unwrap();
490 let (flat, n) = analysis::correlation::joint_psth(a, b, bin_size);
491 if n == 0 {
492 return numpy::PyArray2::zeros(py, [0, 0], false).into();
493 }
494 let rows: Vec<Vec<f64>> = flat.chunks(n).map(|c| c.to_vec()).collect();
495 numpy::PyArray2::from_vec2(py, &rows).unwrap().into()
496}
497
498#[pyfunction]
499#[pyo3(signature = (train_a, train_b, dt=0.001, delta_ms=2.0))]
500fn py_coincidence_index(
501 train_a: PyReadonlyArray1<'_, i32>,
502 train_b: PyReadonlyArray1<'_, i32>,
503 dt: f64,
504 delta_ms: f64,
505) -> f64 {
506 let a = train_a.as_slice().unwrap();
507 let b = train_b.as_slice().unwrap();
508 analysis::correlation::coincidence_index(a, b, dt, delta_ms)
509}
510
511#[pyfunction]
514#[pyo3(signature = (train_a, train_b, dt=0.001, tau_ms=10.0))]
515fn py_van_rossum_distance(
516 train_a: PyReadonlyArray1<'_, i32>,
517 train_b: PyReadonlyArray1<'_, i32>,
518 dt: f64,
519 tau_ms: f64,
520) -> f64 {
521 let a = train_a.as_slice().unwrap();
522 let b = train_b.as_slice().unwrap();
523 analysis::distance::van_rossum_distance(a, b, dt, tau_ms)
524}
525
526#[pyfunction]
527#[pyo3(signature = (times_a, times_b, cost_per_s=1000.0))]
528fn py_victor_purpura_distance(
529 times_a: PyReadonlyArray1<'_, f64>,
530 times_b: PyReadonlyArray1<'_, f64>,
531 cost_per_s: f64,
532) -> f64 {
533 let a = times_a.as_slice().unwrap();
534 let b = times_b.as_slice().unwrap();
535 analysis::distance::victor_purpura_distance(a, b, cost_per_s)
536}
537
538#[pyfunction]
539#[pyo3(signature = (train_a, train_b, dt=0.001))]
540fn py_isi_distance(
541 train_a: PyReadonlyArray1<'_, i32>,
542 train_b: PyReadonlyArray1<'_, i32>,
543 dt: f64,
544) -> f64 {
545 let a = train_a.as_slice().unwrap();
546 let b = train_b.as_slice().unwrap();
547 analysis::distance::isi_distance(a, b, dt)
548}
549
550#[pyfunction]
551#[pyo3(signature = (times_a, times_b, t_start=0.0, t_end=1.0))]
552fn py_spike_distance(
553 times_a: PyReadonlyArray1<'_, f64>,
554 times_b: PyReadonlyArray1<'_, f64>,
555 t_start: f64,
556 t_end: f64,
557) -> f64 {
558 let a = times_a.as_slice().unwrap();
559 let b = times_b.as_slice().unwrap();
560 analysis::distance::spike_distance(a, b, t_start, t_end)
561}
562
563#[pyfunction]
564#[pyo3(signature = (times_a, times_b, t_start=0.0, t_end=1.0))]
565fn py_spike_sync(
566 times_a: PyReadonlyArray1<'_, f64>,
567 times_b: PyReadonlyArray1<'_, f64>,
568 t_start: f64,
569 t_end: f64,
570) -> f64 {
571 let a = times_a.as_slice().unwrap();
572 let b = times_b.as_slice().unwrap();
573 analysis::distance::spike_sync(a, b, t_start, t_end)
574}
575
576#[pyfunction]
577#[pyo3(signature = (times_a, times_b, n_bins=50, t_start=0.0, t_end=1.0))]
578fn py_spike_sync_profile(
579 py: Python<'_>,
580 times_a: PyReadonlyArray1<'_, f64>,
581 times_b: PyReadonlyArray1<'_, f64>,
582 n_bins: usize,
583 t_start: f64,
584 t_end: f64,
585) -> Py<PyArray1<f64>> {
586 let a = times_a.as_slice().unwrap();
587 let b = times_b.as_slice().unwrap();
588 analysis::distance::spike_sync_profile(a, b, n_bins, t_start, t_end)
589 .into_pyarray(py)
590 .into()
591}
592
593#[pyfunction]
594#[pyo3(signature = (times_a, times_b, n_bins=50, t_start=0.0, t_end=1.0))]
595fn py_spike_profile(
596 py: Python<'_>,
597 times_a: PyReadonlyArray1<'_, f64>,
598 times_b: PyReadonlyArray1<'_, f64>,
599 n_bins: usize,
600 t_start: f64,
601 t_end: f64,
602) -> Py<PyArray1<f64>> {
603 let a = times_a.as_slice().unwrap();
604 let b = times_b.as_slice().unwrap();
605 analysis::distance::spike_profile(a, b, n_bins, t_start, t_end)
606 .into_pyarray(py)
607 .into()
608}
609
610#[pyfunction]
611#[pyo3(signature = (binary_train_a, binary_train_b, dt=0.001, n_bins=50))]
612fn py_isi_profile(
613 py: Python<'_>,
614 binary_train_a: PyReadonlyArray1<'_, i32>,
615 binary_train_b: PyReadonlyArray1<'_, i32>,
616 dt: f64,
617 n_bins: usize,
618) -> Py<PyArray1<f64>> {
619 let a = binary_train_a.as_slice().unwrap();
620 let b = binary_train_b.as_slice().unwrap();
621 analysis::distance::isi_profile(a, b, dt, n_bins)
622 .into_pyarray(py)
623 .into()
624}
625
626#[pyfunction]
627#[pyo3(signature = (times_a, times_b, t_start=0.0, t_end=1.0, cost=0.0))]
628fn py_adaptive_spike_distance(
629 times_a: PyReadonlyArray1<'_, f64>,
630 times_b: PyReadonlyArray1<'_, f64>,
631 t_start: f64,
632 t_end: f64,
633 cost: f64,
634) -> f64 {
635 let a = times_a.as_slice().unwrap();
636 let b = times_b.as_slice().unwrap();
637 analysis::distance::adaptive_spike_distance(a, b, t_start, t_end, cost)
638}
639
640#[pyfunction]
641#[pyo3(signature = (train_a, train_b, dt=0.001, sigma_ms=5.0))]
642fn py_schreiber_similarity(
643 train_a: PyReadonlyArray1<'_, i32>,
644 train_b: PyReadonlyArray1<'_, i32>,
645 dt: f64,
646 sigma_ms: f64,
647) -> f64 {
648 let a = train_a.as_slice().unwrap();
649 let b = train_b.as_slice().unwrap();
650 analysis::distance::schreiber_similarity(a, b, dt, sigma_ms)
651}
652
653#[pyfunction]
654#[pyo3(signature = (times_a, times_b, dt_max=0.01))]
655fn py_hunter_milton_similarity(
656 times_a: PyReadonlyArray1<'_, f64>,
657 times_b: PyReadonlyArray1<'_, f64>,
658 dt_max: f64,
659) -> f64 {
660 let a = times_a.as_slice().unwrap();
661 let b = times_b.as_slice().unwrap();
662 analysis::distance::hunter_milton_similarity(a, b, dt_max)
663}
664
665#[pyfunction]
666#[pyo3(signature = (times_a, times_b, t_start=0.0, t_end=1.0, n_bins=100))]
667fn py_earth_movers_distance(
668 times_a: PyReadonlyArray1<'_, f64>,
669 times_b: PyReadonlyArray1<'_, f64>,
670 t_start: f64,
671 t_end: f64,
672 n_bins: usize,
673) -> f64 {
674 let a = times_a.as_slice().unwrap();
675 let b = times_b.as_slice().unwrap();
676 analysis::distance::earth_movers_distance(a, b, t_start, t_end, n_bins)
677}
678
679#[pyfunction]
680#[pyo3(signature = (spike_times_list, cost_per_s=1000.0))]
681fn py_multi_neuron_victor_purpura(
682 py: Python<'_>,
683 spike_times_list: Vec<PyReadonlyArray1<'_, f64>>,
684 cost_per_s: f64,
685) -> Py<PyArray2<f64>> {
686 let vecs: Vec<Vec<f64>> = spike_times_list
687 .iter()
688 .map(|t| t.as_slice().unwrap().to_vec())
689 .collect();
690 let refs: Vec<&[f64]> = vecs.iter().map(|v| v.as_slice()).collect();
691 let mat = analysis::distance::multi_neuron_victor_purpura(&refs, cost_per_s);
692 let rows: Vec<Vec<f64>> = mat.into_iter().collect();
693 let n = rows.len();
694 numpy::PyArray2::from_vec2(py, &rows)
695 .unwrap_or_else(|_| numpy::PyArray2::zeros(py, [n, n], false))
696 .into()
697}
698
699#[pyfunction]
700#[pyo3(signature = (spike_times_list, metric="spike_distance", t_start=0.0, t_end=1.0))]
701fn py_spike_distance_matrix(
702 py: Python<'_>,
703 spike_times_list: Vec<PyReadonlyArray1<'_, f64>>,
704 metric: &str,
705 t_start: f64,
706 t_end: f64,
707) -> Py<PyArray2<f64>> {
708 let vecs: Vec<Vec<f64>> = spike_times_list
709 .iter()
710 .map(|t| t.as_slice().unwrap().to_vec())
711 .collect();
712 let refs: Vec<&[f64]> = vecs.iter().map(|v| v.as_slice()).collect();
713 let mat = analysis::distance::spike_distance_matrix(&refs, metric, t_start, t_end);
714 let rows: Vec<Vec<f64>> = mat.into_iter().collect();
715 let n = rows.len();
716 numpy::PyArray2::from_vec2(py, &rows)
717 .unwrap_or_else(|_| numpy::PyArray2::zeros(py, [n, n], false))
718 .into()
719}
720
721#[pyfunction]
724#[pyo3(signature = (train_a, train_b, bin_size=10))]
725fn py_mutual_information(
726 train_a: PyReadonlyArray1<'_, i32>,
727 train_b: PyReadonlyArray1<'_, i32>,
728 bin_size: usize,
729) -> f64 {
730 let a = train_a.as_slice().unwrap();
731 let b = train_b.as_slice().unwrap();
732 analysis::information::mutual_information(a, b, bin_size)
733}
734
735#[pyfunction]
736#[pyo3(signature = (source, target, bin_size=10, lag=1))]
737fn py_transfer_entropy(
738 source: PyReadonlyArray1<'_, i32>,
739 target: PyReadonlyArray1<'_, i32>,
740 bin_size: usize,
741 lag: usize,
742) -> f64 {
743 let s = source.as_slice().unwrap();
744 let t = target.as_slice().unwrap();
745 analysis::information::transfer_entropy(s, t, bin_size, lag)
746}
747
748#[pyfunction]
749#[pyo3(signature = (binary_train, bin_size=10, word_length=4))]
750fn py_spike_train_entropy(
751 binary_train: PyReadonlyArray1<'_, i32>,
752 bin_size: usize,
753 word_length: usize,
754) -> f64 {
755 let data = binary_train.as_slice().unwrap();
756 analysis::information::spike_train_entropy(data, bin_size, word_length)
757}
758
759#[pyfunction]
760#[pyo3(signature = (binary_train, n_trials=10, bin_size=10, word_length=4))]
761fn py_noise_entropy(
762 binary_train: PyReadonlyArray1<'_, i32>,
763 n_trials: usize,
764 bin_size: usize,
765 word_length: usize,
766) -> f64 {
767 let data = binary_train.as_slice().unwrap();
768 analysis::information::noise_entropy(data, n_trials, bin_size, word_length)
769}
770
771#[pyfunction]
772fn py_stimulus_specific_information(
773 spike_counts: PyReadonlyArray1<'_, f64>,
774 stimulus_ids: PyReadonlyArray1<'_, i64>,
775) -> f64 {
776 let counts = spike_counts.as_slice().unwrap();
777 let ids = stimulus_ids.as_slice().unwrap();
778 analysis::information::stimulus_specific_information(counts, ids)
779}
780
781#[pyfunction]
782#[pyo3(signature = (x, y, k=3))]
783fn py_kozachenko_leonenko_mi(
784 x: PyReadonlyArray1<'_, f64>,
785 y: PyReadonlyArray1<'_, f64>,
786 k: usize,
787) -> f64 {
788 let xd = x.as_slice().unwrap();
789 let yd = y.as_slice().unwrap();
790 analysis::information::kozachenko_leonenko_mi(xd, yd, k)
791}
792
793#[pyfunction]
796#[pyo3(signature = (source, target, bin_size=10, order=5))]
797fn py_pairwise_granger_causality(
798 source: PyReadonlyArray1<'_, i32>,
799 target: PyReadonlyArray1<'_, i32>,
800 bin_size: usize,
801 order: usize,
802) -> f64 {
803 let s = source.as_slice().unwrap();
804 let t = target.as_slice().unwrap();
805 analysis::causality::pairwise_granger_causality(s, t, bin_size, order)
806}
807
808#[pyfunction]
809#[pyo3(signature = (source, target, condition, bin_size=10, order=5))]
810fn py_conditional_granger_causality(
811 source: PyReadonlyArray1<'_, i32>,
812 target: PyReadonlyArray1<'_, i32>,
813 condition: PyReadonlyArray1<'_, i32>,
814 bin_size: usize,
815 order: usize,
816) -> f64 {
817 let s = source.as_slice().unwrap();
818 let t = target.as_slice().unwrap();
819 let c = condition.as_slice().unwrap();
820 analysis::causality::conditional_granger_causality(s, t, c, bin_size, order)
821}
822
823#[pyfunction]
824#[pyo3(signature = (trains, bin_size=10, order=5, n_freqs=64))]
825fn py_spectral_granger_causality(
826 py: Python<'_>,
827 trains: Vec<PyReadonlyArray1<'_, i32>>,
828 bin_size: usize,
829 order: usize,
830 n_freqs: usize,
831) -> (Py<PyArray1<f64>>, usize, usize) {
832 let vecs: Vec<Vec<i32>> = trains
833 .iter()
834 .map(|t| t.as_slice().unwrap().to_vec())
835 .collect();
836 let refs: Vec<&[i32]> = vecs.iter().map(|v| v.as_slice()).collect();
837 let (gc, d) = analysis::causality::spectral_granger_causality(&refs, bin_size, order, n_freqs);
838 (gc.into_pyarray(py).into(), d, n_freqs)
839}
840
841#[pyfunction]
842#[pyo3(signature = (trains, bin_size=10, order=5, n_freqs=64))]
843fn py_partial_directed_coherence(
844 py: Python<'_>,
845 trains: Vec<PyReadonlyArray1<'_, i32>>,
846 bin_size: usize,
847 order: usize,
848 n_freqs: usize,
849) -> (Py<PyArray1<f64>>, usize, usize) {
850 let vecs: Vec<Vec<i32>> = trains
851 .iter()
852 .map(|t| t.as_slice().unwrap().to_vec())
853 .collect();
854 let refs: Vec<&[i32]> = vecs.iter().map(|v| v.as_slice()).collect();
855 let (pdc, d) = analysis::causality::partial_directed_coherence(&refs, bin_size, order, n_freqs);
856 (pdc.into_pyarray(py).into(), d, n_freqs)
857}
858
859#[pyfunction]
860#[pyo3(signature = (trains, bin_size=10, order=5, n_freqs=64))]
861fn py_directed_transfer_function(
862 py: Python<'_>,
863 trains: Vec<PyReadonlyArray1<'_, i32>>,
864 bin_size: usize,
865 order: usize,
866 n_freqs: usize,
867) -> (Py<PyArray1<f64>>, usize, usize) {
868 let vecs: Vec<Vec<i32>> = trains
869 .iter()
870 .map(|t| t.as_slice().unwrap().to_vec())
871 .collect();
872 let refs: Vec<&[i32]> = vecs.iter().map(|v| v.as_slice()).collect();
873 let (dtf, d) = analysis::causality::directed_transfer_function(&refs, bin_size, order, n_freqs);
874 (dtf.into_pyarray(py).into(), d, n_freqs)
875}
876
877#[pyfunction]
880#[pyo3(signature = (trains, preferred_directions, window=50))]
881fn py_population_vector_decode(
882 py: Python<'_>,
883 trains: Vec<PyReadonlyArray1<'_, i32>>,
884 preferred_directions: PyReadonlyArray1<'_, f64>,
885 window: usize,
886) -> Py<PyArray1<f64>> {
887 let vecs: Vec<Vec<i32>> = trains
888 .iter()
889 .map(|t| t.as_slice().unwrap().to_vec())
890 .collect();
891 let refs: Vec<&[i32]> = vecs.iter().map(|v| v.as_slice()).collect();
892 let dirs = preferred_directions.as_slice().unwrap();
893 analysis::decoding::population_vector_decode(&refs, dirs, window)
894 .into_pyarray(py)
895 .into()
896}
897
898#[pyfunction]
899#[pyo3(signature = (spike_counts, tuning_rates, n_stimuli, n_neurons, prior=None))]
900fn py_bayesian_decode(
901 spike_counts: PyReadonlyArray1<'_, f64>,
902 tuning_rates: PyReadonlyArray1<'_, f64>,
903 n_stimuli: usize,
904 n_neurons: usize,
905 prior: Option<PyReadonlyArray1<'_, f64>>,
906) -> usize {
907 let counts = spike_counts.as_slice().unwrap();
908 let rates = tuning_rates.as_slice().unwrap();
909 let p: Vec<f64> = prior
910 .map(|p| p.as_slice().unwrap().to_vec())
911 .unwrap_or_default();
912 analysis::decoding::bayesian_decode(counts, rates, n_stimuli, n_neurons, &p)
913}
914
915#[pyfunction]
916fn py_maximum_likelihood_decode(
917 spike_counts: PyReadonlyArray1<'_, f64>,
918 tuning_rates: PyReadonlyArray1<'_, f64>,
919 n_stimuli: usize,
920 n_neurons: usize,
921) -> usize {
922 let counts = spike_counts.as_slice().unwrap();
923 let rates = tuning_rates.as_slice().unwrap();
924 analysis::decoding::maximum_likelihood_decode(counts, rates, n_stimuli, n_neurons)
925}
926
927#[pyfunction]
928fn py_linear_discriminant_decode(
929 train_data: PyReadonlyArray1<'_, f64>,
930 n_samples: usize,
931 n_features: usize,
932 labels: PyReadonlyArray1<'_, i64>,
933 test_point: PyReadonlyArray1<'_, f64>,
934) -> i64 {
935 let data = train_data.as_slice().unwrap();
936 let lbl = labels.as_slice().unwrap();
937 let tp = test_point.as_slice().unwrap();
938 analysis::decoding::linear_discriminant_decode(data, n_samples, n_features, lbl, tp)
939}
940
941#[pyfunction]
942fn py_naive_bayes_decode(
943 train_data: PyReadonlyArray1<'_, f64>,
944 n_samples: usize,
945 n_features: usize,
946 labels: PyReadonlyArray1<'_, i64>,
947 test_point: PyReadonlyArray1<'_, f64>,
948) -> i64 {
949 let data = train_data.as_slice().unwrap();
950 let lbl = labels.as_slice().unwrap();
951 let tp = test_point.as_slice().unwrap();
952 analysis::decoding::naive_bayes_decode(data, n_samples, n_features, lbl, tp)
953}
954
955#[pyfunction]
958#[pyo3(signature = (trains, dt=1.0))]
959fn py_tokenise_spikes(
960 py: Python<'_>,
961 trains: Vec<PyReadonlyArray1<'_, i32>>,
962 dt: f64,
963) -> (Py<PyArray1<i64>>, Py<PyArray1<f64>>) {
964 let vecs: Vec<Vec<i32>> = trains
965 .iter()
966 .map(|t| t.as_slice().unwrap().to_vec())
967 .collect();
968 let refs: Vec<&[i32]> = vecs.iter().map(|v| v.as_slice()).collect();
969 let tokens = analysis::neural_decoders::tokenise_spikes(&refs, dt);
970 let uids: Vec<i64> = tokens.iter().map(|t| t.0 as i64).collect();
971 let times: Vec<f64> = tokens.iter().map(|t| t.1).collect();
972 (uids.into_pyarray(py).into(), times.into_pyarray(py).into())
973}
974
975#[pyfunction]
976fn py_sinusoidal_position_encode(
977 py: Python<'_>,
978 timestamps: PyReadonlyArray1<'_, f64>,
979 d_model: usize,
980) -> Py<PyArray1<f64>> {
981 let ts = timestamps.as_slice().unwrap();
982 analysis::neural_decoders::sinusoidal_position_encode(ts, d_model)
983 .into_pyarray(py)
984 .into()
985}
986
987#[pyfunction]
988fn py_scaled_dot_product_attention(
989 py: Python<'_>,
990 queries: PyReadonlyArray1<'_, f64>,
991 keys: PyReadonlyArray1<'_, f64>,
992 values: PyReadonlyArray1<'_, f64>,
993 nq: usize,
994 nk: usize,
995 d: usize,
996) -> Py<PyArray1<f64>> {
997 let q = queries.as_slice().unwrap();
998 let k = keys.as_slice().unwrap();
999 let v = values.as_slice().unwrap();
1000 analysis::neural_decoders::scaled_dot_product_attention(q, k, v, nq, nk, d)
1001 .into_pyarray(py)
1002 .into()
1003}
1004
1005#[pyfunction]
1006fn py_gaussian_attention(
1007 py: Python<'_>,
1008 queries: PyReadonlyArray1<'_, f64>,
1009 keys: PyReadonlyArray1<'_, f64>,
1010 values: PyReadonlyArray1<'_, f64>,
1011 nq: usize,
1012 nk: usize,
1013 d: usize,
1014 sigma: f64,
1015) -> Py<PyArray1<f64>> {
1016 let q = queries.as_slice().unwrap();
1017 let k = keys.as_slice().unwrap();
1018 let v = values.as_slice().unwrap();
1019 analysis::neural_decoders::gaussian_attention(q, k, v, nq, nk, d, sigma)
1020 .into_pyarray(py)
1021 .into()
1022}
1023
1024#[pyfunction]
1025fn py_infonce_loss(
1026 anchors: PyReadonlyArray1<'_, f64>,
1027 positives: PyReadonlyArray1<'_, f64>,
1028 n: usize,
1029 d: usize,
1030 temperature: f64,
1031) -> f64 {
1032 let a = anchors.as_slice().unwrap();
1033 let p = positives.as_slice().unwrap();
1034 analysis::neural_decoders::infonce_loss(a, p, n, d, temperature)
1035}
1036
1037#[pyfunction]
1040#[pyo3(signature = (trains, max_lag_ms=20.0, dt=0.001))]
1041fn py_functional_connectivity(
1042 py: Python<'_>,
1043 trains: Vec<PyReadonlyArray1<'_, i32>>,
1044 max_lag_ms: f64,
1045 dt: f64,
1046) -> Py<PyArray2<f64>> {
1047 let vecs: Vec<Vec<i32>> = trains
1048 .iter()
1049 .map(|t| t.as_slice().unwrap().to_vec())
1050 .collect();
1051 let refs: Vec<&[i32]> = vecs.iter().map(|v| v.as_slice()).collect();
1052 let mat = analysis::network::functional_connectivity(&refs, max_lag_ms, dt);
1053 let n = refs.len();
1054 let rows: Vec<Vec<f64>> = mat.chunks(n).map(|c| c.to_vec()).collect();
1055 numpy::PyArray2::from_vec2(py, &rows)
1056 .unwrap_or_else(|_| numpy::PyArray2::zeros(py, [n, n], false))
1057 .into()
1058}
1059
1060#[pyfunction]
1061#[pyo3(signature = (trains, bin_size=5, alpha=0.05))]
1062fn py_unitary_events(
1063 py: Python<'_>,
1064 trains: Vec<PyReadonlyArray1<'_, i32>>,
1065 bin_size: usize,
1066 alpha: f64,
1067) -> Py<PyArray1<i64>> {
1068 let vecs: Vec<Vec<i32>> = trains
1069 .iter()
1070 .map(|t| t.as_slice().unwrap().to_vec())
1071 .collect();
1072 let refs: Vec<&[i32]> = vecs.iter().map(|v| v.as_slice()).collect();
1073 let result = analysis::network::unitary_events(&refs, bin_size, alpha);
1074 let as_i64: Vec<i64> = result.into_iter().map(|v| v as i64).collect();
1075 as_i64.into_pyarray(py).into()
1076}
1077
1078#[pyfunction]
1079#[pyo3(signature = (trains, bin_size=5, threshold=2.0))]
1080fn py_cell_assembly_detection(
1081 trains: Vec<PyReadonlyArray1<'_, i32>>,
1082 bin_size: usize,
1083 threshold: f64,
1084) -> Vec<Vec<usize>> {
1085 let vecs: Vec<Vec<i32>> = trains
1086 .iter()
1087 .map(|t| t.as_slice().unwrap().to_vec())
1088 .collect();
1089 let refs: Vec<&[i32]> = vecs.iter().map(|v| v.as_slice()).collect();
1090 analysis::network::cell_assembly_detection(&refs, bin_size, threshold)
1091}
1092
1093#[pyfunction]
1094#[pyo3(signature = (trains, dt=0.001, max_delay_ms=20.0, min_chain_length=3))]
1095fn py_synfire_chain_detection(
1096 trains: Vec<PyReadonlyArray1<'_, i32>>,
1097 dt: f64,
1098 max_delay_ms: f64,
1099 min_chain_length: usize,
1100) -> Vec<Vec<usize>> {
1101 let vecs: Vec<Vec<i32>> = trains
1102 .iter()
1103 .map(|t| t.as_slice().unwrap().to_vec())
1104 .collect();
1105 let refs: Vec<&[i32]> = vecs.iter().map(|v| v.as_slice()).collect();
1106 analysis::network::synfire_chain_detection(&refs, dt, max_delay_ms, min_chain_length)
1107}
1108
1109#[pyfunction]
1112#[pyo3(signature = (binary_train, seed=0))]
1113fn py_surrogate_isi_shuffle(
1114 py: Python<'_>,
1115 binary_train: PyReadonlyArray1<'_, i32>,
1116 seed: u64,
1117) -> Py<PyArray1<i32>> {
1118 let data = binary_train.as_slice().unwrap();
1119 analysis::surrogates::surrogate_isi_shuffle(data, seed)
1120 .into_pyarray(py)
1121 .into()
1122}
1123
1124#[pyfunction]
1125#[pyo3(signature = (binary_train, dither_ms=5.0, dt=0.001, seed=0))]
1126fn py_surrogate_dither(
1127 py: Python<'_>,
1128 binary_train: PyReadonlyArray1<'_, i32>,
1129 dither_ms: f64,
1130 dt: f64,
1131 seed: u64,
1132) -> Py<PyArray1<i32>> {
1133 let data = binary_train.as_slice().unwrap();
1134 analysis::surrogates::surrogate_dither(data, dither_ms, dt, seed)
1135 .into_pyarray(py)
1136 .into()
1137}
1138
1139#[pyfunction]
1140#[pyo3(signature = (rate_hz, duration_s, dt=0.001, seed=0))]
1141fn py_homogeneous_poisson(
1142 py: Python<'_>,
1143 rate_hz: f64,
1144 duration_s: f64,
1145 dt: f64,
1146 seed: u64,
1147) -> Py<PyArray1<f64>> {
1148 analysis::surrogates::homogeneous_poisson(rate_hz, duration_s, dt, seed)
1149 .into_pyarray(py)
1150 .into()
1151}
1152
1153#[pyfunction]
1154#[pyo3(signature = (rate_hz, shape, duration_s, dt=0.001, seed=0))]
1155fn py_gamma_process(
1156 py: Python<'_>,
1157 rate_hz: f64,
1158 shape: f64,
1159 duration_s: f64,
1160 dt: f64,
1161 seed: u64,
1162) -> Py<PyArray1<f64>> {
1163 analysis::surrogates::gamma_process(rate_hz, shape, duration_s, dt, seed)
1164 .into_pyarray(py)
1165 .into()
1166}
1167
1168#[pyfunction]
1169#[pyo3(signature = (rate_hz, burst_mean, duration_s, dt=0.001, seed=0))]
1170fn py_compound_poisson_process(
1171 py: Python<'_>,
1172 rate_hz: f64,
1173 burst_mean: f64,
1174 duration_s: f64,
1175 dt: f64,
1176 seed: u64,
1177) -> Py<PyArray1<f64>> {
1178 analysis::surrogates::compound_poisson_process(rate_hz, burst_mean, duration_s, dt, seed)
1179 .into_pyarray(py)
1180 .into()
1181}
1182
1183#[pyfunction]
1184#[pyo3(signature = (binary_train, seed=0))]
1185fn py_surrogate_joint_isi(
1186 py: Python<'_>,
1187 binary_train: PyReadonlyArray1<'_, i32>,
1188 seed: u64,
1189) -> Py<PyArray1<i32>> {
1190 let data = binary_train.as_slice().unwrap();
1191 analysis::surrogates::surrogate_joint_isi(data, seed)
1192 .into_pyarray(py)
1193 .into()
1194}
1195
1196#[pyfunction]
1197#[pyo3(signature = (binary_train, bin_size=10, seed=0))]
1198fn py_surrogate_bin_shuffling(
1199 py: Python<'_>,
1200 binary_train: PyReadonlyArray1<'_, i32>,
1201 bin_size: usize,
1202 seed: u64,
1203) -> Py<PyArray1<i32>> {
1204 let data = binary_train.as_slice().unwrap();
1205 analysis::surrogates::surrogate_bin_shuffling(data, bin_size, seed)
1206 .into_pyarray(py)
1207 .into()
1208}
1209
1210#[pyfunction]
1211#[pyo3(signature = (binary_train, max_shift=50, seed=0))]
1212fn py_surrogate_spike_train_shifting(
1213 py: Python<'_>,
1214 binary_train: PyReadonlyArray1<'_, i32>,
1215 max_shift: usize,
1216 seed: u64,
1217) -> Py<PyArray1<i32>> {
1218 let data = binary_train.as_slice().unwrap();
1219 analysis::surrogates::surrogate_spike_train_shifting(data, max_shift, seed)
1220 .into_pyarray(py)
1221 .into()
1222}
1223
1224#[pyfunction]
1227#[pyo3(signature = (binary_train, dt=0.001, max_isi_ms=10.0, min_spikes=3))]
1228fn py_burst_detection(
1229 binary_train: PyReadonlyArray1<'_, i32>,
1230 dt: f64,
1231 max_isi_ms: f64,
1232 min_spikes: usize,
1233) -> Vec<(f64, f64, usize)> {
1234 let data = binary_train.as_slice().unwrap();
1235 analysis::temporal::burst_detection(data, dt, max_isi_ms, min_spikes)
1236}
1237
1238#[pyfunction]
1239#[pyo3(signature = (binary_train, dt=0.001))]
1240fn py_first_spike_latency(binary_train: PyReadonlyArray1<'_, i32>, dt: f64) -> f64 {
1241 let data = binary_train.as_slice().unwrap();
1242 analysis::temporal::first_spike_latency(data, dt)
1243}
1244
1245#[pyfunction]
1246#[pyo3(signature = (binary_train, baseline_steps=100, dt=0.001, threshold_sigma=3.0))]
1247fn py_response_onset(
1248 binary_train: PyReadonlyArray1<'_, i32>,
1249 baseline_steps: usize,
1250 dt: f64,
1251 threshold_sigma: f64,
1252) -> f64 {
1253 let data = binary_train.as_slice().unwrap();
1254 analysis::temporal::response_onset(data, baseline_steps, dt, threshold_sigma)
1255}
1256
1257#[pyfunction]
1258#[pyo3(signature = (binary_train, bin_size=50, threshold=3.0))]
1259fn py_change_point_detection(
1260 py: Python<'_>,
1261 binary_train: PyReadonlyArray1<'_, i32>,
1262 bin_size: usize,
1263 threshold: f64,
1264) -> Py<PyArray1<i64>> {
1265 let data = binary_train.as_slice().unwrap();
1266 let cps = analysis::temporal::change_point_detection(data, bin_size, threshold);
1267 let as_i64: Vec<i64> = cps.into_iter().map(|v| v as i64).collect();
1268 as_i64.into_pyarray(py).into()
1269}
1270
1271#[pyfunction]
1274#[pyo3(signature = (times_a, times_b, t_start=0.0, t_end=1.0))]
1275fn py_spike_directionality(
1276 times_a: PyReadonlyArray1<'_, f64>,
1277 times_b: PyReadonlyArray1<'_, f64>,
1278 t_start: f64,
1279 t_end: f64,
1280) -> f64 {
1281 let a = times_a.as_slice().unwrap();
1282 let b = times_b.as_slice().unwrap();
1283 analysis::patterns::spike_directionality(a, b, t_start, t_end)
1284}
1285
1286#[pyfunction]
1287#[pyo3(signature = (times_list, t_start=0.0, t_end=1.0))]
1288fn py_spike_train_order(
1289 py: Python<'_>,
1290 times_list: Vec<PyReadonlyArray1<'_, f64>>,
1291 t_start: f64,
1292 t_end: f64,
1293) -> Py<PyArray2<f64>> {
1294 let vecs: Vec<Vec<f64>> = times_list
1295 .iter()
1296 .map(|t| t.as_slice().unwrap().to_vec())
1297 .collect();
1298 let refs: Vec<&[f64]> = vecs.iter().map(|v| v.as_slice()).collect();
1299 let mat = analysis::patterns::spike_train_order(&refs, t_start, t_end);
1300 let n = refs.len();
1301 let rows: Vec<Vec<f64>> = mat.chunks(n).map(|c| c.to_vec()).collect();
1302 numpy::PyArray2::from_vec2(py, &rows)
1303 .unwrap_or_else(|_| numpy::PyArray2::zeros(py, [n, n], false))
1304 .into()
1305}
1306
1307#[pyfunction]
1308#[pyo3(signature = (binary_train, dt=0.001, max_lag=20))]
1309fn py_cubic_higher_order(
1310 py: Python<'_>,
1311 binary_train: PyReadonlyArray1<'_, i32>,
1312 dt: f64,
1313 max_lag: usize,
1314) -> Py<PyArray2<f64>> {
1315 let data = binary_train.as_slice().unwrap();
1316 let c3 = analysis::patterns::cubic_higher_order(data, dt, max_lag);
1317 let rows: Vec<Vec<f64>> = c3.chunks(max_lag).map(|c| c.to_vec()).collect();
1318 numpy::PyArray2::from_vec2(py, &rows)
1319 .unwrap_or_else(|_| numpy::PyArray2::zeros(py, [max_lag, max_lag], false))
1320 .into()
1321}
1322
1323#[pyfunction]
1326#[pyo3(signature = (binary_train, dt=0.001))]
1327fn py_power_spectrum(
1328 py: Python<'_>,
1329 binary_train: PyReadonlyArray1<'_, i32>,
1330 dt: f64,
1331) -> (Py<PyArray1<f64>>, Py<PyArray1<f64>>) {
1332 let data = binary_train.as_slice().unwrap();
1333 let (psd, freqs) = analysis::spectral::power_spectrum(data, dt);
1334 (psd.into_pyarray(py).into(), freqs.into_pyarray(py).into())
1335}
1336
1337#[pyfunction]
1340#[pyo3(signature = (waveform, dt=3.3333333333333335e-05))]
1341fn py_waveform_width(waveform: PyReadonlyArray1<'_, f64>, dt: f64) -> f64 {
1342 analysis::waveform::waveform_width(waveform.as_slice().unwrap(), dt)
1343}
1344
1345#[pyfunction]
1346fn py_waveform_amplitude(waveform: PyReadonlyArray1<'_, f64>) -> f64 {
1347 analysis::waveform::waveform_amplitude(waveform.as_slice().unwrap())
1348}
1349
1350#[pyfunction]
1351#[pyo3(signature = (waveform, dt=3.3333333333333335e-05))]
1352fn py_waveform_repolarization_slope(waveform: PyReadonlyArray1<'_, f64>, dt: f64) -> f64 {
1353 analysis::waveform::waveform_repolarization_slope(waveform.as_slice().unwrap(), dt)
1354}
1355
1356#[pyfunction]
1357#[pyo3(signature = (waveform, dt=3.3333333333333335e-05))]
1358fn py_waveform_recovery_slope(waveform: PyReadonlyArray1<'_, f64>, dt: f64) -> f64 {
1359 analysis::waveform::waveform_recovery_slope(waveform.as_slice().unwrap(), dt)
1360}
1361
1362#[pyfunction]
1363#[pyo3(signature = (waveform, dt=3.3333333333333335e-05))]
1364fn py_waveform_halfwidth(waveform: PyReadonlyArray1<'_, f64>, dt: f64) -> f64 {
1365 analysis::waveform::waveform_halfwidth(waveform.as_slice().unwrap(), dt)
1366}
1367
1368#[pyfunction]
1369fn py_waveform_pt_ratio(waveform: PyReadonlyArray1<'_, f64>) -> f64 {
1370 analysis::waveform::waveform_pt_ratio(waveform.as_slice().unwrap())
1371}
1372
1373#[pyfunction]
1376#[pyo3(signature = (binary_train, dt=0.001, window_ms=50.0))]
1377fn py_conditional_intensity(
1378 py: Python<'_>,
1379 binary_train: PyReadonlyArray1<'_, i32>,
1380 dt: f64,
1381 window_ms: f64,
1382) -> Py<PyArray1<f64>> {
1383 let data = binary_train.as_slice().unwrap();
1384 analysis::point_process::conditional_intensity(data, dt, window_ms)
1385 .into_pyarray(py)
1386 .into()
1387}
1388
1389#[pyfunction]
1390#[pyo3(signature = (binary_train, dt=0.001, bins=30))]
1391fn py_isi_hazard_function(
1392 py: Python<'_>,
1393 binary_train: PyReadonlyArray1<'_, i32>,
1394 dt: f64,
1395 bins: usize,
1396) -> (Py<PyArray1<f64>>, Py<PyArray1<f64>>) {
1397 let data = binary_train.as_slice().unwrap();
1398 let (hazard, centres) = analysis::point_process::isi_hazard_function(data, dt, bins);
1399 (
1400 hazard.into_pyarray(py).into(),
1401 centres.into_pyarray(py).into(),
1402 )
1403}
1404
1405#[pyfunction]
1406#[pyo3(signature = (binary_train, dt=0.001, bins=30))]
1407fn py_isi_survivor_function(
1408 py: Python<'_>,
1409 binary_train: PyReadonlyArray1<'_, i32>,
1410 dt: f64,
1411 bins: usize,
1412) -> (Py<PyArray1<f64>>, Py<PyArray1<f64>>) {
1413 let data = binary_train.as_slice().unwrap();
1414 let (surv, centres) = analysis::point_process::isi_survivor_function(data, dt, bins);
1415 (
1416 surv.into_pyarray(py).into(),
1417 centres.into_pyarray(py).into(),
1418 )
1419}
1420
1421#[pyfunction]
1422#[pyo3(signature = (binary_train, dt=0.001, bins=30))]
1423fn py_renewal_density(
1424 py: Python<'_>,
1425 binary_train: PyReadonlyArray1<'_, i32>,
1426 dt: f64,
1427 bins: usize,
1428) -> (Py<PyArray1<f64>>, Py<PyArray1<f64>>) {
1429 let data = binary_train.as_slice().unwrap();
1430 let (dens, centres) = analysis::point_process::renewal_density(data, dt, bins);
1431 (
1432 dens.into_pyarray(py).into(),
1433 centres.into_pyarray(py).into(),
1434 )
1435}
1436
1437#[pyfunction]
1440#[pyo3(signature = (stimulus, binary_train, window_steps=50))]
1441fn py_spike_triggered_average(
1442 py: Python<'_>,
1443 stimulus: PyReadonlyArray1<'_, f64>,
1444 binary_train: PyReadonlyArray1<'_, i32>,
1445 window_steps: usize,
1446) -> Py<PyArray1<f64>> {
1447 analysis::stimulus::spike_triggered_average(
1448 stimulus.as_slice().unwrap(),
1449 binary_train.as_slice().unwrap(),
1450 window_steps,
1451 )
1452 .into_pyarray(py)
1453 .into()
1454}
1455
1456#[pyfunction]
1457#[pyo3(signature = (stimulus, binary_train, window_steps=50))]
1458fn py_spike_triggered_covariance(
1459 py: Python<'_>,
1460 stimulus: PyReadonlyArray1<'_, f64>,
1461 binary_train: PyReadonlyArray1<'_, i32>,
1462 window_steps: usize,
1463) -> Py<PyArray2<f64>> {
1464 let cov = analysis::stimulus::spike_triggered_covariance(
1465 stimulus.as_slice().unwrap(),
1466 binary_train.as_slice().unwrap(),
1467 window_steps,
1468 );
1469 let rows: Vec<Vec<f64>> = cov.chunks(window_steps).map(|c| c.to_vec()).collect();
1470 numpy::PyArray2::from_vec2(py, &rows)
1471 .unwrap_or_else(|_| numpy::PyArray2::zeros(py, [window_steps, window_steps], false))
1472 .into()
1473}
1474
1475#[pyfunction]
1476#[pyo3(signature = (binary_train, positions, n_bins=20, dt=0.001))]
1477fn py_spatial_information(
1478 binary_train: PyReadonlyArray1<'_, i32>,
1479 positions: PyReadonlyArray1<'_, f64>,
1480 n_bins: usize,
1481 dt: f64,
1482) -> f64 {
1483 analysis::stimulus::spatial_information(
1484 binary_train.as_slice().unwrap(),
1485 positions.as_slice().unwrap(),
1486 n_bins,
1487 dt,
1488 )
1489}
1490
1491#[pyfunction]
1492#[pyo3(signature = (binary_train, positions, n_bins=50, threshold_std=2.0, dt=0.001))]
1493fn py_place_field_detection(
1494 binary_train: PyReadonlyArray1<'_, i32>,
1495 positions: PyReadonlyArray1<'_, f64>,
1496 n_bins: usize,
1497 threshold_std: f64,
1498 dt: f64,
1499) -> Vec<(f64, f64)> {
1500 analysis::stimulus::place_field_detection(
1501 binary_train.as_slice().unwrap(),
1502 positions.as_slice().unwrap(),
1503 n_bins,
1504 threshold_std,
1505 dt,
1506 )
1507}
1508
1509#[pyfunction]
1510#[pyo3(signature = (binary_train, stimulus_values, n_bins=20, dt=0.001))]
1511fn py_tuning_curve(
1512 py: Python<'_>,
1513 binary_train: PyReadonlyArray1<'_, i32>,
1514 stimulus_values: PyReadonlyArray1<'_, f64>,
1515 n_bins: usize,
1516 dt: f64,
1517) -> (Py<PyArray1<f64>>, Py<PyArray1<f64>>) {
1518 let (rates, centres) = analysis::stimulus::tuning_curve(
1519 binary_train.as_slice().unwrap(),
1520 stimulus_values.as_slice().unwrap(),
1521 n_bins,
1522 dt,
1523 );
1524 (
1525 rates.into_pyarray(py).into(),
1526 centres.into_pyarray(py).into(),
1527 )
1528}
1529
1530#[pyfunction]
1533fn py_phase_locking_value(
1534 binary_train: PyReadonlyArray1<'_, i32>,
1535 lfp_signal: PyReadonlyArray1<'_, f64>,
1536) -> f64 {
1537 analysis::lfp::phase_locking_value(
1538 binary_train.as_slice().unwrap(),
1539 lfp_signal.as_slice().unwrap(),
1540 )
1541}
1542
1543#[pyfunction]
1544#[pyo3(signature = (binary_train, lfp_signal, dt=0.001))]
1545fn py_spike_field_coherence(
1546 py: Python<'_>,
1547 binary_train: PyReadonlyArray1<'_, i32>,
1548 lfp_signal: PyReadonlyArray1<'_, f64>,
1549 dt: f64,
1550) -> (Py<PyArray1<f64>>, Py<PyArray1<f64>>) {
1551 let (sfc, freqs) = analysis::lfp::spike_field_coherence(
1552 binary_train.as_slice().unwrap(),
1553 lfp_signal.as_slice().unwrap(),
1554 dt,
1555 );
1556 (sfc.into_pyarray(py).into(), freqs.into_pyarray(py).into())
1557}
1558
1559#[pyfunction]
1560#[pyo3(signature = (binary_train, lfp_signal, n_bins=36))]
1561fn py_spike_phase_histogram(
1562 py: Python<'_>,
1563 binary_train: PyReadonlyArray1<'_, i32>,
1564 lfp_signal: PyReadonlyArray1<'_, f64>,
1565 n_bins: usize,
1566) -> (Py<PyArray1<i64>>, Py<PyArray1<f64>>) {
1567 let (hist, centres) = analysis::lfp::spike_phase_histogram(
1568 binary_train.as_slice().unwrap(),
1569 lfp_signal.as_slice().unwrap(),
1570 n_bins,
1571 );
1572 (
1573 hist.into_pyarray(py).into(),
1574 centres.into_pyarray(py).into(),
1575 )
1576}
1577
1578#[pyfunction]
1581fn py_isolation_distance(
1582 cluster: PyReadonlyArray2<'_, f64>,
1583 noise: PyReadonlyArray2<'_, f64>,
1584) -> f64 {
1585 let c_shape = cluster.shape();
1586 let n_shape = noise.shape();
1587 let d = c_shape[1];
1588 let c_data: Vec<f64> = cluster.as_slice().unwrap().to_vec();
1589 let n_data: Vec<f64> = noise.as_slice().unwrap().to_vec();
1590 analysis::sorting_quality::isolation_distance(&c_data, c_shape[0], &n_data, n_shape[0], d)
1591}
1592
1593#[pyfunction]
1594fn py_l_ratio(cluster: PyReadonlyArray2<'_, f64>, noise: PyReadonlyArray2<'_, f64>) -> f64 {
1595 let c_shape = cluster.shape();
1596 let n_shape = noise.shape();
1597 let d = c_shape[1];
1598 let c_data: Vec<f64> = cluster.as_slice().unwrap().to_vec();
1599 let n_data: Vec<f64> = noise.as_slice().unwrap().to_vec();
1600 analysis::sorting_quality::l_ratio(&c_data, c_shape[0], &n_data, n_shape[0], d)
1601}
1602
1603#[pyfunction]
1604fn py_silhouette_score(
1605 features: PyReadonlyArray2<'_, f64>,
1606 labels: PyReadonlyArray1<'_, i64>,
1607) -> f64 {
1608 let shape = features.shape();
1609 let f_data: Vec<f64> = features.as_slice().unwrap().to_vec();
1610 let l_data: Vec<i64> = labels.as_slice().unwrap().to_vec();
1611 analysis::sorting_quality::silhouette_score(&f_data, shape[0], shape[1], &l_data)
1612}
1613
1614#[pyfunction]
1615fn py_d_prime(cluster_a: PyReadonlyArray2<'_, f64>, cluster_b: PyReadonlyArray2<'_, f64>) -> f64 {
1616 let a_shape = cluster_a.shape();
1617 let b_shape = cluster_b.shape();
1618 let d = a_shape[1];
1619 let a_data: Vec<f64> = cluster_a.as_slice().unwrap().to_vec();
1620 let b_data: Vec<f64> = cluster_b.as_slice().unwrap().to_vec();
1621 analysis::sorting_quality::d_prime(&a_data, a_shape[0], &b_data, b_shape[0], d)
1622}
1623
1624#[pyfunction]
1625#[pyo3(signature = (binary_train, dt=0.001, refractory_ms=1.5))]
1626fn py_isi_violation_rate(
1627 binary_train: PyReadonlyArray1<'_, i32>,
1628 dt: f64,
1629 refractory_ms: f64,
1630) -> f64 {
1631 analysis::sorting_quality::isi_violation_rate(
1632 binary_train.as_slice().unwrap(),
1633 dt,
1634 refractory_ms,
1635 )
1636}
1637
1638#[pyfunction]
1639#[pyo3(signature = (binary_train, n_bins=100))]
1640fn py_presence_ratio(binary_train: PyReadonlyArray1<'_, i32>, n_bins: usize) -> f64 {
1641 analysis::sorting_quality::presence_ratio(binary_train.as_slice().unwrap(), n_bins)
1642}
1643
1644#[pyfunction]
1645#[pyo3(signature = (amplitudes, bins=100))]
1646fn py_amplitude_cutoff(amplitudes: PyReadonlyArray1<'_, f64>, bins: usize) -> f64 {
1647 analysis::sorting_quality::amplitude_cutoff(amplitudes.as_slice().unwrap(), bins)
1648}
1649
1650#[pyfunction]
1651fn py_snr(waveforms: PyReadonlyArray2<'_, f64>) -> f64 {
1652 let shape = waveforms.shape();
1653 let data: Vec<f64> = waveforms.as_slice().unwrap().to_vec();
1654 analysis::sorting_quality::snr(&data, shape[0], shape[1])
1655}
1656
1657#[pyfunction]
1658#[pyo3(signature = (cluster, noise, k=4))]
1659fn py_nn_hit_rate(
1660 cluster: PyReadonlyArray2<'_, f64>,
1661 noise: PyReadonlyArray2<'_, f64>,
1662 k: usize,
1663) -> f64 {
1664 let c_shape = cluster.shape();
1665 let n_shape = noise.shape();
1666 let d = c_shape[1];
1667 let c_data: Vec<f64> = cluster.as_slice().unwrap().to_vec();
1668 let n_data: Vec<f64> = noise.as_slice().unwrap().to_vec();
1669 analysis::sorting_quality::nn_hit_rate(&c_data, c_shape[0], &n_data, n_shape[0], d, k)
1670}
1671
1672#[pyfunction]
1673#[pyo3(signature = (waveforms, timestamps, n_bins=10))]
1674fn py_drift_metric(
1675 waveforms: PyReadonlyArray2<'_, f64>,
1676 timestamps: PyReadonlyArray1<'_, f64>,
1677 n_bins: usize,
1678) -> f64 {
1679 let shape = waveforms.shape();
1680 let data: Vec<f64> = waveforms.as_slice().unwrap().to_vec();
1681 let ts: Vec<f64> = timestamps.as_slice().unwrap().to_vec();
1682 analysis::sorting_quality::drift_metric(&data, shape[0], shape[1], &ts, n_bins)
1683}
1684
1685#[pyfunction]
1688#[pyo3(signature = (trains, n_components=3, bin_size=10))]
1689fn py_spike_train_pca(
1690 py: Python<'_>,
1691 trains: Vec<PyReadonlyArray1<'_, i32>>,
1692 n_components: usize,
1693 bin_size: usize,
1694) -> (Py<PyArray1<f64>>, Py<PyArray1<f64>>) {
1695 let vecs: Vec<Vec<i32>> = trains
1696 .iter()
1697 .map(|t| t.as_slice().unwrap().to_vec())
1698 .collect();
1699 let refs: Vec<&[i32]> = vecs.iter().map(|v| v.as_slice()).collect();
1700 let (proj, expl) = analysis::dimensionality::spike_train_pca(&refs, n_components, bin_size);
1701 (proj.into_pyarray(py).into(), expl.into_pyarray(py).into())
1702}
1703
1704#[pyfunction]
1705#[pyo3(signature = (conditions, n_components=3, bin_size=10))]
1706fn py_demixed_pca(
1707 py: Python<'_>,
1708 conditions: Vec<Vec<PyReadonlyArray1<'_, i32>>>,
1709 n_components: usize,
1710 bin_size: usize,
1711) -> (Py<PyArray1<f64>>, Py<PyArray1<f64>>) {
1712 let vecs: Vec<Vec<Vec<i32>>> = conditions
1713 .iter()
1714 .map(|cond| {
1715 cond.iter()
1716 .map(|t| t.as_slice().unwrap().to_vec())
1717 .collect()
1718 })
1719 .collect();
1720 let refs: Vec<Vec<&[i32]>> = vecs
1721 .iter()
1722 .map(|cond| cond.iter().map(|v| v.as_slice()).collect())
1723 .collect();
1724 let (proj, expl) = analysis::dimensionality::demixed_pca(&refs, n_components, bin_size);
1725 (proj.into_pyarray(py).into(), expl.into_pyarray(py).into())
1726}
1727
1728#[pyfunction]
1729#[pyo3(signature = (trains, n_factors=3, bin_size=10, n_iter=50))]
1730fn py_factor_analysis(
1731 py: Python<'_>,
1732 trains: Vec<PyReadonlyArray1<'_, i32>>,
1733 n_factors: usize,
1734 bin_size: usize,
1735 n_iter: usize,
1736) -> (Py<PyArray1<f64>>, Py<PyArray1<f64>>) {
1737 let vecs: Vec<Vec<i32>> = trains
1738 .iter()
1739 .map(|t| t.as_slice().unwrap().to_vec())
1740 .collect();
1741 let refs: Vec<&[i32]> = vecs.iter().map(|v| v.as_slice()).collect();
1742 let (loadings, psi) =
1743 analysis::dimensionality::factor_analysis(&refs, n_factors, bin_size, n_iter);
1744 (
1745 loadings.into_pyarray(py).into(),
1746 psi.into_pyarray(py).into(),
1747 )
1748}
1749
1750#[pyfunction]
1755#[pyo3(signature = (mat, n_components=3))]
1756fn py_pca_components(
1757 py: Python<'_>,
1758 mat: PyReadonlyArray2<'_, f64>,
1759 n_components: usize,
1760) -> (Py<PyArray1<f64>>, Py<PyArray1<f64>>) {
1761 let shape = mat.shape();
1762 let data: Vec<f64> = mat.as_slice().unwrap().to_vec();
1763 let (proj, expl) =
1764 analysis::dimensionality::pca_from_centered(&data, shape[0], shape[1], n_components);
1765 (proj.into_pyarray(py).into(), expl.into_pyarray(py).into())
1766}
1767
1768#[pyfunction]
1769#[pyo3(signature = (mean_mat, n_components=3))]
1770fn py_demixed_components(
1771 py: Python<'_>,
1772 mean_mat: PyReadonlyArray2<'_, f64>,
1773 n_components: usize,
1774) -> (Py<PyArray1<f64>>, Py<PyArray1<f64>>) {
1775 let shape = mean_mat.shape();
1776 let data: Vec<f64> = mean_mat.as_slice().unwrap().to_vec();
1777 let (proj, expl) =
1778 analysis::dimensionality::demixed_from_centered(&data, shape[0], shape[1], n_components);
1779 (proj.into_pyarray(py).into(), expl.into_pyarray(py).into())
1780}
1781
1782#[pyfunction]
1783#[pyo3(signature = (mat, n_factors=3, n_iter=50))]
1784fn py_factor_loadings(
1785 py: Python<'_>,
1786 mat: PyReadonlyArray2<'_, f64>,
1787 n_factors: usize,
1788 n_iter: usize,
1789) -> (Py<PyArray1<f64>>, Py<PyArray1<f64>>) {
1790 let shape = mat.shape();
1791 let data: Vec<f64> = mat.as_slice().unwrap().to_vec();
1792 let (loadings, psi) =
1793 analysis::dimensionality::fa_from_centered(&data, shape[0], shape[1], n_factors, n_iter);
1794 (
1795 loadings.into_pyarray(py).into(),
1796 psi.into_pyarray(py).into(),
1797 )
1798}
1799
1800#[pyfunction]
1803#[pyo3(signature = (trains, n_latents=3, bin_ms=20.0, dt=0.001, max_iter=50, tol=1e-4, seed=42))]
1804fn py_gpfa<'py>(
1805 py: Python<'py>,
1806 trains: Vec<PyReadonlyArray1<'py, i32>>,
1807 n_latents: usize,
1808 bin_ms: f64,
1809 dt: f64,
1810 max_iter: usize,
1811 tol: f64,
1812 seed: u64,
1813) -> PyResult<Bound<'py, PyDict>> {
1814 let vecs: Vec<Vec<i32>> = trains
1815 .iter()
1816 .map(|t| t.as_slice().unwrap().to_vec())
1817 .collect();
1818 let refs: Vec<&[i32]> = vecs.iter().map(|v| v.as_slice()).collect();
1819 let result = analysis::gpfa::gpfa(&refs, n_latents, bin_ms, dt, max_iter, tol, seed);
1820
1821 let dict = PyDict::new(py);
1822 dict.set_item("trajectories", result.trajectories.into_pyarray(py))?;
1823 dict.set_item("C", result.c.into_pyarray(py))?;
1824 dict.set_item("d", result.d.into_pyarray(py))?;
1825 dict.set_item("R", result.r.into_pyarray(py))?;
1826 dict.set_item("tau", result.tau.into_pyarray(py))?;
1827 dict.set_item("log_likelihoods", result.log_likelihoods.into_pyarray(py))?;
1828 dict.set_item("n_latents", result.n_latents)?;
1829 dict.set_item("n_bins", result.n_bins)?;
1830 dict.set_item("n_neurons", result.n_neurons)?;
1831 Ok(dict)
1832}
1833
1834#[pyfunction]
1840#[pyo3(signature = (y, n_neurons, n_bins, c0, d0, r0_diag, tau, n_latents, max_iter, tol))]
1841#[allow(clippy::too_many_arguments, clippy::type_complexity)]
1842fn py_gpfa_em<'py>(
1843 py: Python<'py>,
1844 y: PyReadonlyArray1<'py, f64>,
1845 n_neurons: usize,
1846 n_bins: usize,
1847 c0: PyReadonlyArray1<'py, f64>,
1848 d0: PyReadonlyArray1<'py, f64>,
1849 r0_diag: PyReadonlyArray1<'py, f64>,
1850 tau: PyReadonlyArray1<'py, f64>,
1851 n_latents: usize,
1852 max_iter: usize,
1853 tol: f64,
1854) -> PyResult<(
1855 Bound<'py, PyArray1<f64>>,
1856 Bound<'py, PyArray1<f64>>,
1857 Bound<'py, PyArray1<f64>>,
1858 Bound<'py, PyArray1<f64>>,
1859 Vec<f64>,
1860)> {
1861 let (x_post, c, d, r, log_liks) = analysis::gpfa::gpfa_em_from_init(
1862 y.as_slice()?,
1863 c0.as_slice()?,
1864 d0.as_slice()?,
1865 r0_diag.as_slice()?,
1866 tau.as_slice()?,
1867 n_neurons,
1868 n_bins,
1869 n_latents,
1870 max_iter,
1871 tol,
1872 );
1873 Ok((
1874 x_post.into_pyarray(py),
1875 c.into_pyarray(py),
1876 d.into_pyarray(py),
1877 r.into_pyarray(py),
1878 log_liks,
1879 ))
1880}
1881
1882#[pyfunction]
1883#[pyo3(signature = (new_trains, c, d, r, tau, n_latents, bin_ms=20.0, dt=0.001))]
1884fn py_gpfa_transform(
1885 py: Python<'_>,
1886 new_trains: Vec<PyReadonlyArray1<'_, i32>>,
1887 c: PyReadonlyArray1<'_, f64>,
1888 d: PyReadonlyArray1<'_, f64>,
1889 r: PyReadonlyArray1<'_, f64>,
1890 tau: PyReadonlyArray1<'_, f64>,
1891 n_latents: usize,
1892 bin_ms: f64,
1893 dt: f64,
1894) -> Py<PyArray1<f64>> {
1895 let vecs: Vec<Vec<i32>> = new_trains
1896 .iter()
1897 .map(|t| t.as_slice().unwrap().to_vec())
1898 .collect();
1899 let refs: Vec<&[i32]> = vecs.iter().map(|v| v.as_slice()).collect();
1900 let proj = analysis::gpfa::gpfa_transform(
1901 &refs,
1902 c.as_slice().unwrap(),
1903 d.as_slice().unwrap(),
1904 r.as_slice().unwrap(),
1905 tau.as_slice().unwrap(),
1906 n_latents,
1907 bin_ms,
1908 dt,
1909 );
1910 proj.into_pyarray(py).into()
1911}
1912
1913#[pyfunction]
1916#[pyo3(signature = (trains, bin_ms=5.0, dt=0.001, min_support=3, max_pattern_size=5, n_surrogates=100, alpha=0.05, seed=42))]
1917fn py_spade_detect<'py>(
1918 py: Python<'py>,
1919 trains: Vec<PyReadonlyArray1<'py, i32>>,
1920 bin_ms: f64,
1921 dt: f64,
1922 min_support: usize,
1923 max_pattern_size: usize,
1924 n_surrogates: usize,
1925 alpha: f64,
1926 seed: u64,
1927) -> PyResult<Vec<Bound<'py, PyDict>>> {
1928 let vecs: Vec<Vec<i32>> = trains
1929 .iter()
1930 .map(|t| t.as_slice().unwrap().to_vec())
1931 .collect();
1932 let refs: Vec<&[i32]> = vecs.iter().map(|v| v.as_slice()).collect();
1933 let results = analysis::spade::spade_detect(
1934 &refs,
1935 bin_ms,
1936 dt,
1937 min_support,
1938 max_pattern_size,
1939 n_surrogates,
1940 alpha,
1941 seed,
1942 );
1943 let mut dicts = Vec::new();
1944 for pat in results {
1945 let dict = PyDict::new(py);
1946 dict.set_item(
1947 "neurons",
1948 pat.neurons.iter().map(|&n| n as i64).collect::<Vec<_>>(),
1949 )?;
1950 dict.set_item("lags", pat.lags.clone())?;
1951 dict.set_item("count", pat.count as i64)?;
1952 dict.set_item("p_value", pat.p_value)?;
1953 dicts.push(dict);
1954 }
1955 Ok(dicts)
1956}