"""
Fit FLI image cubes with Torch-based GPU optimization and optional CRLB estimates.
This module belongs to :mod:`pyfli.solver` and is part of PyFLI least-squares, maximum-
likelihood, CPU, GPU, binned, and global FLI fitting routines. Public API includes
classes :class:`FLIGPUProcessor`.
"""
import math
import os
import time
from typing import Any
import h5py
import numpy as np
import torch
from tqdm import tqdm
from pyfli import logging
from .shared_metrics import pearson_chi_square, reduced_poisson_deviance
[docs]
class FLIGPUProcessor:
"""
Fit FLI image cubes with Torch on GPU or CPU fallback. It vectorizes parameter
transforms, model evaluation, optimization, CRLB estimation, reconstruction, and
result saving.
Parameters
----------
freq : float
Acquisition frequency information used to derive timing constants.
fitter_class : Any | None
Fitter class instantiated by the processor.
device : Any | None
Execution device, such as a Torch device or device string.
"""
def __init__(
self, freq: float, fitter_class: Any | None = None, device: Any | None = None
) -> None:
self.device = (
device if device else ("cuda" if torch.cuda.is_available() else "cpu")
)
self.freq = freq
self.fitter_class = fitter_class
self.T_acq = 1000.0 / freq[1]
self.T_laser = 1000.0 / freq[0]
logging.info(f"Using Device: {self.device}")
def _transform_params(self, raw_p: np.ndarray, model_type: str) -> Any:
"""
Run the transform params routine.
Parameters
----------
raw_p : np.ndarray
Unconstrained optimizer parameters before physical transformation.
model_type : str
FLI model family, such as mono- or bi-exponential.
Returns
-------
Any
Object produced by transform params.
"""
shift_bound = self.T_acq / 4.0
S = torch.exp(raw_p[:, 0:1])
b = torch.exp(raw_p[:, -2:-1])
h_shift = torch.tanh(raw_p[:, -1:]) * shift_bound
if model_type == "bi-exponential":
a1 = torch.sigmoid(raw_p[:, 1:2])
t1 = torch.exp(raw_p[:, 2:3])
t2 = t1 + torch.exp(raw_p[:, 3:4])
return torch.cat([S, a1, t1, t2, b, h_shift], dim=1)
else:
tau = torch.exp(raw_p[:, 1:2])
return torch.cat([S, tau, b, h_shift], dim=1)
def _model_kernel(
self, params: Any, t: np.ndarray, irf: np.ndarray, model_type: str
) -> Any:
"""
Batched forward model, identical to :func:`pyfli.solver.model_numpy`: the
gate-integrated decay with onset ``h_shift`` (evaluated on extra negative-lag
gates so an earlier onset shifts the curve instead of truncating it),
convolved with the normalized IRF, plus ``v_shift``. ``S`` is the total
photon count of the decay.
Parameters
----------
params : Any
Model, detector, or plotting parameters used by the routine.
t : np.ndarray
Time axis or acquisition period used by the calculation.
irf : np.ndarray
Instrument response function aligned with the decay signal.
model_type : str
FLI model family, such as mono- or bi-exponential.
Returns
-------
Any
Object produced by model kernel.
"""
T = t.shape[-1]
dt = self.T_acq / T
n_neg = math.ceil((self.T_acq / 4.0) / dt) + 1
lags = torch.arange(-n_neg, T, device=t.device, dtype=params.dtype) * dt
a = lags[None, :]
b_edge = a + dt
h_shift = params[:, -1:]
a_eff = torch.maximum(a, h_shift) - h_shift
b_eff = torch.maximum(b_edge, h_shift) - h_shift
def gate_fraction(tau: Any) -> Any:
return torch.exp(-a_eff / tau) - torch.exp(-b_eff / tau)
if model_type == "mono-exponential":
S, tau, b = params[:, 0:1], params[:, 1:2], params[:, 2:3]
decay = S * gate_fraction(tau)
else:
S, a1, t1, t2, b = (
params[:, 0:1],
params[:, 1:2],
params[:, 2:3],
params[:, 3:4],
params[:, 4:5],
)
decay = S * (a1 * gate_fraction(t1) + (1.0 - a1) * gate_fraction(t2))
irf_norm = irf / irf.sum(dim=1, keepdim=True).clamp(min=1e-9)
n_fft = 2 * (T + n_neg)
decay_fft = torch.fft.rfft(decay, n=n_fft)
irf_fft = torch.fft.rfft(irf_norm, n=n_fft)
convolved = torch.fft.irfft(decay_fft * irf_fft, n=n_fft)[
..., n_neg : n_neg + T
]
return convolved + b
def _compute_crlb_errors(
self, p_phys: np.ndarray, t: np.ndarray, irf: np.ndarray, model_type: str
) -> Any:
"""
Compute crlb errors.
Parameters
----------
p_phys : np.ndarray
Physical parameter tensor after constrained transformation.
t : np.ndarray
Time axis or acquisition period used by the calculation.
irf : np.ndarray
Instrument response function aligned with the decay signal.
model_type : str
FLI model family, such as mono- or bi-exponential.
Returns
-------
Any
Object produced by compute CRLB errors.
"""
p_phys = p_phys.detach().clone().requires_grad_(True)
def model_func(p: Any) -> Any:
"""
Run the model func routine.
Parameters
----------
p : Any
Detector parameter object or fitted parameter vector.
Returns
-------
Any
Object produced by model func.
"""
return self._model_kernel(p, t, irf, model_type)
jac = torch.autograd.functional.jacobian(model_func, p_phys, vectorize=True)
jac = torch.diagonal(jac, dim1=0, dim2=2).permute(2, 0, 1)
with torch.no_grad():
pred = model_func(p_phys)
W = 1.0 / torch.clamp(pred, min=1.0)
jt_w = jac.transpose(1, 2) * W.unsqueeze(1)
fim = torch.bmm(jt_w, jac)
trace = (
torch.diagonal(fim, dim1=1, dim2=2)
.sum(dim=1, keepdim=True)
.unsqueeze(-1)
)
eps_mat = (1e-6 * trace / fim.shape[-1]) * torch.eye(
fim.shape[-1], device=self.device
)
try:
cov = torch.inverse(fim + eps_mat)
return torch.sqrt(torch.abs(torch.diagonal(cov, dim1=1, dim2=2)))
except RuntimeError:
return torch.zeros_like(p_phys)
[docs]
def fit_image(
self,
image_cube: np.ndarray,
irf_cube: np.ndarray,
mask: np.ndarray | None = None,
mode: str = "MLE",
model_type: str = "bi-exponential",
max_iter: int = 500,
CRLB: bool = False,
data_name: str = "Torch_Fit",
p0: Any | None = None,
fit_indices: tuple[int, int] | None = None,
weighting: str = "irls",
**kwargs: Any,
) -> Any:
# Normalise mode tag: NLSF/LSE variants → 'NLSF', everything else → 'MLE'
"""
Fit image.
Parameters
----------
image_cube : np.ndarray
Time-resolved decay image cube.
irf_cube : np.ndarray
Instrument response cube aligned with the decay image cube.
mask : np.ndarray | None
Boolean or labeled mask selecting pixels for the operation.
mode : str
Mode selector used by the fitting, loading, or plotting routine.
model_type : str
FLI model family, such as mono- or bi-exponential.
max_iter : int
Maximum number of optimization iterations.
CRLB : bool
If ``True``, compute Cramer-Rao lower-bound uncertainty estimates.
data_name : str
Label assigned to the fitted or processed dataset.
p0 : Any | None
Initial parameter vector supplied to the optimizer.
fit_indices : tuple[int, int] | None
Optional (gate_num_start, gate_num_end) gate range to fit over, e.g. to
focus on the tail of the decay. The forward model is still evaluated over
the full trace (needed for correct IRF convolution); only the loss and fit
statistics are restricted to this gate range. ``None`` fits the full trace.
weighting : str
Residual weights for ``mode="NLSF"`` (ignored for MLE, which uses the
Poisson deviance): ``"irls"`` (default) divides each squared residual by
the current model value held constant in the gradient -- the batched
counterpart of iteratively reweighted least squares, whose converged
solution solves the Poisson likelihood equations; ``"none"`` is
unweighted; ``"neyman"`` divides by the measured counts (former default,
biased towards short lifetimes at low counts). ``mode="NEYMAN"`` implies
``"neyman"``. ``variance_floor`` (kwarg, default 1.0) floors the IRLS
variance.
**kwargs : Any
Additional keyword options forwarded to the underlying implementation.
Returns
-------
Any
Object produced by fit image.
"""
_NLSF_MODES = {"NLSF", "LSE", "WLS", "NEYMAN", "LEAST_SQUARES"}
if mode.upper() == "NEYMAN":
weighting = "neyman"
if weighting not in ("irls", "none", "neyman"):
raise ValueError(
f"weighting must be 'irls', 'none' or 'neyman', got {weighting!r}"
)
mode = "NLSF" if mode.upper() in _NLSF_MODES else "MLE"
variance_floor = kwargs.get("variance_floor", 1.0)
start_time = time.time()
H, W, T = image_cube.shape
t_axis = torch.arange(T, device=self.device) * (self.T_acq / T)
if fit_indices is not None:
gate_start, gate_end = fit_indices
gate_start, gate_end = max(gate_start, 0), min(gate_end, T)
else:
gate_start, gate_end = 0, T
irf_tensor = torch.tensor(
irf_cube, device=self.device, dtype=torch.float32
).reshape(-1, T)
if mask is None:
mask = np.sum(image_cube, axis=2) > 20
valid_idx = np.where(mask.flatten())[0]
if len(valid_idx) == 0:
logging.warning("No valid pixels found.")
return None
flat_data = torch.tensor(
image_cube.reshape(-1, T)[valid_idx],
device=self.device,
dtype=torch.float32,
)
flat_irf = irf_tensor[valid_idx]
if p0 is not None:
p_guess = self._p0_to_tensor(p0, len(valid_idx), model_type)
else:
p_guess = self._get_sophisticated_guess(flat_data, flat_irf, model_type)
raw_p = torch.zeros_like(p_guess)
with torch.no_grad():
raw_p[:, 0] = torch.log(torch.clamp(p_guess[:, 0], min=1e-3)) # log(S)
raw_p[:, -2] = torch.log(
torch.clamp(p_guess[:, -2], min=1e-6)
) # log(v_shift)
raw_p[:, -1] = math.atanh(0.5 * (self.T_acq / T) / (self.T_acq / 4.0))
if model_type == "bi-exponential":
raw_p[:, 1] = torch.logit(torch.clamp(p_guess[:, 1], 0.001, 0.999))
raw_p[:, 2] = torch.log(torch.clamp(p_guess[:, 2], min=1e-3))
raw_p[:, 3] = torch.log(
torch.clamp(p_guess[:, 3] - p_guess[:, 2], min=1e-3)
)
else:
raw_p[:, 1] = torch.log(torch.clamp(p_guess[:, 1], min=0.1))
raw_p.requires_grad_(True)
pixel_health_map = np.ones(H * W, dtype=np.float32)
if mode == "NLSF":
def objective_fn(p_raw: np.ndarray) -> Any:
"""
Run the objective fn routine.
Parameters
----------
p_raw : np.ndarray
Raw unconstrained parameter tensor optimized by the objective.
Returns
-------
Any
Object produced by objective fn.
"""
p_phys = self._transform_params(p_raw, model_type)
pred = self._model_kernel(p_phys, t_axis, flat_irf, model_type)
pred_sel = pred[:, gate_start:gate_end]
data_sel = flat_data[:, gate_start:gate_end]
if weighting == "neyman":
variance = torch.clamp(data_sel, min=1.0)
elif weighting == "none":
variance = torch.ones_like(data_sel)
else:
variance = torch.clamp(pred_sel.detach(), min=variance_floor)
per_px = torch.sum((pred_sel - data_sel) ** 2 / variance, dim=1)
return per_px[torch.isfinite(per_px)].sum()
else:
# Poisson MLE (C-statistic): matches CPU MLEFLIFitter
def objective_fn(p_raw: np.ndarray) -> Any:
"""
Run the objective fn routine.
Parameters
----------
p_raw : np.ndarray
Raw unconstrained parameter tensor optimized by the objective.
Returns
-------
Any
Object produced by objective fn.
"""
p_phys = self._transform_params(p_raw, model_type)
pred = self._model_kernel(p_phys, t_axis, flat_irf, model_type)
pred_sel = pred[:, gate_start:gate_end]
data_sel = flat_data[:, gate_start:gate_end]
pred_safe = torch.clamp(pred_sel, min=1e-9)
per_px = 2.0 * torch.sum(
pred_safe
- data_sel
+ data_sel * torch.log(data_sel.clamp(min=1e-9) / pred_safe),
dim=1,
)
return per_px[torch.isfinite(per_px)].sum()
# Adam operates per-parameter independently — unlike LBFGS it does not maintain
# a single global Hessian across all pixels, correctly handling the joint space.
optimizer = torch.optim.Adam([raw_p], lr=kwargs.get("lr", 0.05))
logging.info(f"--- GPU {mode} Processing ({len(valid_idx)} pixels) ---")
pbar = tqdm(total=max_iter, desc=f"Optimizing ({mode})")
prev_loss = float("inf")
patience_count = 0
patience = kwargs.get("patience", 50)
try:
for step in range(max_iter):
optimizer.zero_grad()
loss = objective_fn(raw_p)
loss.backward()
optimizer.step()
pbar.update(1)
cur = loss.item()
if abs(prev_loss - cur) < 1e-7 * (abs(prev_loss) + 1e-10):
patience_count += 1
if patience_count >= patience:
pbar.update(max_iter - pbar.n)
break
else:
patience_count = 0
prev_loss = cur
except Exception as e:
logging.warning(f"Optimization interrupted: {e}")
pixel_health_map[valid_idx] = 0
pbar.close()
with torch.no_grad():
p_final = self._transform_params(raw_p, model_type)
fit_flat = self._model_kernel(p_final, t_axis, flat_irf, model_type)
res_flat = flat_data - fit_flat
dof = max((gate_end - gate_start) - p_final.shape[1], 1)
fit_sel = fit_flat[:, gate_start:gate_end]
data_sel = flat_data[:, gate_start:gate_end]
res_sel = res_flat[:, gate_start:gate_end]
fit_sel_np = fit_sel.detach().cpu().numpy().astype(np.float64)
data_sel_np = data_sel.detach().cpu().numpy().astype(np.float64)
chi2_raw_flat, chi2_red_flat = reduced_poisson_deviance(
fit_sel_np, data_sel_np, p_final.shape[1]
)
pearson_flat = pearson_chi_square(fit_sel_np, data_sel_np)
ss_tot = torch.sum(
(data_sel - data_sel.mean(dim=1, keepdim=True)) ** 2, dim=1
)
ss_res = torch.sum(res_sel**2, dim=1)
r2_flat = torch.where(
ss_tot > 0, 1.0 - ss_res / ss_tot, torch.zeros_like(ss_tot)
)
rmse_flat = torch.sqrt(torch.mean(res_sel**2, dim=1))
perr_flat = torch.zeros_like(p_final)
if CRLB:
perr_flat = self._compute_crlb_errors(
p_final, t_axis, flat_irf, model_type
)
full_popt = np.zeros((H * W, p_final.shape[1]))
full_perr = np.zeros((H * W, p_final.shape[1]))
full_fit = np.zeros((H * W, T))
full_res = np.zeros((H * W, T))
full_chi2_raw = np.zeros(H * W)
full_chi2_red = np.zeros(H * W)
full_r2 = np.zeros(H * W)
full_rmse = np.zeros(H * W)
full_popt[valid_idx] = p_final.detach().cpu().numpy()
full_perr[valid_idx] = perr_flat.detach().cpu().numpy()
full_fit[valid_idx] = fit_flat.detach().cpu().numpy()
full_res[valid_idx] = res_flat.detach().cpu().numpy()
full_chi2_raw[valid_idx] = chi2_raw_flat
full_chi2_red[valid_idx] = chi2_red_flat
full_pearson = np.zeros(H * W)
full_pearson[valid_idx] = pearson_flat
full_r2[valid_idx] = r2_flat.detach().cpu().numpy()
full_rmse[valid_idx] = rmse_flat.detach().cpu().numpy()
logging.info(f"Fit Finished in {time.time() - start_time:.2f}s")
health_mask = np.zeros(H * W)
health_mask[valid_idx] = 1.0
tau_lo, tau_hi = 1e-4, self.T_laser
p_np = p_final.detach().cpu().numpy()
chi2_red_np = full_chi2_red[valid_idx]
if model_type == "bi-exponential":
at_bound = (
(p_np[:, 2] <= tau_lo * 1.01)
| (p_np[:, 2] >= tau_hi * 0.99)
| (p_np[:, 3] <= tau_lo * 1.01)
| (p_np[:, 3] >= tau_hi * 0.99)
)
else:
at_bound = (p_np[:, 1] <= tau_lo * 1.01) | (p_np[:, 1] >= tau_hi * 0.99)
health_mask[valid_idx] = np.where(at_bound | (chi2_red_np > 5.0), 0.0, 1.0)
dataset = self._reconstruct_dataset(
full_popt.reshape(H, W, -1),
full_perr.reshape(H, W, -1),
full_fit.reshape(H, W, T),
full_res.reshape(H, W, T),
full_chi2_raw.reshape(H, W),
full_chi2_red.reshape(H, W),
full_r2.reshape(H, W),
full_rmse.reshape(H, W),
health_mask.reshape(H, W),
model_type,
mode,
data_name,
)
dataset["results"]["maps"]["pearson_chi2_map"] = full_pearson.reshape(
H, W
).astype(np.float32)
dataset["results"]["maps"]["pearson_reduced_chi2_map"] = (
full_pearson.reshape(H, W) / dof
).astype(np.float32)
return dataset
def _reconstruct_dataset(
self,
p_maps: np.ndarray,
e_maps: np.ndarray,
fit_map: np.ndarray,
res_map: np.ndarray,
chi2_raw: np.ndarray,
chi2_reduced: np.ndarray,
r2_map: np.ndarray,
rmse_map: np.ndarray,
health_map: np.ndarray,
model_type: str,
mode: str,
name: str,
) -> dict[Any, Any]:
"""
Reconstruct dataset.
Parameters
----------
p_maps : np.ndarray
Physical parameter maps reconstructed from flattened fit results.
e_maps : np.ndarray
Parameter-error maps reconstructed from flattened fit results.
fit_map : np.ndarray
Parameter or mask map processed by the routine.
res_map : np.ndarray
Parameter or mask map processed by the routine.
chi2_raw : np.ndarray
Raw chi-square map from the fit reconstruction.
chi2_reduced : np.ndarray
Reduced chi-square map from the fit reconstruction.
r2_map : np.ndarray
Parameter or mask map processed by the routine.
rmse_map : np.ndarray
Root-mean-square-error map from the fit reconstruction.
health_map : np.ndarray
Parameter or mask map processed by the routine.
model_type : str
FLI model family, such as mono- or bi-exponential.
mode : str
Mode selector used by the fitting, loading, or plotting routine.
name : str
Dataset, experiment, figure, or output name.
Returns
-------
dict[Any, Any]
Dictionary containing the data produced by reconstruct dataset.
"""
S = p_maps[..., 0]
common = {
"chi2_map": chi2_raw,
"reduced_chi2_map": chi2_reduced,
"R2_map": r2_map,
"rmse_map": rmse_map,
"pixel_health_map": health_map,
"convergence_map": health_map,
}
if model_type == "bi-exponential":
tau1_m, tau2_m = p_maps[..., 2], p_maps[..., 3]
alpha1_m = p_maps[..., 1]
maps = {
"photon_count_map": S,
"alpha1_map": alpha1_m,
"tau1_map": tau1_m,
"tau2_map": tau2_m,
"tau_mean_map": (alpha1_m * tau1_m + (1.0 - alpha1_m) * tau2_m).astype(
np.float32
),
"v_shift_map": p_maps[..., 4].astype(np.float32),
"h_shift_map": p_maps[..., 5].astype(np.float32),
"fret_efficiency_map": (
1.0
- np.divide(
tau1_m,
tau2_m,
out=np.zeros_like(tau2_m, dtype=np.float32),
where=tau2_m > 0,
)
).astype(np.float32),
**common,
}
else:
maps = {
"photon_count_map": S,
"tau_map": p_maps[..., 1],
"v_shift_map": p_maps[..., 2].astype(np.float32),
"h_shift_map": p_maps[..., 3].astype(np.float32),
**common,
}
return {
"name": name,
"method": f"GPU_{mode}",
"results": {
"maps": maps,
"error_maps": e_maps,
"TR_maps": {
"fit_map": fit_map.astype(np.float32),
"residual_map": res_map.astype(np.float32),
},
},
}
def _p0_to_tensor(self, p0: Any, n_pixels: int, model_type: str) -> Any:
"""
Run the p0 to tensor routine.
Parameters
----------
p0 : Any
Initial parameter vector supplied to the optimizer.
n_pixels : int
Number of samples, components, gates, or iterations used by the routine.
model_type : str
FLI model family, such as mono- or bi-exponential.
Returns
-------
Any
Object produced by p0 to tensor.
"""
if isinstance(p0, dict):
if model_type == "bi-exponential":
vals = [
p0.get("amp", 1000.0),
p0.get("alpha1", 0.2),
p0.get("tau1", 0.5),
p0.get("tau2", 1.1),
p0.get("v_shift", 10.0),
p0.get("h_shift", 0.0),
]
else:
vals = [
p0.get("amp", 1000.0),
p0.get("tau", 0.9),
p0.get("v_shift", 10.0),
p0.get("h_shift", 0.0),
]
arr = np.tile(vals, (n_pixels, 1)).astype(np.float32)
else:
row = np.asarray(p0, dtype=np.float32).ravel()
arr = np.tile(row, (n_pixels, 1))
return torch.tensor(arr, device=self.device, dtype=torch.float32)
def _get_sophisticated_guess(
self, data: np.ndarray, irf: np.ndarray, model_type: str
) -> Any:
"""
Return sophisticated guess.
Parameters
----------
data : np.ndarray
Data array or mapping processed by the routine.
irf : np.ndarray
Instrument response function aligned with the decay signal.
model_type : str
FLI model family, such as mono- or bi-exponential.
Returns
-------
Any
Object produced by get sophisticated guess.
"""
cpu_data = data.detach().cpu().numpy().astype(np.float64)
P, T = cpu_data.shape
t_axis = np.linspace(0, self.T_acq, T, endpoint=False)
dt = (t_axis[1] - t_axis[0]) if T > 1 else 1.0
offset_guess = np.percentile(cpu_data, 5, axis=1)
clean_d = np.clip(cpu_data - offset_guess[:, None], 1e-6, None)
idx_max = np.argmax(clean_d, axis=1)
col_idx = np.arange(T)[None, :]
post_peak = col_idx >= idx_max[:, None]
d_post = clean_d * post_peak
t_peak = t_axis[idx_max]
t_rel = np.maximum(t_axis[None, :] - t_peak[:, None], 0.0) * post_peak
m0 = np.trapezoid(d_post, dx=dt, axis=1).clip(min=1e-12)
m1 = np.trapezoid(t_rel * d_post, dx=dt, axis=1)
tau_mean = np.clip(m1 / m0, 0.05, self.T_laser * 0.8)
inside = -np.expm1(-self.T_acq / tau_mean)
s_guess = np.clip(clean_d.sum(axis=1) / inside, 1e-3, None)
offset_safe = np.clip(offset_guess, 0.0, None)
h_shift_guess = np.zeros(P)
if model_type == "mono-exponential":
guesses = np.stack([s_guess, tau_mean, offset_safe, h_shift_guess], axis=1)
else:
tau1 = np.clip(tau_mean * 0.5, 1e-4, self.T_laser * 0.99)
tau2 = np.clip(tau_mean * 1.5, tau1 * 1.01, self.T_laser)
# Neutral alpha1=0.5: the area-ratio estimate is unreliable when one
# lifetime is near the IRF width (fast component area ≈ slow area in
# any early/late split), so starting at 0.5 is always safer.
alpha1 = np.full(P, 0.5)
guesses = np.stack(
[s_guess, alpha1, tau1, tau2, offset_safe, h_shift_guess], axis=1
)
return torch.tensor(guesses.astype(np.float32), device=self.device)
[docs]
def save_results(self, dataset: np.ndarray, folder: str = "results") -> None:
"""
Save results.
Parameters
----------
dataset : np.ndarray
Dataset dictionary or fit result collection to save.
folder : str
Output directory used when saving results.
Returns
-------
None
No object is returned; the function save results.
"""
if not os.path.exists(folder):
os.makedirs(folder)
h5_path = os.path.join(folder, f"{dataset['name']}_GPU_results.h5")
with h5py.File(h5_path, "w") as f:
f.attrs["method"] = dataset["method"]
res_grp = f.create_group("results")
maps_grp = res_grp.create_group("maps")
for k, v in dataset["results"]["maps"].items():
maps_grp.create_dataset(
k, data=v.astype(np.float32), compression="gzip"
)
err_grp = res_grp.create_group("error_maps")
err_grp.create_dataset(
"errors", data=dataset["results"]["error_maps"], compression="gzip"
)
tr_grp = res_grp.create_group("TR_maps")
for k, v in dataset["results"]["TR_maps"].items():
tr_grp.create_dataset(k, data=v, compression="gzip")
logging.info(f"Dataset successfully saved to: {h5_path}")