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
flimkit/GPU/_base.py ADDED
@@ -0,0 +1,391 @@
1
+ import multiprocessing
2
+ import numpy as np
3
+ from concurrent.futures import ThreadPoolExecutor
4
+ from scipy.optimize import least_squares
5
+ from ..FLIM.fit_tools import calibrated_chi2, calibrated_from_terms, chi2_terms
6
+ from ..FLIM.models import reconvolution_model
7
+
8
+ def fit_window(fit_idx, n_bins):
9
+ if fit_idx is None:
10
+ return None
11
+ raw = np.asarray(fit_idx)
12
+ if raw.ndim != 1 or raw.size == 0:
13
+ raise ValueError('fit_idx must be a non-empty one-dimensional array')
14
+ if not np.issubdtype(raw.dtype, np.integer):
15
+ raise ValueError('fit_idx must contain integer bin indices')
16
+ idx = raw.astype(int, copy=False)
17
+ if np.any(idx < 0) or np.any(idx >= n_bins):
18
+ raise ValueError(f'fit_idx entries must be between 0 and {n_bins - 1}')
19
+ if np.unique(idx).size != idx.size:
20
+ raise ValueError('fit_idx must not contain duplicate bins')
21
+ if np.array_equal(idx, np.arange(n_bins)):
22
+ return None
23
+ return idx
24
+
25
+
26
+ GPU_BLOCK_BYTES = 256 * 1024 * 1024
27
+
28
+ def gpu_block_bytes():
29
+ import os
30
+ override = os.environ.get('FLIMKIT_GPU_BLOCK_BYTES')
31
+ if override:
32
+ try:
33
+ return max(1, int(override))
34
+ except ValueError:
35
+ pass
36
+ return GPU_BLOCK_BYTES
37
+
38
+ def pixel_blocks(n_pixels, bytes_per_pixel, budget=None):
39
+ budget = gpu_block_bytes() if budget is None else budget
40
+ step = max(1, int(budget // max(int(bytes_per_pixel), 1)))
41
+ for start in range(0, n_pixels, step):
42
+ yield start, min(start + step, n_pixels)
43
+
44
+
45
+ class GPUBackend:
46
+
47
+ def batch_fixed_tau(
48
+ self,
49
+ stack,
50
+ A,
51
+ taus_fixed,
52
+ min_photons,
53
+ correct_pileup,
54
+ n_sync_px,
55
+ progress_callback,
56
+ fit_idx=None,
57
+ ):
58
+ raise NotImplementedError
59
+
60
+ def batch_grid_scan_1exp(
61
+ self,
62
+ stack,
63
+ basis_grid,
64
+ bb_grid,
65
+ tau_grid,
66
+ min_photons,
67
+ correct_pileup,
68
+ n_sync_px,
69
+ progress_callback,
70
+ fit_idx=None,
71
+ ):
72
+ raise NotImplementedError
73
+
74
+ def batch_free_tau_fit(
75
+ self,
76
+ stack,
77
+ irf_array,
78
+ tcspc_res,
79
+ taus_init,
80
+ tau_min_s,
81
+ tau_max_s,
82
+ n_exp,
83
+ min_photons,
84
+ correct_pileup,
85
+ n_sync_px,
86
+ n_steps,
87
+ lr,
88
+ fit_idx=None,
89
+ ):
90
+ raise NotImplementedError
91
+
92
+ def batch_dist_scan_unimodal(
93
+ self,
94
+ stack,
95
+ basis,
96
+ bb_grid,
97
+ param_pairs,
98
+ irf_fixed,
99
+ tcspc_res,
100
+ n_bins,
101
+ dist_type,
102
+ min_photons,
103
+ progress_callback,
104
+ tvb_profile=None,
105
+ fit_tvb=False,
106
+ fit_idx=None,
107
+ ):
108
+ raise NotImplementedError
109
+
110
+ class _BackendMixin:
111
+
112
+ @staticmethod
113
+ def _estimate_bg_batch(flat, valid_mask, pre_gap=5):
114
+ n_pix, n_bins = flat.shape
115
+ bg = np.zeros(n_pix, dtype=np.float32)
116
+ wanted = np.asarray(valid_mask, dtype=bool)
117
+ if not wanted.any():
118
+ return bg
119
+ ends = np.maximum(flat.argmax(axis=1) - pre_gap, 0)
120
+ for end in np.unique(ends[wanted]):
121
+ rows = np.flatnonzero(wanted & (ends == end))
122
+ if end >= 5:
123
+ region = flat[rows, :end]
124
+ else:
125
+ region = flat[rows, -30:]
126
+ bg[rows] = np.maximum(np.median(region, axis=1), 0.0)
127
+ return bg
128
+ @staticmethod
129
+ def _init_maps(ny, nx, n_exp, intensity, taus_fixed_ns, free_tau):
130
+ maps = dict(
131
+ intensity = intensity,
132
+ tau_mean_int = np.full((ny, nx), np.nan),
133
+ tau_mean_amp = np.full((ny, nx), np.nan),
134
+ chi2_r = np.full((ny, nx), np.nan),
135
+ calibrated_chi2_r = np.full((ny, nx), np.nan),
136
+ )
137
+ for i in range(n_exp):
138
+ maps[f"alpha_{i+1}"] = np.full((ny, nx), np.nan)
139
+ maps[f"frac_{i+1}"] = np.full((ny, nx), np.nan)
140
+ maps[f"tau_{i+1}"] = (np.full((ny, nx), np.nan) if (n_exp == 1 or free_tau)
141
+ else np.full((ny, nx), taus_fixed_ns[i]))
142
+ maps[f"a{i+1}"] = maps[f"alpha_{i+1}"]
143
+ return maps
144
+
145
+ @staticmethod
146
+ def _scatter_fixed_tau(
147
+ maps,
148
+ valid_idx,
149
+ amps,
150
+ bg,
151
+ decay_valid,
152
+ A,
153
+ taus_ns,
154
+ ny, nx,
155
+ tvb=None,
156
+ tvb_profile=None,
157
+ ):
158
+ n_exp = A.shape[1]
159
+ n_bins = A.shape[0]
160
+ amp_sum = amps.sum(axis=1)
161
+ good = amp_sum > 0
162
+ if not good.any():
163
+ return
164
+ fracs = np.where(good[:, None], amps / np.maximum(amp_sum[:, None], 1e-30), 0.0)
165
+ taus_ns2 = taus_ns ** 2
166
+ tau_amp = (fracs * taus_ns[None, :]).sum(axis=1)
167
+ denom = (amps * taus_ns[None, :]).sum(axis=1)
168
+ tau_int = np.where(denom > 0,(amps * taus_ns2[None, :]).sum(axis=1) / np.maximum(denom, 1e-30), np.nan)
169
+ model = amps @ A.T + bg[:, None]
170
+ if tvb is not None and tvb_profile is not None:
171
+ model = model + tvb[:, None] * tvb_profile[None, :]
172
+ numerator, expected, row_ok = chi2_terms(decay_valid, model, axis=1)
173
+ chi2 = numerator
174
+ calibrated = calibrated_from_terms(numerator, expected, row_ok)
175
+ dof = max(n_bins - n_exp, 1)
176
+
177
+ yi_arr, xi_arr = np.unravel_index(valid_idx, (ny, nx))
178
+
179
+ maps['tau_mean_amp'][yi_arr[good], xi_arr[good]] = tau_amp[good]
180
+ maps['tau_mean_int'][yi_arr[good], xi_arr[good]] = tau_int[good]
181
+ maps['chi2_r'][yi_arr[good], xi_arr[good]] = chi2[good] / dof
182
+ maps['calibrated_chi2_r'][yi_arr[good], xi_arr[good]] = calibrated[good]
183
+ for i in range(n_exp):
184
+ maps[f"alpha_{i+1}"][yi_arr[good], xi_arr[good]] = amps[good, i]
185
+ maps[f"frac_{i+1}"][yi_arr[good], xi_arr[good]] = fracs[good, i]
186
+ if tvb is not None:
187
+ maps.setdefault('tvb_scale', np.full((ny, nx), np.nan))
188
+ maps['tvb_scale'][yi_arr[good], xi_arr[good]] = tvb[good]
189
+
190
+ @staticmethod
191
+ def _tvb_grid_prep(basis_grid, tvb_profile, n_bins):
192
+ U = np.column_stack([np.asarray(tvb_profile, dtype=np.float64), np.ones(n_bins)])
193
+ U_pinv = np.linalg.pinv(U)
194
+ basis_perp = basis_grid - (basis_grid @ U_pinv.T) @ U.T
195
+ bb_perp = np.maximum((basis_perp ** 2).sum(axis=1), 1e-20)
196
+ return U, U_pinv, basis_perp.astype(np.float32), bb_perp.astype(np.float32)
197
+
198
+ @staticmethod
199
+ def _tvb_project_data(data, U, U_pinv):
200
+ return data - (data @ U_pinv.T) @ U.T
201
+
202
+ @staticmethod
203
+ def _scatter_1exp(
204
+ maps,
205
+ valid_idx,
206
+ tau_v,
207
+ amp_v,
208
+ bg_v,
209
+ decay_valid,
210
+ basis_best,
211
+ ny, nx,
212
+ n_bins,
213
+ tvb=None,
214
+ tvb_profile=None,
215
+ ):
216
+ good = amp_v > 0
217
+ tau_ns = tau_v * 1e9
218
+ model = amp_v[:, None] * basis_best + bg_v[:, None]
219
+ if tvb is not None and tvb_profile is not None:
220
+ model = model + tvb[:, None] * tvb_profile[None, :]
221
+ numerator, expected, row_ok = chi2_terms(decay_valid, model, axis=1)
222
+ chi2 = numerator / max(n_bins - 2, 1)
223
+ calibrated = calibrated_from_terms(numerator, expected, row_ok)
224
+
225
+ yi_arr, xi_arr = np.unravel_index(valid_idx, (ny, nx))
226
+ maps['tau_1'][yi_arr[good], xi_arr[good]] = tau_ns[good]
227
+ maps['tau_mean_amp'][yi_arr[good], xi_arr[good]] = tau_ns[good]
228
+ maps['tau_mean_int'][yi_arr[good], xi_arr[good]] = tau_ns[good]
229
+ maps['alpha_1'][yi_arr[good], xi_arr[good]] = amp_v[good]
230
+ maps['frac_1'][yi_arr[good], xi_arr[good]] = 1.0
231
+ maps['chi2_r'][yi_arr[good], xi_arr[good]] = chi2[good]
232
+ maps['calibrated_chi2_r'][yi_arr[good], xi_arr[good]] = calibrated[good]
233
+ if tvb is not None:
234
+ maps.setdefault('tvb_scale', np.full((ny, nx), np.nan))
235
+ maps['tvb_scale'][yi_arr[good], xi_arr[good]] = tvb[good]
236
+
237
+ @staticmethod
238
+ def _scatter_free_tau(
239
+ maps,
240
+ valid_idx,
241
+ taus_s,
242
+ amps,
243
+ chi2_r,
244
+ calibrated_values,
245
+ ny, nx,
246
+ n_exp,
247
+ tvb=None,
248
+ ):
249
+ amp_sum = amps.sum(axis=1)
250
+ good = amp_sum > 0
251
+ if not good.any():
252
+ return
253
+
254
+ fracs = np.where(good[:, None], amps / np.maximum(amp_sum[:, None], 1e-30), 0.0)
255
+ taus_ns = taus_s * 1e9
256
+ tau_amp = (fracs * taus_ns).sum(axis=1)
257
+ denom = (amps * taus_ns).sum(axis=1)
258
+ tau_int = np.where(
259
+ denom > 0,
260
+ (amps * taus_ns ** 2).sum(axis=1) / np.maximum(denom, 1e-30),
261
+ np.nan,
262
+ )
263
+
264
+ yi_arr, xi_arr = np.unravel_index(valid_idx, (ny, nx))
265
+ maps['tau_mean_amp'][yi_arr[good], xi_arr[good]] = tau_amp[good]
266
+ maps['tau_mean_int'][yi_arr[good], xi_arr[good]] = tau_int[good]
267
+ maps['chi2_r'][yi_arr[good], xi_arr[good]] = chi2_r[good]
268
+ maps['calibrated_chi2_r'][yi_arr[good], xi_arr[good]] = calibrated_values[good]
269
+ for i in range(n_exp):
270
+ maps[f"tau_{i+1}"][yi_arr[good], xi_arr[good]] = taus_ns[good, i]
271
+ maps[f"alpha_{i+1}"][yi_arr[good], xi_arr[good]] = amps[good, i]
272
+ maps[f"frac_{i+1}"][yi_arr[good], xi_arr[good]] = fracs[good, i]
273
+ if tvb is not None:
274
+ maps.setdefault('tvb_scale', np.full((ny, nx), np.nan))
275
+ maps['tvb_scale'][yi_arr[good], xi_arr[good]] = tvb[good]
276
+
277
+ @staticmethod
278
+ def _scipy_parallel_free_tau_fit(
279
+ raw_valid,
280
+ bg_valid,
281
+ irf_array,
282
+ tcspc_res,
283
+ taus_init,
284
+ tau_min_s,
285
+ tau_max_s,
286
+ n_exp,
287
+ n_bins,
288
+ tvb_profile=None,
289
+ fit_tvb=False,
290
+ fit_idx=None,
291
+ ):
292
+ B = raw_valid.shape[0]
293
+ win = fit_window(fit_idx, n_bins)
294
+ n_fit = n_bins if win is None else len(win)
295
+
296
+ amp0 = float(raw_valid.max()) / n_exp
297
+ # Use the same bounds as the CPU free-tau path in fit_per_pixel
298
+ amp_hi = float(raw_valid.max()) * 10.0
299
+ lo_px = np.array([float(tau_min_s)] * n_exp + [0.0] * n_exp)
300
+ hi_px = np.array([float(tau_max_s)] * n_exp + [amp_hi] * n_exp)
301
+ if fit_tvb:
302
+ tvb_hi = float(raw_valid.sum(axis=1).max())
303
+ lo_px = np.concatenate([lo_px, [0.0]])
304
+ hi_px = np.concatenate([hi_px, [tvb_hi]])
305
+
306
+ def _fit_pixel(b):
307
+ decay_b = raw_valid[b].astype(np.float64)
308
+ bg_b = float(bg_valid[b])
309
+ wt = np.sqrt(np.maximum(decay_b, 1.0))
310
+ p0 = np.concatenate([taus_init,
311
+ np.full(n_exp, amp0)])
312
+ if fit_tvb:
313
+ p0 = np.concatenate([p0, [bg_b * n_bins]])
314
+
315
+ def _resid(p):
316
+ if fit_tvb:
317
+ full_p = np.concatenate([p[:n_exp], p[n_exp:2 * n_exp], [0.0], [p[2 * n_exp]]])
318
+ model = reconvolution_model(
319
+ full_p, tcspc_res, n_bins, irf_array,
320
+ n_exp, 0.0, False, False, False,
321
+ tvb_profile=tvb_profile, fit_tvb=True)
322
+ else:
323
+ full_p = np.concatenate([p[:n_exp], p[n_exp:], [0.0]])
324
+ model = reconvolution_model(
325
+ full_p, tcspc_res, n_bins, irf_array,
326
+ n_exp, bg_b, False, False, False)
327
+ if win is None:
328
+ return (model - decay_b) / wt
329
+ return (model[win] - decay_b[win]) / wt[win]
330
+
331
+ try:
332
+ res = least_squares(_resid, p0, bounds=(lo_px, hi_px),
333
+ method='trf', max_nfev=500,
334
+ ftol=1e-8, xtol=1e-8, gtol=1e-8)
335
+ return res.x
336
+ except Exception:
337
+ return None
338
+
339
+ n_workers = min(B, max(1, multiprocessing.cpu_count()))
340
+ with ThreadPoolExecutor(max_workers=n_workers) as pool:
341
+ solutions = list(pool.map(_fit_pixel, range(B)))
342
+
343
+ taus_out = np.zeros((B, n_exp), dtype=np.float32)
344
+ amps_out = np.zeros((B, n_exp), dtype=np.float32)
345
+ chi2r_out = np.full(B, np.nan, dtype=np.float64)
346
+ chi2c_out = np.full(B, np.nan, dtype=np.float64)
347
+ model_out = np.zeros((B, n_bins), dtype=np.float32)
348
+ tvb_out = np.zeros(B, dtype=np.float32)
349
+ valid_b = np.zeros(B, dtype=bool)
350
+
351
+ for b, p_sol in enumerate(solutions):
352
+ if p_sol is None:
353
+ continue
354
+ taus_b = p_sol[:n_exp]; amps_b = p_sol[n_exp:2 * n_exp]
355
+ if amps_b.sum() <= 0:
356
+ continue
357
+ tvb_b = float(p_sol[2 * n_exp]) if fit_tvb else 0.0
358
+ # Sort ascending for output (matches CPU convention)
359
+ order = np.argsort(taus_b)
360
+ taus_b = taus_b[order]; amps_b = amps_b[order]
361
+ bg_b = float(bg_valid[b])
362
+ if fit_tvb:
363
+ full_p = np.concatenate([taus_b, amps_b, [0.0], [tvb_b]])
364
+ model_b = reconvolution_model(
365
+ full_p, tcspc_res, n_bins, irf_array,
366
+ n_exp, 0.0, False, False, False,
367
+ tvb_profile=tvb_profile, fit_tvb=True)
368
+ else:
369
+ full_p = np.concatenate([taus_b, amps_b, [0.0]])
370
+ model_b = reconvolution_model(
371
+ full_p, tcspc_res, n_bins, irf_array,
372
+ n_exp, bg_b, False, False, False)
373
+ decay_fit = raw_valid[b].astype(np.float64)
374
+ if win is None:
375
+ model_fit = model_b
376
+ else:
377
+ decay_fit = decay_fit[win]
378
+ model_fit = model_b[win]
379
+ resid_b = decay_fit - model_fit
380
+ chi2_b = (resid_b ** 2 / np.maximum(model_fit, 1.0)).sum()
381
+ dof = max(n_fit - 2 * n_exp, 1)
382
+
383
+ taus_out[b] = taus_b.astype(np.float32)
384
+ amps_out[b] = amps_b.astype(np.float32)
385
+ model_out[b] = model_b.astype(np.float32)
386
+ tvb_out[b] = tvb_b
387
+ chi2r_out[b] = chi2_b / dof
388
+ chi2c_out[b] = calibrated_chi2(decay_fit, model_fit)
389
+ valid_b[b] = True
390
+
391
+ return taus_out, amps_out, chi2r_out, chi2c_out, model_out, valid_b, tvb_out
flimkit/GPU/cuda.py ADDED
@@ -0,0 +1,10 @@
1
+ from flimkit.GPU.torch_backend import TorchBackend
2
+
3
+
4
+ class CUDABackend(TorchBackend):
5
+
6
+ def __init__(self):
7
+ super().__init__(device='cuda')
8
+
9
+ def __repr__(self):
10
+ return "CUDABackend(device='cuda')"