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.
Files changed (104) hide show
  1. flimkit/FLIM/__init__.py +0 -0
  2. flimkit/FLIM/assemble.py +254 -0
  3. flimkit/FLIM/batch.py +681 -0
  4. flimkit/FLIM/bg_tools.py +51 -0
  5. flimkit/FLIM/fit_tools.py +244 -0
  6. flimkit/FLIM/fitters.py +1471 -0
  7. flimkit/FLIM/irf_tools.py +617 -0
  8. flimkit/FLIM/models.py +391 -0
  9. flimkit/GPU/__init__.py +85 -0
  10. flimkit/GPU/_base.py +391 -0
  11. flimkit/GPU/cuda.py +10 -0
  12. flimkit/GPU/mlx_backend.py +381 -0
  13. flimkit/GPU/mps.py +10 -0
  14. flimkit/GPU/rocm.py +10 -0
  15. flimkit/GPU/torch_backend.py +385 -0
  16. flimkit/UI/app_state.py +10 -0
  17. flimkit/UI/controller.py +139 -0
  18. flimkit/UI/expert_settings.py +248 -0
  19. flimkit/UI/fit_help.py +206 -0
  20. flimkit/UI/fov_preview.py +1085 -0
  21. flimkit/UI/gui.py +3919 -0
  22. flimkit/UI/icon.icns +0 -0
  23. flimkit/UI/icon.ico +0 -0
  24. flimkit/UI/icon.png +0 -0
  25. flimkit/UI/irf_widget.py +103 -0
  26. flimkit/UI/mode_controller.py +118 -0
  27. flimkit/UI/modes/__init__.py +0 -0
  28. flimkit/UI/modes/base.py +3 -0
  29. flimkit/UI/modes/batch_mode.py +312 -0
  30. flimkit/UI/modes/fov_mode.py +164 -0
  31. flimkit/UI/modes/irf_mode.py +80 -0
  32. flimkit/UI/modes/phasor_mode.py +131 -0
  33. flimkit/UI/modes/stitch_mode.py +254 -0
  34. flimkit/UI/phasor_panel.py +1087 -0
  35. flimkit/UI/progress_window.py +113 -0
  36. flimkit/UI/project_panel.py +262 -0
  37. flimkit/UI/results_panel.py +332 -0
  38. flimkit/UI/roi_tools.py +794 -0
  39. flimkit/UI/utils.py +217 -0
  40. flimkit/__init__.py +0 -0
  41. flimkit/_version.py +41 -0
  42. flimkit/cli.py +120 -0
  43. flimkit/configs.py +148 -0
  44. flimkit/dialogs.py +46 -0
  45. flimkit/formats/BH/__init__.py +0 -0
  46. flimkit/formats/BH/reader.py +296 -0
  47. flimkit/formats/BH/writer.py +86 -0
  48. flimkit/formats/ISS/__init__.py +0 -0
  49. flimkit/formats/ISS/fdflim.py +86 -0
  50. flimkit/formats/ISS/image.py +114 -0
  51. flimkit/formats/ISS/reader.py +223 -0
  52. flimkit/formats/PS/__init__.py +0 -0
  53. flimkit/formats/PS/reader.py +202 -0
  54. flimkit/formats/PTU/__init__.py +0 -0
  55. flimkit/formats/PTU/decode.py +27 -0
  56. flimkit/formats/PTU/phu.py +85 -0
  57. flimkit/formats/PTU/reader.py +235 -0
  58. flimkit/formats/PTU/series.py +258 -0
  59. flimkit/formats/PTU/stitch.py +1182 -0
  60. flimkit/formats/PTU/tools.py +94 -0
  61. flimkit/formats/__init__.py +2 -0
  62. flimkit/formats/flim_file.py +232 -0
  63. flimkit/formats/phasor.py +132 -0
  64. flimkit/formats/signal.py +170 -0
  65. flimkit/image/tools.py +124 -0
  66. flimkit/interactive.py +1857 -0
  67. flimkit/mpl_backend.py +22 -0
  68. flimkit/phasor/__init__.py +40 -0
  69. flimkit/phasor/filters.py +127 -0
  70. flimkit/phasor/fret.py +654 -0
  71. flimkit/phasor/interactive.py +556 -0
  72. flimkit/phasor/peaks.py +186 -0
  73. flimkit/phasor/signal.py +90 -0
  74. flimkit/phasor_launcher.py +314 -0
  75. flimkit/plugins/__init__.py +137 -0
  76. flimkit/plugins/bindings.py +116 -0
  77. flimkit/plugins/builtin/__init__.py +3 -0
  78. flimkit/plugins/builtin/core_tools.py +28 -0
  79. flimkit/plugins/loader.py +371 -0
  80. flimkit/plugins/registry.py +406 -0
  81. flimkit/project.py +197 -0
  82. flimkit/synth.py +145 -0
  83. flimkit/utils/__init__.py +0 -0
  84. flimkit/utils/batch_fit.py +301 -0
  85. flimkit/utils/config_manager.py +119 -0
  86. flimkit/utils/config_snapshot.py +30 -0
  87. flimkit/utils/crash_handler.py +183 -0
  88. flimkit/utils/display.py +197 -0
  89. flimkit/utils/enhanced_outputs.py +345 -0
  90. flimkit/utils/fancy.py +103 -0
  91. flimkit/utils/lifetime_image.py +243 -0
  92. flimkit/utils/misc.py +111 -0
  93. flimkit/utils/plotting.py +190 -0
  94. flimkit/utils/roi.py +370 -0
  95. flimkit/utils/session.py +51 -0
  96. flimkit/utils/update_check.py +198 -0
  97. flimkit/utils/xlsx_tools.py +97 -0
  98. flimkit/utils/xml_utils.py +219 -0
  99. flimkit-0.12.0.dist-info/METADATA +356 -0
  100. flimkit-0.12.0.dist-info/RECORD +104 -0
  101. flimkit-0.12.0.dist-info/WHEEL +5 -0
  102. flimkit-0.12.0.dist-info/entry_points.txt +2 -0
  103. flimkit-0.12.0.dist-info/licenses/LICENSE.md +11 -0
  104. 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
@@ -0,0 +1,10 @@
1
+ from flimkit.GPU.torch_backend import TorchBackend
2
+
3
+
4
+ class MPSBackend(TorchBackend):
5
+
6
+ def __init__(self):
7
+ super().__init__(device='mps')
8
+
9
+ def __repr__(self):
10
+ return "MPSBackend(device='mps')"
flimkit/GPU/rocm.py ADDED
@@ -0,0 +1,10 @@
1
+ from flimkit.GPU.torch_backend import TorchBackend
2
+
3
+
4
+ class ROCmBackend(TorchBackend):
5
+
6
+ def __init__(self):
7
+ super().__init__(device='cuda')
8
+
9
+ def __repr__(self):
10
+ return "ROCmBackend(device='cuda/rocm')"