pylottone 0.2.2__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- pylottone/__init__.py +58 -0
- pylottone/constants.py +8 -0
- pylottone/editer.py +252 -0
- pylottone/model_selection.py +94 -0
- pylottone/mrdhelper.py +476 -0
- pylottone/pt.py +735 -0
- pylottone/reconstruction/Body6Spine18.xml +89 -0
- pylottone/reconstruction/GIRF.py +275 -0
- pylottone/reconstruction/GIRF_20200221_Duyn_method_coil2.mat +0 -0
- pylottone/reconstruction/client.py +380 -0
- pylottone/reconstruction/coils.py +499 -0
- pylottone/reconstruction/connection.py +434 -0
- pylottone/reconstruction/constants.py +55 -0
- pylottone/reconstruction/send_to_recon_server.py +112 -0
- pylottone/resources/__init__.py +1 -0
- pylottone/resources/rocket_pipeline.pkl +0 -0
- pylottone/selectionui.py +72 -0
- pylottone/selfnav.py +244 -0
- pylottone/signal.py +454 -0
- pylottone/sobi/__init__.py +1 -0
- pylottone/sobi/sobi.py +216 -0
- pylottone/sobi/utils.py +78 -0
- pylottone/trajectory.py +105 -0
- pylottone/triggering.py +1327 -0
- pylottone/vis.py +641 -0
- pylottone-0.2.2.dist-info/METADATA +142 -0
- pylottone-0.2.2.dist-info/RECORD +30 -0
- pylottone-0.2.2.dist-info/WHEEL +4 -0
- pylottone-0.2.2.dist-info/entry_points.txt +2 -0
- pylottone-0.2.2.dist-info/licenses/LICENSE +21 -0
pylottone/__init__.py
ADDED
|
@@ -0,0 +1,58 @@
|
|
|
1
|
+
from importlib import import_module
|
|
2
|
+
|
|
3
|
+
_PT_EXPORTS = (
|
|
4
|
+
"est_dtft",
|
|
5
|
+
"dtft_sum",
|
|
6
|
+
"extract_raw_pt",
|
|
7
|
+
"sniffer_sub",
|
|
8
|
+
"plot_multich_comparison",
|
|
9
|
+
"pickcoilsbycorr",
|
|
10
|
+
"check_waveform_polarity",
|
|
11
|
+
"extract_pilottone_navs",
|
|
12
|
+
"calibrate_pt",
|
|
13
|
+
"apply_pt_calib",
|
|
14
|
+
"process_cplx_pt",
|
|
15
|
+
"pick_cardiac_source",
|
|
16
|
+
"pick_source_bypeak",
|
|
17
|
+
"pick_navigators_from_sources",
|
|
18
|
+
"pred_scan",
|
|
19
|
+
)
|
|
20
|
+
|
|
21
|
+
_SELFNAV_EXPORTS = (
|
|
22
|
+
"extract_selfnav_navs",
|
|
23
|
+
)
|
|
24
|
+
|
|
25
|
+
_TRIGGERING_EXPORTS = (
|
|
26
|
+
"repair_cardiac_triggers_rr",
|
|
27
|
+
"repair_ecg_triggers_with_pt",
|
|
28
|
+
)
|
|
29
|
+
|
|
30
|
+
__all__ = [*_PT_EXPORTS, *_SELFNAV_EXPORTS, *_TRIGGERING_EXPORTS, "main"]
|
|
31
|
+
|
|
32
|
+
|
|
33
|
+
def __getattr__(name: str):
|
|
34
|
+
if name in _PT_EXPORTS:
|
|
35
|
+
module_name = ".model_selection" if name in {"pick_navigators_from_sources", "pred_scan"} else ".pt"
|
|
36
|
+
module = import_module(module_name, __name__)
|
|
37
|
+
value = getattr(module, name)
|
|
38
|
+
globals()[name] = value
|
|
39
|
+
return value
|
|
40
|
+
if name in _SELFNAV_EXPORTS:
|
|
41
|
+
module = import_module(".selfnav", __name__)
|
|
42
|
+
value = getattr(module, name)
|
|
43
|
+
globals()[name] = value
|
|
44
|
+
return value
|
|
45
|
+
if name in _TRIGGERING_EXPORTS:
|
|
46
|
+
module = import_module(".triggering", __name__)
|
|
47
|
+
value = getattr(module, name)
|
|
48
|
+
globals()[name] = value
|
|
49
|
+
return value
|
|
50
|
+
raise AttributeError(f"module {__name__!r} has no attribute {name!r}")
|
|
51
|
+
|
|
52
|
+
|
|
53
|
+
def __dir__() -> list[str]:
|
|
54
|
+
return sorted(set(globals()) | set(__all__))
|
|
55
|
+
|
|
56
|
+
|
|
57
|
+
def main() -> None:
|
|
58
|
+
print("Hello from pylottone!")
|
pylottone/constants.py
ADDED
|
@@ -0,0 +1,8 @@
|
|
|
1
|
+
PILOTTONE_WAVEFORM_ID = 1025
|
|
2
|
+
PILOTTONE_CH = {'RESP': 0, 'CARDIAC': 1, 'CARDIAC_TRIGGERS': 2, 'CARDIAC_DERIVATIVE': 3, 'DERIVATIVE_TRIGGERS': 4}
|
|
3
|
+
ECG_WAVEFORM_ID = 0
|
|
4
|
+
PULSEOX_WAVEFORM_ID = 1
|
|
5
|
+
RESP_WAVEFORM_ID = 2
|
|
6
|
+
EXT1_WAVEFORM_ID = 3
|
|
7
|
+
EXT2_WAVEFORM_ID = 4
|
|
8
|
+
RESPPT_WAVEFORM_ID = 12
|
pylottone/editer.py
ADDED
|
@@ -0,0 +1,252 @@
|
|
|
1
|
+
from warnings import warn
|
|
2
|
+
import numpy as np
|
|
3
|
+
import numpy.typing as npt
|
|
4
|
+
from numpy.fft import ifft, ifftshift
|
|
5
|
+
import math
|
|
6
|
+
|
|
7
|
+
try:
|
|
8
|
+
import cupy as cp
|
|
9
|
+
except Exception as exc:
|
|
10
|
+
cp = None
|
|
11
|
+
_cupy_import_error = exc
|
|
12
|
+
warn(f'CuPy is unavailable; EDITER will use the CPU path. ({exc!r})')
|
|
13
|
+
|
|
14
|
+
import scipy as sp
|
|
15
|
+
|
|
16
|
+
|
|
17
|
+
def _cupy_ready() -> bool:
|
|
18
|
+
return cp is not None
|
|
19
|
+
|
|
20
|
+
|
|
21
|
+
def _to_numpy(array):
|
|
22
|
+
if cp is not None and hasattr(array, 'get'):
|
|
23
|
+
return array.get()
|
|
24
|
+
return array
|
|
25
|
+
|
|
26
|
+
def to_hybrid_kspace(indata):
|
|
27
|
+
'''Centered ifft on first dimension. Does not do fftshift before ifft, as it treats data as time signal.'''
|
|
28
|
+
return ifftshift(ifft(indata, None, axis=0), axes=0)
|
|
29
|
+
|
|
30
|
+
def autopick_sensing_coils(data, f_emi, bw_emi, bw_sig, f_samp, ratio_th=None, n_sensing=None):
|
|
31
|
+
|
|
32
|
+
# Ensure CPU array for operations (Numba or SciPy paths expect numpy arrays)
|
|
33
|
+
if cp is not None and hasattr(data, '__cuda_array_interface__'):
|
|
34
|
+
data = _to_numpy(data)
|
|
35
|
+
|
|
36
|
+
n_samp = data.shape[0]
|
|
37
|
+
df = f_samp/n_samp
|
|
38
|
+
freq_axis = np.arange(0, f_samp, df) - (f_samp - (n_samp % 2)*df)/2 # Handles both even and odd length signals.
|
|
39
|
+
signal_mask = (freq_axis < bw_sig/2) & (freq_axis > -bw_sig/2)
|
|
40
|
+
signal_region = to_hybrid_kspace(data[:, 0,:].squeeze())*signal_mask[:,None]
|
|
41
|
+
|
|
42
|
+
emi_mask = (freq_axis < (f_emi+bw_emi/2)) & (freq_axis > (f_emi-bw_emi/2))
|
|
43
|
+
emi_region = to_hybrid_kspace(data[:, 0,:].squeeze())*emi_mask[:,None]
|
|
44
|
+
|
|
45
|
+
emi_energy = np.sum(np.abs(emi_region), axis=0)
|
|
46
|
+
signal_energy = np.sum(np.abs(signal_region), axis=0)
|
|
47
|
+
sratio = signal_energy/emi_energy
|
|
48
|
+
|
|
49
|
+
Isig = np.argsort(signal_energy)
|
|
50
|
+
Irat = np.argsort(sratio)
|
|
51
|
+
|
|
52
|
+
sratio /= sratio[Isig[-1]] # Normalize with the largest Signal coil to reduce dep on PT amplitude.
|
|
53
|
+
|
|
54
|
+
# plt.figure()
|
|
55
|
+
# plt.plot(np.abs(to_hybrid_kspace(data[:,0,:].squeeze())))
|
|
56
|
+
|
|
57
|
+
# fig, ax1 = plt.subplots()
|
|
58
|
+
# ax2 = ax1.twinx()
|
|
59
|
+
# ax1.plot(emi_energy, label='emi_energy')
|
|
60
|
+
# ax1.plot(signal_energy, label='signal_energy')
|
|
61
|
+
# ax1.legend()
|
|
62
|
+
# ax2.plot(sratio, 'green', label='ratio')
|
|
63
|
+
# ax2.legend()
|
|
64
|
+
# plt.show()
|
|
65
|
+
# plt.figure()
|
|
66
|
+
# plt.plot(np.abs(signal_region))
|
|
67
|
+
# plt.show()
|
|
68
|
+
|
|
69
|
+
# plt.figure()
|
|
70
|
+
# plt.plot(np.abs(emi_region))
|
|
71
|
+
# plt.show()
|
|
72
|
+
|
|
73
|
+
# print(f'Coil with highest signal {coil_name[Isig[-1]]}.')
|
|
74
|
+
|
|
75
|
+
# print(coil_name[Isig])
|
|
76
|
+
|
|
77
|
+
# print(coil_name[Irat])
|
|
78
|
+
|
|
79
|
+
# print(f'Coils with smallest ratio {coil_name[sratio < ratio_th]}')
|
|
80
|
+
|
|
81
|
+
n_ch = data.shape[2]
|
|
82
|
+
if n_sensing is not None and ratio_th is not None:
|
|
83
|
+
warn('Both n_sensing and ratio_th are set. Using n_sensing.')
|
|
84
|
+
if n_sensing is not None:
|
|
85
|
+
sensing_coils = Irat[:n_sensing]
|
|
86
|
+
elif ratio_th is not None:
|
|
87
|
+
sensing_coils = np.nonzero(sratio < ratio_th)[0]
|
|
88
|
+
else:
|
|
89
|
+
raise ValueError('Either n_sensing or ratio_th must be set.')
|
|
90
|
+
|
|
91
|
+
mri_coils = np.arange(n_ch)
|
|
92
|
+
mri_coils = mri_coils[~np.isin(mri_coils, sensing_coils)]
|
|
93
|
+
return mri_coils, sensing_coils
|
|
94
|
+
|
|
95
|
+
|
|
96
|
+
|
|
97
|
+
def get_noise_mtx(
|
|
98
|
+
line_grp: npt.NDArray[np.complex64],
|
|
99
|
+
dk: list[int],
|
|
100
|
+
xp=None,
|
|
101
|
+
):
|
|
102
|
+
"""
|
|
103
|
+
Creates the shifted noise matrix for a given line group and kernel sizes.
|
|
104
|
+
|
|
105
|
+
Args:
|
|
106
|
+
line_grp (numpy.ndarray): Line group data of shape (Nsamples, Nlines, Nchannels).
|
|
107
|
+
dk (list or tuple): Kernel size [d_kx, d_ky].
|
|
108
|
+
|
|
109
|
+
Returns:
|
|
110
|
+
numpy.ndarray: Noise matrix of shape ((Nsamples * Nlines) x (Nchannels * (d_kx * 2 + 1) * (d_ky * 2 + 1))).
|
|
111
|
+
"""
|
|
112
|
+
if xp is None:
|
|
113
|
+
if cp is not None and hasattr(line_grp, '__cuda_array_interface__'):
|
|
114
|
+
xp = cp.get_array_module(line_grp)
|
|
115
|
+
else:
|
|
116
|
+
xp = np
|
|
117
|
+
if cp is None and xp is not np:
|
|
118
|
+
xp = np
|
|
119
|
+
d_kx = dk[0]
|
|
120
|
+
d_ky = dk[1]
|
|
121
|
+
|
|
122
|
+
# noise_mat = np.zeros((line_grp.shape[0], line_grp.shape[1], n_ch*(d_kx*2+1)*(d_ky*2+1)), dtype=line_grp.dtype)
|
|
123
|
+
noise_mat = []
|
|
124
|
+
|
|
125
|
+
dfp = xp.pad(line_grp, ((d_kx, d_kx), (d_ky, d_ky), (0, 0)), mode='constant')
|
|
126
|
+
if d_ky == 0:
|
|
127
|
+
end_slc = None
|
|
128
|
+
else:
|
|
129
|
+
end_slc = d_ky
|
|
130
|
+
ii = 0
|
|
131
|
+
for col_shift in range(-d_kx, d_kx + 1):
|
|
132
|
+
for lin_shift in range(-d_ky, d_ky + 1):
|
|
133
|
+
dftmp = xp.roll(dfp, shift=(col_shift, lin_shift), axis=(0, 1))
|
|
134
|
+
cropped = dftmp[d_kx:-d_kx, d_ky:end_slc, :]
|
|
135
|
+
# noise_mat[:,:,(ii*n_ch):((ii+1)*n_ch)] = dftmp[d_kx:-d_kx, d_ky:end_slc, :]
|
|
136
|
+
ii += 1
|
|
137
|
+
|
|
138
|
+
noise_mat.append(cropped)
|
|
139
|
+
|
|
140
|
+
noise_mat = xp.concatenate(noise_mat, axis=2)
|
|
141
|
+
noise_mat = noise_mat.reshape(noise_mat.shape[0] * noise_mat.shape[1], -1)
|
|
142
|
+
|
|
143
|
+
return noise_mat
|
|
144
|
+
|
|
145
|
+
|
|
146
|
+
def _denoise_sniffer_window(
|
|
147
|
+
snf,
|
|
148
|
+
denoise_rank: int,
|
|
149
|
+
xp,
|
|
150
|
+
):
|
|
151
|
+
"""Low-rank denoise of a sniffer window along coil dimension.
|
|
152
|
+
|
|
153
|
+
Args:
|
|
154
|
+
snf: Array of shape (Ncol, Nlin_window, Nsniffer_coils).
|
|
155
|
+
denoise_rank: Number of singular values/components to keep.
|
|
156
|
+
xp: Backend module (numpy or cupy).
|
|
157
|
+
"""
|
|
158
|
+
if denoise_rank <= 0:
|
|
159
|
+
return snf
|
|
160
|
+
|
|
161
|
+
n_col, n_lin, n_ch = snf.shape
|
|
162
|
+
snf_mat = snf.reshape(n_col * n_lin, n_ch)
|
|
163
|
+
min_dim = min(snf_mat.shape)
|
|
164
|
+
keep_rank = min(denoise_rank, min_dim)
|
|
165
|
+
|
|
166
|
+
if keep_rank >= min_dim:
|
|
167
|
+
return snf
|
|
168
|
+
|
|
169
|
+
u, s, vh = xp.linalg.svd(snf_mat, full_matrices=False)
|
|
170
|
+
snf_denoised = (u[:, :keep_rank] * s[:keep_rank]) @ vh[:keep_rank, :]
|
|
171
|
+
return snf_denoised.reshape(n_col, n_lin, n_ch)
|
|
172
|
+
|
|
173
|
+
|
|
174
|
+
def _est_emi_impl(signal_in, sniffer, line_grps, dk, w, denoise_sniffers, denoise_rank, xp, to_numpy):
|
|
175
|
+
Ncol, Nlin, Nc = sniffer.shape
|
|
176
|
+
emi_hat = xp.zeros((Ncol, Nlin), dtype=np.complex64)
|
|
177
|
+
kern = []
|
|
178
|
+
|
|
179
|
+
for cwin, pe_rng in enumerate(line_grps):
|
|
180
|
+
# Denoise sniffer channels per window before constructing the noise matrix.
|
|
181
|
+
snf = xp.asarray(sniffer[:, pe_rng, :])
|
|
182
|
+
if denoise_sniffers:
|
|
183
|
+
snf = _denoise_sniffer_window(snf, denoise_rank=denoise_rank, xp=xp)
|
|
184
|
+
|
|
185
|
+
# Build noise matrix and input vectors in the chosen backend
|
|
186
|
+
noise_mat = get_noise_mtx(snf, dk, xp=xp)
|
|
187
|
+
|
|
188
|
+
init_mat_sub = xp.reshape(xp.asarray(signal_in[:, pe_rng]), (Ncol * len(pe_rng), 1))
|
|
189
|
+
ww = xp.reshape(xp.asarray(w[:, pe_rng]), (Ncol * len(pe_rng), 1))
|
|
190
|
+
|
|
191
|
+
# Use SciPy's lstsq on CPU for consistency with previous implementation,
|
|
192
|
+
# and CuPy's lstsq on GPU when available.
|
|
193
|
+
if xp is np:
|
|
194
|
+
kern_ ,_,_,_ = sp.linalg.lstsq(ww * noise_mat, ww * init_mat_sub, cond=None, check_finite=False)
|
|
195
|
+
else:
|
|
196
|
+
kern_ ,_,_,_ = xp.linalg.lstsq(ww * noise_mat, ww * init_mat_sub, rcond=None)
|
|
197
|
+
|
|
198
|
+
kern.append((pe_rng, to_numpy(kern_)))
|
|
199
|
+
emi_hat[:, pe_rng] = xp.reshape(xp.dot(noise_mat, kern_), (Ncol, len(pe_rng)))
|
|
200
|
+
|
|
201
|
+
return to_numpy(emi_hat), kern
|
|
202
|
+
|
|
203
|
+
def est_emi(
|
|
204
|
+
signal_in: npt.NDArray[np.complex64],
|
|
205
|
+
sniffer: npt.NDArray[np.complex64],
|
|
206
|
+
line_grps: list[npt.NDArray],
|
|
207
|
+
dk: list[int],
|
|
208
|
+
w: npt.NDArray[np.float32],
|
|
209
|
+
denoise_sniffers: bool = False,
|
|
210
|
+
denoise_rank: int = 1,
|
|
211
|
+
):
|
|
212
|
+
return _est_emi_impl(signal_in, sniffer, line_grps, dk, w, denoise_sniffers, denoise_rank, np, lambda x: x)
|
|
213
|
+
|
|
214
|
+
def est_emi_gpu(
|
|
215
|
+
signal_in: npt.NDArray[np.complex64],
|
|
216
|
+
sniffer: npt.NDArray[np.complex64],
|
|
217
|
+
line_grps: list[npt.NDArray],
|
|
218
|
+
dk: list[int],
|
|
219
|
+
w: npt.NDArray[np.float32],
|
|
220
|
+
denoise_sniffers: bool = False,
|
|
221
|
+
denoise_rank: int = 1,
|
|
222
|
+
):
|
|
223
|
+
if cp is None:
|
|
224
|
+
warn('CuPy is unavailable; falling back to the CPU EDITER path.')
|
|
225
|
+
return est_emi(signal_in, sniffer, line_grps, dk, w, denoise_sniffers, denoise_rank)
|
|
226
|
+
return _est_emi_impl(signal_in, sniffer, line_grps, dk, w, denoise_sniffers, denoise_rank, cp, _to_numpy)
|
|
227
|
+
|
|
228
|
+
def apply_editer(signal_in: npt.NDArray[np.complex64], sniffer: npt.NDArray[np.complex64], params, w) -> tuple[npt.NDArray[np.complex64], npt.NDArray[np.complex64]]:
|
|
229
|
+
max_lines = params['max_lines_per_group']
|
|
230
|
+
nlin = signal_in.shape[1]
|
|
231
|
+
if params['grouping_method'] == "uniform":
|
|
232
|
+
Ngrp = math.ceil(nlin/max_lines)
|
|
233
|
+
line_grps = []
|
|
234
|
+
|
|
235
|
+
for grp_i in range(Ngrp):
|
|
236
|
+
line_grps.append(np.arange(((grp_i)*max_lines), min(max_lines*(grp_i+1), nlin)))
|
|
237
|
+
|
|
238
|
+
|
|
239
|
+
denoise_sniffers = bool(params.get('denoise_sniffers', False))
|
|
240
|
+
denoise_rank = int(params.get('denoise_rank', 1))
|
|
241
|
+
|
|
242
|
+
if params['gpu'] == -1 or not _cupy_ready():
|
|
243
|
+
emi_hat, kernels = est_emi(signal_in, sniffer, line_grps, params['dk'], w, denoise_sniffers, denoise_rank)
|
|
244
|
+
else:
|
|
245
|
+
try:
|
|
246
|
+
with cp.cuda.Device(params['gpu']):
|
|
247
|
+
emi_hat, kernels = est_emi_gpu(signal_in, sniffer, line_grps, params['dk'], w, denoise_sniffers, denoise_rank)
|
|
248
|
+
except Exception as exc:
|
|
249
|
+
warn(f'CuPy/CUDA EDITER execution failed; falling back to CPU. ({exc!r})')
|
|
250
|
+
emi_hat, kernels = est_emi(signal_in, sniffer, line_grps, params['dk'], w, denoise_sniffers, denoise_rank)
|
|
251
|
+
return emi_hat, kernels
|
|
252
|
+
|
|
@@ -0,0 +1,94 @@
|
|
|
1
|
+
import joblib
|
|
2
|
+
import numpy as np
|
|
3
|
+
from scipy.signal import firwin, convolve
|
|
4
|
+
from scipy.interpolate import pchip_interpolate
|
|
5
|
+
import warnings
|
|
6
|
+
|
|
7
|
+
def pred_scan(rocket_pipeline, scan, force_navpred=False):
|
|
8
|
+
''' Predict navigators in the input scan using a pre-trained ROCKET classifier, given a set of sources.
|
|
9
|
+
A label of 0 indicates non-navigator, 1 indicates respiratory navigator, and 2 indicates cardiac navigator.
|
|
10
|
+
Parameters
|
|
11
|
+
----------
|
|
12
|
+
rocket_pipeline : sklearn.pipeline.Pipeline
|
|
13
|
+
Pre-trained ROCKET classifier pipeline.
|
|
14
|
+
scan : np.ndarray
|
|
15
|
+
Input scan sources, shape (n_sources, n_samples)
|
|
16
|
+
force_navpred : bool
|
|
17
|
+
If True, forces the function to return a respiratory and cardiac navigator even if the classifier does not predict any.
|
|
18
|
+
Returns
|
|
19
|
+
-------
|
|
20
|
+
y_pred_ : np.ndarray
|
|
21
|
+
Predicted labels for each source, shape (n_sources,).
|
|
22
|
+
confs_ : np.ndarray
|
|
23
|
+
Confidence scores for each source, shape (n_sources, n_classes).
|
|
24
|
+
'''
|
|
25
|
+
with warnings.catch_warnings():
|
|
26
|
+
warnings.simplefilter("ignore") # Workaround, ignore sklearn tag warnings.
|
|
27
|
+
confs_ = rocket_pipeline.decision_function(scan[:,None,:])
|
|
28
|
+
|
|
29
|
+
y_pred_ = np.argmax(confs_, axis=1)
|
|
30
|
+
|
|
31
|
+
# Resolve multiple positive predictions
|
|
32
|
+
if np.sum(y_pred_ == 1) > 1:
|
|
33
|
+
resp_preds = (y_pred_ == 1).nonzero()[0]
|
|
34
|
+
conf_r_ = confs_[resp_preds, 1]
|
|
35
|
+
top_resp_idx = resp_preds[np.argmax(conf_r_)]
|
|
36
|
+
other_idxs = np.setdiff1d(resp_preds, np.array([top_resp_idx]))
|
|
37
|
+
y_pred_[other_idxs] = 0
|
|
38
|
+
warnings.warn(f"Multiple respiratory predictions found at indices {resp_preds}, keeping index {top_resp_idx} only.")
|
|
39
|
+
if np.sum(y_pred_ == 2) > 1:
|
|
40
|
+
card_preds = (y_pred_ == 2).nonzero()[0]
|
|
41
|
+
conf_r_ = confs_[card_preds, 2]
|
|
42
|
+
top_card_idx = card_preds[np.argmax(conf_r_)]
|
|
43
|
+
other_idxs = np.setdiff1d(card_preds, np.array([top_card_idx]))
|
|
44
|
+
y_pred_[other_idxs] = 0
|
|
45
|
+
warnings.warn(f"Multiple cardiac predictions found at indices {card_preds}, keeping index {top_card_idx} only.")
|
|
46
|
+
if force_navpred:
|
|
47
|
+
if np.sum(y_pred_ == 1) == 0:
|
|
48
|
+
y_pred_[np.argmax(confs_[:,1])] = 1
|
|
49
|
+
warnings.warn("No respiratory prediction was found, but force_navpred is True. Forcing the highest confidence prediction as respiratory.")
|
|
50
|
+
if np.sum(y_pred_ == 2) == 0:
|
|
51
|
+
y_pred_[np.argmax(confs_[:,2])] = 2
|
|
52
|
+
warnings.warn("No cardiac prediction was found, but force_navpred is True. Forcing the highest confidence prediction as cardiac.")
|
|
53
|
+
return y_pred_, confs_
|
|
54
|
+
|
|
55
|
+
|
|
56
|
+
def pick_navigators_from_sources(sources, time_vec, classifier_path='rocket_pipeline.pkl', force_navpred=False):
|
|
57
|
+
''' Pick respiratory and cardiac navigators from input sources using a pre-trained ROCKET classifier.
|
|
58
|
+
Parameters
|
|
59
|
+
----------
|
|
60
|
+
sources : np.ndarray
|
|
61
|
+
Input sources, shape (n_sources, n_samples)
|
|
62
|
+
time_vec : np.ndarray
|
|
63
|
+
Time vector corresponding to the sources, unit is seconds, shape (n_samples,)
|
|
64
|
+
classifier_path : str
|
|
65
|
+
Path to the pre-trained ROCKET classifier. Must be compatible with joblib.load().
|
|
66
|
+
force_navpred : bool
|
|
67
|
+
If True, forces the function to return a respiratory and cardiac navigator even if the classifier does not predict any.
|
|
68
|
+
Returns
|
|
69
|
+
-------
|
|
70
|
+
resp_idx : int
|
|
71
|
+
Index of the respiratory navigator source in the input sources.
|
|
72
|
+
card_idx : int
|
|
73
|
+
Index of the cardiac navigator source in the input sources.
|
|
74
|
+
confs : np.ndarray
|
|
75
|
+
Confidence scores for each source, shape (n_sources, n_classes).
|
|
76
|
+
'''
|
|
77
|
+
rocket_pipeline = joblib.load(classifier_path)
|
|
78
|
+
n_samp = sources.shape[1]
|
|
79
|
+
dt_samp = time_vec[1] - time_vec[0]
|
|
80
|
+
|
|
81
|
+
h_denoise = firwin(2*(n_samp//8)-1, [0.1, 6], fs=1/dt_samp, window=('tukey', 1), pass_zero=False)
|
|
82
|
+
sources_filt = convolve(sources, h_denoise[None, :], mode='same')
|
|
83
|
+
|
|
84
|
+
dt_new = 10e-3 # 10 ms
|
|
85
|
+
n_samp_new = int(np.ceil(n_samp * dt_samp / dt_new))
|
|
86
|
+
time_new = np.arange(0, n_samp_new)*dt_new
|
|
87
|
+
sources_resampled = pchip_interpolate(time_vec, sources_filt, time_new, axis=1)
|
|
88
|
+
sources_resampled -= np.mean(sources_resampled, axis=1, keepdims=True)
|
|
89
|
+
sources_resampled /= np.std(sources_resampled, axis=1, keepdims=True)
|
|
90
|
+
|
|
91
|
+
y_pred, confs = pred_scan(rocket_pipeline, sources_resampled, force_navpred=force_navpred)
|
|
92
|
+
resp_idx = np.where(y_pred == 1)[0]
|
|
93
|
+
card_idx = np.where(y_pred == 2)[0]
|
|
94
|
+
return resp_idx, card_idx, confs
|