flimkit 0.12.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- flimkit/FLIM/__init__.py +0 -0
- flimkit/FLIM/assemble.py +254 -0
- flimkit/FLIM/batch.py +681 -0
- flimkit/FLIM/bg_tools.py +51 -0
- flimkit/FLIM/fit_tools.py +244 -0
- flimkit/FLIM/fitters.py +1471 -0
- flimkit/FLIM/irf_tools.py +617 -0
- flimkit/FLIM/models.py +391 -0
- flimkit/GPU/__init__.py +85 -0
- flimkit/GPU/_base.py +391 -0
- flimkit/GPU/cuda.py +10 -0
- flimkit/GPU/mlx_backend.py +381 -0
- flimkit/GPU/mps.py +10 -0
- flimkit/GPU/rocm.py +10 -0
- flimkit/GPU/torch_backend.py +385 -0
- flimkit/UI/app_state.py +10 -0
- flimkit/UI/controller.py +139 -0
- flimkit/UI/expert_settings.py +248 -0
- flimkit/UI/fit_help.py +206 -0
- flimkit/UI/fov_preview.py +1085 -0
- flimkit/UI/gui.py +3919 -0
- flimkit/UI/icon.icns +0 -0
- flimkit/UI/icon.ico +0 -0
- flimkit/UI/icon.png +0 -0
- flimkit/UI/irf_widget.py +103 -0
- flimkit/UI/mode_controller.py +118 -0
- flimkit/UI/modes/__init__.py +0 -0
- flimkit/UI/modes/base.py +3 -0
- flimkit/UI/modes/batch_mode.py +312 -0
- flimkit/UI/modes/fov_mode.py +164 -0
- flimkit/UI/modes/irf_mode.py +80 -0
- flimkit/UI/modes/phasor_mode.py +131 -0
- flimkit/UI/modes/stitch_mode.py +254 -0
- flimkit/UI/phasor_panel.py +1087 -0
- flimkit/UI/progress_window.py +113 -0
- flimkit/UI/project_panel.py +262 -0
- flimkit/UI/results_panel.py +332 -0
- flimkit/UI/roi_tools.py +794 -0
- flimkit/UI/utils.py +217 -0
- flimkit/__init__.py +0 -0
- flimkit/_version.py +41 -0
- flimkit/cli.py +120 -0
- flimkit/configs.py +148 -0
- flimkit/dialogs.py +46 -0
- flimkit/formats/BH/__init__.py +0 -0
- flimkit/formats/BH/reader.py +296 -0
- flimkit/formats/BH/writer.py +86 -0
- flimkit/formats/ISS/__init__.py +0 -0
- flimkit/formats/ISS/fdflim.py +86 -0
- flimkit/formats/ISS/image.py +114 -0
- flimkit/formats/ISS/reader.py +223 -0
- flimkit/formats/PS/__init__.py +0 -0
- flimkit/formats/PS/reader.py +202 -0
- flimkit/formats/PTU/__init__.py +0 -0
- flimkit/formats/PTU/decode.py +27 -0
- flimkit/formats/PTU/phu.py +85 -0
- flimkit/formats/PTU/reader.py +235 -0
- flimkit/formats/PTU/series.py +258 -0
- flimkit/formats/PTU/stitch.py +1182 -0
- flimkit/formats/PTU/tools.py +94 -0
- flimkit/formats/__init__.py +2 -0
- flimkit/formats/flim_file.py +232 -0
- flimkit/formats/phasor.py +132 -0
- flimkit/formats/signal.py +170 -0
- flimkit/image/tools.py +124 -0
- flimkit/interactive.py +1857 -0
- flimkit/mpl_backend.py +22 -0
- flimkit/phasor/__init__.py +40 -0
- flimkit/phasor/filters.py +127 -0
- flimkit/phasor/fret.py +654 -0
- flimkit/phasor/interactive.py +556 -0
- flimkit/phasor/peaks.py +186 -0
- flimkit/phasor/signal.py +90 -0
- flimkit/phasor_launcher.py +314 -0
- flimkit/plugins/__init__.py +137 -0
- flimkit/plugins/bindings.py +116 -0
- flimkit/plugins/builtin/__init__.py +3 -0
- flimkit/plugins/builtin/core_tools.py +28 -0
- flimkit/plugins/loader.py +371 -0
- flimkit/plugins/registry.py +406 -0
- flimkit/project.py +197 -0
- flimkit/synth.py +145 -0
- flimkit/utils/__init__.py +0 -0
- flimkit/utils/batch_fit.py +301 -0
- flimkit/utils/config_manager.py +119 -0
- flimkit/utils/config_snapshot.py +30 -0
- flimkit/utils/crash_handler.py +183 -0
- flimkit/utils/display.py +197 -0
- flimkit/utils/enhanced_outputs.py +345 -0
- flimkit/utils/fancy.py +103 -0
- flimkit/utils/lifetime_image.py +243 -0
- flimkit/utils/misc.py +111 -0
- flimkit/utils/plotting.py +190 -0
- flimkit/utils/roi.py +370 -0
- flimkit/utils/session.py +51 -0
- flimkit/utils/update_check.py +198 -0
- flimkit/utils/xlsx_tools.py +97 -0
- flimkit/utils/xml_utils.py +219 -0
- flimkit-0.12.0.dist-info/METADATA +356 -0
- flimkit-0.12.0.dist-info/RECORD +104 -0
- flimkit-0.12.0.dist-info/WHEEL +5 -0
- flimkit-0.12.0.dist-info/entry_points.txt +2 -0
- flimkit-0.12.0.dist-info/licenses/LICENSE.md +11 -0
- flimkit-0.12.0.dist-info/top_level.txt +1 -0
|
@@ -0,0 +1,381 @@
|
|
|
1
|
+
import numpy as np
|
|
2
|
+
from flimkit.GPU._base import _BackendMixin, fit_window, pixel_blocks
|
|
3
|
+
from flimkit.FLIM.fit_tools import (calibrated_chi2, distribution_dof,
|
|
4
|
+
estimate_bg, coates_pileup_correction)
|
|
5
|
+
|
|
6
|
+
class MLXBackend(_BackendMixin):
|
|
7
|
+
|
|
8
|
+
def __init__(self):
|
|
9
|
+
import mlx.core as mx
|
|
10
|
+
self._mx = mx
|
|
11
|
+
self.device = mx.Device(mx.gpu)
|
|
12
|
+
|
|
13
|
+
def __repr__(self):
|
|
14
|
+
return "MLXBackend(device='metal')"
|
|
15
|
+
|
|
16
|
+
def batch_fixed_tau(
|
|
17
|
+
self,
|
|
18
|
+
stack,
|
|
19
|
+
A,
|
|
20
|
+
taus_fixed,
|
|
21
|
+
min_photons,
|
|
22
|
+
correct_pileup,
|
|
23
|
+
n_sync_px,
|
|
24
|
+
progress_callback=None,
|
|
25
|
+
tvb_profile=None,
|
|
26
|
+
fit_tvb=False,
|
|
27
|
+
fit_idx=None,
|
|
28
|
+
):
|
|
29
|
+
mx = self._mx
|
|
30
|
+
ny, nx, n_bins = stack.shape
|
|
31
|
+
n_exp = A.shape[1]
|
|
32
|
+
taus_ns = taus_fixed * 1e9
|
|
33
|
+
raw = stack.reshape(ny * nx, n_bins)
|
|
34
|
+
win = fit_window(fit_idx, n_bins)
|
|
35
|
+
A = A if win is None else A[win]
|
|
36
|
+
valid_idx = np.where(raw.sum(axis=1) >= min_photons)[0]
|
|
37
|
+
maps = self._init_maps(
|
|
38
|
+
ny, nx, n_exp,
|
|
39
|
+
intensity=stack.sum(axis=2),
|
|
40
|
+
taus_fixed_ns=taus_ns,
|
|
41
|
+
free_tau=False,
|
|
42
|
+
)
|
|
43
|
+
if valid_idx.size == 0:
|
|
44
|
+
return maps
|
|
45
|
+
with_tvb = fit_tvb and tvb_profile is not None
|
|
46
|
+
B_col = None
|
|
47
|
+
if with_tvb:
|
|
48
|
+
B_col = np.asarray(tvb_profile, dtype=np.float32)
|
|
49
|
+
B_col = B_col if win is None else B_col[win]
|
|
50
|
+
A_aug = np.column_stack(
|
|
51
|
+
[A, B_col, np.ones(A.shape[0], dtype=np.float32)]).astype(np.float32)
|
|
52
|
+
A_pinv_mx = mx.linalg.pinv(mx.array(A_aug), stream=mx.cpu)
|
|
53
|
+
else:
|
|
54
|
+
A_pinv_mx = mx.linalg.pinv(mx.array(A.astype(np.float32)), stream=mx.cpu)
|
|
55
|
+
n_fit = A.shape[0]
|
|
56
|
+
for first, last in pixel_blocks(valid_idx.size, 4 * (2 * n_bins + n_fit + n_exp)):
|
|
57
|
+
block = valid_idx[first:last]
|
|
58
|
+
decay = raw[block].astype(np.float32)
|
|
59
|
+
if with_tvb:
|
|
60
|
+
data_in = decay if win is None else decay[:, win]
|
|
61
|
+
data_in = data_in.copy()
|
|
62
|
+
if correct_pileup and n_sync_px > 0:
|
|
63
|
+
for row in range(data_in.shape[0]):
|
|
64
|
+
data_in[row] = coates_pileup_correction(data_in[row], n_sync_px)
|
|
65
|
+
coeffs_mx = mx.maximum(mx.array(data_in) @ A_pinv_mx.T, 0.0)
|
|
66
|
+
mx.eval(coeffs_mx)
|
|
67
|
+
coeffs = np.array(coeffs_mx)
|
|
68
|
+
amps = coeffs[:, :n_exp]
|
|
69
|
+
tvb = coeffs[:, n_exp]
|
|
70
|
+
bg = coeffs[:, n_exp + 1]
|
|
71
|
+
decay_fit = data_in
|
|
72
|
+
else:
|
|
73
|
+
bg = self._estimate_bg_batch(decay, np.ones(decay.shape[0], dtype=bool))
|
|
74
|
+
corrected = np.maximum(decay - bg[:, None], 0.0)
|
|
75
|
+
corrected = corrected if win is None else corrected[:, win]
|
|
76
|
+
if correct_pileup and n_sync_px > 0:
|
|
77
|
+
for row in range(corrected.shape[0]):
|
|
78
|
+
corrected[row] = coates_pileup_correction(corrected[row], n_sync_px)
|
|
79
|
+
amps_mx = mx.maximum(mx.array(corrected) @ A_pinv_mx.T, 0.0)
|
|
80
|
+
mx.eval(amps_mx)
|
|
81
|
+
amps = np.array(amps_mx)
|
|
82
|
+
tvb = None
|
|
83
|
+
decay_fit = decay if win is None else decay[:, win]
|
|
84
|
+
self._scatter_fixed_tau(
|
|
85
|
+
maps,
|
|
86
|
+
valid_idx = block,
|
|
87
|
+
amps = amps,
|
|
88
|
+
bg = bg,
|
|
89
|
+
decay_valid = decay_fit,
|
|
90
|
+
A = A,
|
|
91
|
+
taus_ns = taus_ns,
|
|
92
|
+
ny=ny, nx=nx,
|
|
93
|
+
tvb = tvb,
|
|
94
|
+
tvb_profile = B_col if with_tvb else None,
|
|
95
|
+
)
|
|
96
|
+
if progress_callback is not None:
|
|
97
|
+
progress_callback(last, valid_idx.size)
|
|
98
|
+
return maps
|
|
99
|
+
|
|
100
|
+
def batch_grid_scan_1exp(
|
|
101
|
+
self,
|
|
102
|
+
stack,
|
|
103
|
+
basis_grid,
|
|
104
|
+
bb_grid,
|
|
105
|
+
tau_grid,
|
|
106
|
+
min_photons,
|
|
107
|
+
correct_pileup,
|
|
108
|
+
n_sync_px,
|
|
109
|
+
progress_callback=None,
|
|
110
|
+
tvb_profile=None,
|
|
111
|
+
fit_tvb=False,
|
|
112
|
+
fit_idx=None,
|
|
113
|
+
):
|
|
114
|
+
mx = self._mx
|
|
115
|
+
ny, nx, n_bins = stack.shape
|
|
116
|
+
N_GRID = len(tau_grid)
|
|
117
|
+
raw = stack.reshape(ny * nx, n_bins)
|
|
118
|
+
win = fit_window(fit_idx, n_bins)
|
|
119
|
+
if win is not None:
|
|
120
|
+
if fit_tvb and tvb_profile is not None:
|
|
121
|
+
raise ValueError('a fit window with time-varying background is not '
|
|
122
|
+
'supported on the GPU for one-exponential fits')
|
|
123
|
+
basis_grid = basis_grid[:, win]
|
|
124
|
+
bb_grid = np.maximum((basis_grid ** 2).sum(axis=1), 1e-20)
|
|
125
|
+
n_fit = len(win)
|
|
126
|
+
else:
|
|
127
|
+
n_fit = n_bins
|
|
128
|
+
valid_idx = np.where(raw.sum(axis=1) >= min_photons)[0]
|
|
129
|
+
maps = self._init_maps(
|
|
130
|
+
ny, nx, n_exp=1,
|
|
131
|
+
intensity=stack.sum(axis=2),
|
|
132
|
+
taus_fixed_ns=np.array([tau_grid[N_GRID // 2] * 1e9]),
|
|
133
|
+
free_tau=True,
|
|
134
|
+
)
|
|
135
|
+
if valid_idx.size == 0:
|
|
136
|
+
return maps
|
|
137
|
+
with_tvb = fit_tvb and tvb_profile is not None
|
|
138
|
+
if with_tvb:
|
|
139
|
+
U, U_pinv, basis_perp, bb_perp = self._tvb_grid_prep(
|
|
140
|
+
basis_grid, tvb_profile, n_bins)
|
|
141
|
+
basis_mx = mx.array(basis_perp)
|
|
142
|
+
bbp_mx = mx.array(bb_perp)
|
|
143
|
+
else:
|
|
144
|
+
basis_mx = mx.array(basis_grid.astype(np.float32))
|
|
145
|
+
bb_mx = mx.array(bb_grid.astype(np.float32))
|
|
146
|
+
per_pixel = 4 * (2 * n_bins + n_fit + N_GRID)
|
|
147
|
+
for first, last in pixel_blocks(valid_idx.size, per_pixel):
|
|
148
|
+
block = valid_idx[first:last]
|
|
149
|
+
decay = raw[block].astype(np.float32)
|
|
150
|
+
if with_tvb:
|
|
151
|
+
data_in = decay.copy()
|
|
152
|
+
if correct_pileup and n_sync_px > 0:
|
|
153
|
+
for row in range(data_in.shape[0]):
|
|
154
|
+
data_in[row] = coates_pileup_correction(data_in[row], n_sync_px)
|
|
155
|
+
d_perp = self._tvb_project_data(
|
|
156
|
+
data_in.astype(np.float64), U, U_pinv).astype(np.float32)
|
|
157
|
+
dperp_mx = mx.array(d_perp)
|
|
158
|
+
bd_mx = dperp_mx @ basis_mx.T
|
|
159
|
+
dsq_mx = (dperp_mx ** 2).sum(axis=1)
|
|
160
|
+
costs_mx = dsq_mx[:, None] - mx.maximum(bd_mx, 0.0) ** 2 / bbp_mx[None, :]
|
|
161
|
+
best_g_mx = costs_mx.argmin(axis=1)
|
|
162
|
+
mx.eval(best_g_mx, bd_mx)
|
|
163
|
+
best_g = np.array(best_g_mx)
|
|
164
|
+
bd_np = np.array(bd_mx)
|
|
165
|
+
amp_v = np.maximum(
|
|
166
|
+
bd_np[np.arange(block.size), best_g] / bb_perp[best_g], 0.0)
|
|
167
|
+
basis_best = basis_grid[best_g]
|
|
168
|
+
resid_after = data_in.astype(np.float64) - amp_v[:, None] * basis_best
|
|
169
|
+
vz = resid_after @ U_pinv.T
|
|
170
|
+
self._scatter_1exp(
|
|
171
|
+
maps, valid_idx=block, tau_v=tau_grid[best_g], amp_v=amp_v,
|
|
172
|
+
bg_v=vz[:, 1].astype(np.float32), decay_valid=data_in,
|
|
173
|
+
basis_best=basis_best, ny=ny, nx=nx, n_bins=n_bins,
|
|
174
|
+
tvb=np.maximum(vz[:, 0], 0.0).astype(np.float32),
|
|
175
|
+
tvb_profile=np.asarray(tvb_profile, dtype=np.float32),
|
|
176
|
+
)
|
|
177
|
+
else:
|
|
178
|
+
bg = self._estimate_bg_batch(
|
|
179
|
+
decay, np.ones(decay.shape[0], dtype=bool))
|
|
180
|
+
corrected = np.maximum(decay - bg[:, None], 0.0)
|
|
181
|
+
if correct_pileup and n_sync_px > 0:
|
|
182
|
+
for row in range(corrected.shape[0]):
|
|
183
|
+
corrected[row] = coates_pileup_correction(
|
|
184
|
+
corrected[row], n_sync_px)
|
|
185
|
+
corrected = corrected if win is None else corrected[:, win]
|
|
186
|
+
dc_mx = mx.array(corrected)
|
|
187
|
+
bd_mx = dc_mx @ basis_mx.T
|
|
188
|
+
dc_sq_mx = (dc_mx ** 2).sum(axis=1)
|
|
189
|
+
costs_mx = dc_sq_mx[:, None] - mx.maximum(bd_mx, 0.0) ** 2 / bb_mx[None, :]
|
|
190
|
+
best_g_mx = costs_mx.argmin(axis=1)
|
|
191
|
+
mx.eval(best_g_mx, bd_mx)
|
|
192
|
+
best_g = np.array(best_g_mx)
|
|
193
|
+
bd_np = np.array(bd_mx)
|
|
194
|
+
amp_v = np.maximum(
|
|
195
|
+
bd_np[np.arange(block.size), best_g] / bb_grid[best_g], 0.0)
|
|
196
|
+
self._scatter_1exp(
|
|
197
|
+
maps,
|
|
198
|
+
valid_idx = block,
|
|
199
|
+
tau_v = tau_grid[best_g],
|
|
200
|
+
amp_v = amp_v,
|
|
201
|
+
bg_v = bg,
|
|
202
|
+
decay_valid = decay if win is None else decay[:, win],
|
|
203
|
+
basis_best = basis_grid[best_g],
|
|
204
|
+
ny = ny, nx = nx,
|
|
205
|
+
n_bins = n_fit,
|
|
206
|
+
)
|
|
207
|
+
if progress_callback is not None:
|
|
208
|
+
progress_callback(last, valid_idx.size)
|
|
209
|
+
return maps
|
|
210
|
+
|
|
211
|
+
def batch_dist_scan_unimodal(
|
|
212
|
+
self,
|
|
213
|
+
stack,
|
|
214
|
+
basis,
|
|
215
|
+
bb_grid,
|
|
216
|
+
param_pairs,
|
|
217
|
+
irf_fixed,
|
|
218
|
+
tcspc_res,
|
|
219
|
+
n_bins,
|
|
220
|
+
dist_type,
|
|
221
|
+
min_photons,
|
|
222
|
+
progress_callback=None,
|
|
223
|
+
tvb_profile=None,
|
|
224
|
+
fit_tvb=False,
|
|
225
|
+
fit_idx=None,
|
|
226
|
+
):
|
|
227
|
+
mx = self._mx
|
|
228
|
+
ny, nx, _ = stack.shape
|
|
229
|
+
flat = stack.reshape(ny * nx, n_bins).astype(np.float32)
|
|
230
|
+
win = fit_window(fit_idx, n_bins)
|
|
231
|
+
n_fit = n_bins if win is None else len(win)
|
|
232
|
+
flat_fit = flat if win is None else flat[:, win]
|
|
233
|
+
intensity_flat = flat.sum(axis=1)
|
|
234
|
+
valid_mask = intensity_flat >= min_photons
|
|
235
|
+
valid_idx = np.where(valid_mask)[0]
|
|
236
|
+
maps = dict(
|
|
237
|
+
intensity = stack.sum(axis=2),
|
|
238
|
+
tau_mean_amp = np.full((ny, nx), np.nan),
|
|
239
|
+
tau_mean_int = np.full((ny, nx), np.nan),
|
|
240
|
+
chi2_r = np.full((ny, nx), np.nan),
|
|
241
|
+
calibrated_chi2_r = np.full((ny, nx), np.nan),
|
|
242
|
+
tau_center_1 = np.full((ny, nx), np.nan),
|
|
243
|
+
width_1 = np.full((ny, nx), np.nan),
|
|
244
|
+
alpha_1 = np.full((ny, nx), np.nan),
|
|
245
|
+
frac_1 = np.full((ny, nx), np.nan),
|
|
246
|
+
)
|
|
247
|
+
if valid_idx.size == 0:
|
|
248
|
+
return maps
|
|
249
|
+
if fit_tvb and tvb_profile is not None:
|
|
250
|
+
maps['tvb_scale'] = np.full((ny, nx), np.nan)
|
|
251
|
+
tvb_fit = np.asarray(tvb_profile) if win is None else np.asarray(tvb_profile)[win]
|
|
252
|
+
U, U_pinv, basis_perp, bb_perp = self._tvb_grid_prep(basis, tvb_fit, n_fit)
|
|
253
|
+
d_valid = flat_fit[valid_idx]
|
|
254
|
+
d_perp = self._tvb_project_data(d_valid.astype(np.float64), U, U_pinv).astype(np.float32)
|
|
255
|
+
basis_pmx = mx.array(basis_perp)
|
|
256
|
+
bbp_mx = mx.array(bb_perp)
|
|
257
|
+
dperp_mx = mx.array(d_perp)
|
|
258
|
+
bd_mx = dperp_mx @ basis_pmx.T
|
|
259
|
+
dsq_mx = (dperp_mx ** 2).sum(axis=1)
|
|
260
|
+
costs_mx = dsq_mx[:, None] - mx.maximum(bd_mx, 0.0) ** 2 / bbp_mx[None, :]
|
|
261
|
+
best_g_mx = costs_mx.argmin(axis=1)
|
|
262
|
+
mx.eval(best_g_mx, bd_mx)
|
|
263
|
+
best_g = np.array(best_g_mx)
|
|
264
|
+
bd_np = np.array(bd_mx)
|
|
265
|
+
tau_v = param_pairs[best_g, 0]
|
|
266
|
+
w_v = param_pairs[best_g, 1]
|
|
267
|
+
amp_v = np.maximum(bd_np[np.arange(len(valid_idx)), best_g] / bb_perp[best_g], 0.0)
|
|
268
|
+
good = amp_v > 0
|
|
269
|
+
tau_amp_ns = tau_v * 1e9
|
|
270
|
+
tau_int_ns = (tau_v + w_v ** 2 / np.maximum(tau_v, 1e-15)) * 1e9
|
|
271
|
+
basis_best = basis[best_g].astype(np.float64)
|
|
272
|
+
resid_after = d_valid.astype(np.float64) - amp_v[:, None] * basis_best
|
|
273
|
+
vz = resid_after @ U_pinv.T
|
|
274
|
+
tvb_v = np.maximum(vz[:, 0], 0.0)
|
|
275
|
+
bg_z = vz[:, 1]
|
|
276
|
+
B_arr = np.asarray(tvb_fit, dtype=np.float64)
|
|
277
|
+
model_v = amp_v[:, None] * basis_best + tvb_v[:, None] * B_arr[None, :] + bg_z[:, None]
|
|
278
|
+
resid_v = d_valid.astype(np.float64) - model_v
|
|
279
|
+
chi2_v = (resid_v ** 2 / np.maximum(model_v, 1.0)).sum(axis=1) / distribution_dof(n_fit, 1, True)
|
|
280
|
+
chi2_cal_v = calibrated_chi2(d_valid, model_v, axis=1)
|
|
281
|
+
yi_arr, xi_arr = np.unravel_index(valid_idx, (ny, nx))
|
|
282
|
+
maps['tau_center_1'][yi_arr[good], xi_arr[good]] = tau_amp_ns[good]
|
|
283
|
+
maps['width_1'][yi_arr[good], xi_arr[good]] = w_v[good] * 1e9
|
|
284
|
+
maps['alpha_1'][yi_arr[good], xi_arr[good]] = amp_v[good]
|
|
285
|
+
maps['frac_1'][yi_arr[good], xi_arr[good]] = 1.0
|
|
286
|
+
maps['tau_mean_amp'][yi_arr[good], xi_arr[good]] = tau_amp_ns[good]
|
|
287
|
+
maps['tau_mean_int'][yi_arr[good], xi_arr[good]] = tau_int_ns[good]
|
|
288
|
+
maps['chi2_r'][yi_arr[good], xi_arr[good]] = chi2_v[good]
|
|
289
|
+
maps['calibrated_chi2_r'][yi_arr[good], xi_arr[good]] = chi2_cal_v[good]
|
|
290
|
+
maps['tvb_scale'][yi_arr[good], xi_arr[good]] = tvb_v[good]
|
|
291
|
+
return maps
|
|
292
|
+
bg_flat = self._estimate_bg_batch(flat, valid_mask)
|
|
293
|
+
dc_flat = np.maximum(flat_fit - bg_flat[:, None], 0.0)
|
|
294
|
+
dc_valid = dc_flat[valid_idx]
|
|
295
|
+
basis_mx = mx.array(basis)
|
|
296
|
+
bb_mx = mx.array(bb_grid)
|
|
297
|
+
dc_mx = mx.array(dc_valid)
|
|
298
|
+
bd_mx = dc_mx @ basis_mx.T
|
|
299
|
+
dc_sq_mx = (dc_mx ** 2).sum(axis=1)
|
|
300
|
+
costs_mx = dc_sq_mx[:, None] - mx.maximum(bd_mx, 0.0) ** 2 / bb_mx[None, :]
|
|
301
|
+
best_g_mx = costs_mx.argmin(axis=1)
|
|
302
|
+
mx.eval(best_g_mx, bd_mx)
|
|
303
|
+
best_g = np.array(best_g_mx)
|
|
304
|
+
bd_np = np.array(bd_mx)
|
|
305
|
+
tau_v = param_pairs[best_g, 0]
|
|
306
|
+
w_v = param_pairs[best_g, 1]
|
|
307
|
+
amp_v = np.maximum(bd_np[np.arange(len(valid_idx)), best_g] / bb_grid[best_g].astype(np.float64), 0.0)
|
|
308
|
+
good = amp_v > 0
|
|
309
|
+
tau_amp_ns = tau_v * 1e9
|
|
310
|
+
tau_int_ns = (tau_v + w_v ** 2 / np.maximum(tau_v, 1e-15)) * 1e9
|
|
311
|
+
basis_best = basis[best_g].astype(np.float64)
|
|
312
|
+
model_v = amp_v[:, None] * basis_best + bg_flat[valid_idx, None]
|
|
313
|
+
resid_v = flat_fit[valid_idx].astype(np.float64) - model_v
|
|
314
|
+
chi2_v = (resid_v ** 2 / np.maximum(model_v, 1.0)).sum(axis=1) / distribution_dof(n_fit, 1, False)
|
|
315
|
+
chi2_cal_v = calibrated_chi2(flat_fit[valid_idx], model_v, axis=1)
|
|
316
|
+
yi_arr, xi_arr = np.unravel_index(valid_idx, (ny, nx))
|
|
317
|
+
maps['tau_center_1'][yi_arr[good], xi_arr[good]] = tau_amp_ns[good]
|
|
318
|
+
maps['width_1'][yi_arr[good], xi_arr[good]] = w_v[good] * 1e9
|
|
319
|
+
maps['alpha_1'][yi_arr[good], xi_arr[good]] = amp_v[good]
|
|
320
|
+
maps['frac_1'][yi_arr[good], xi_arr[good]] = 1.0
|
|
321
|
+
maps['tau_mean_amp'][yi_arr[good], xi_arr[good]] = tau_amp_ns[good]
|
|
322
|
+
maps['tau_mean_int'][yi_arr[good], xi_arr[good]] = tau_int_ns[good]
|
|
323
|
+
maps['chi2_r'][yi_arr[good], xi_arr[good]] = chi2_v[good]
|
|
324
|
+
maps['calibrated_chi2_r'][yi_arr[good], xi_arr[good]] = chi2_cal_v[good]
|
|
325
|
+
return maps
|
|
326
|
+
|
|
327
|
+
|
|
328
|
+
def batch_free_tau_fit(
|
|
329
|
+
self,
|
|
330
|
+
stack,
|
|
331
|
+
irf_array,
|
|
332
|
+
tcspc_res,
|
|
333
|
+
taus_init,
|
|
334
|
+
tau_min_s,
|
|
335
|
+
tau_max_s,
|
|
336
|
+
n_exp,
|
|
337
|
+
min_photons,
|
|
338
|
+
correct_pileup,
|
|
339
|
+
n_sync_px,
|
|
340
|
+
n_steps=50,
|
|
341
|
+
lr=None,
|
|
342
|
+
tvb_profile=None,
|
|
343
|
+
fit_tvb=False,
|
|
344
|
+
fit_idx=None,
|
|
345
|
+
):
|
|
346
|
+
mx = self._mx
|
|
347
|
+
ny, nx, n_bins = stack.shape
|
|
348
|
+
taus_ns_init = taus_init * 1e9
|
|
349
|
+
flat = stack.reshape(ny * nx, n_bins).astype(np.float32)
|
|
350
|
+
intensity_flat = flat.sum(axis=1)
|
|
351
|
+
valid_mask = intensity_flat >= min_photons
|
|
352
|
+
valid_idx = np.where(valid_mask)[0]
|
|
353
|
+
maps = self._init_maps(
|
|
354
|
+
ny, nx, n_exp,
|
|
355
|
+
intensity=stack.sum(axis=2),
|
|
356
|
+
taus_fixed_ns=taus_ns_init,
|
|
357
|
+
free_tau=True,
|
|
358
|
+
)
|
|
359
|
+
if valid_idx.size == 0:
|
|
360
|
+
return maps
|
|
361
|
+
bg_flat = self._estimate_bg_batch(flat, valid_mask)
|
|
362
|
+
dc_flat = np.maximum(flat - bg_flat[:, None], 0.0)
|
|
363
|
+
if correct_pileup and n_sync_px > 0:
|
|
364
|
+
for idx in valid_idx:
|
|
365
|
+
dc_flat[idx] = coates_pileup_correction(dc_flat[idx], n_sync_px)
|
|
366
|
+
raw_valid = flat[valid_idx].astype(np.float32)
|
|
367
|
+
bg_valid = bg_flat[valid_idx].astype(np.float32)
|
|
368
|
+
B = len(valid_idx)
|
|
369
|
+
taus_out, amps_out, chi2r_out, chi2c_out, _, valid_b, tvb_out = self._scipy_parallel_free_tau_fit(
|
|
370
|
+
raw_valid, bg_valid, irf_array, tcspc_res,
|
|
371
|
+
taus_init, tau_min_s, tau_max_s, n_exp, n_bins,
|
|
372
|
+
tvb_profile=tvb_profile, fit_tvb=fit_tvb, fit_idx=fit_idx,
|
|
373
|
+
)
|
|
374
|
+
self._scatter_free_tau(
|
|
375
|
+
maps, valid_idx=valid_idx[valid_b],
|
|
376
|
+
taus_s=taus_out[valid_b], amps=amps_out[valid_b],
|
|
377
|
+
chi2_r=chi2r_out[valid_b], calibrated_values=chi2c_out[valid_b],
|
|
378
|
+
ny=ny, nx=nx, n_exp=n_exp,
|
|
379
|
+
tvb=tvb_out[valid_b] if fit_tvb else None,
|
|
380
|
+
)
|
|
381
|
+
return maps
|
flimkit/GPU/mps.py
ADDED