pyfli.bayes_utils#

Provide Bayesian-inference tooling for PyFLI.

This module belongs to pyfli.bayes_utils and covers the direct-inference (BayesFlow/Keras) side of FLI decay fitting, downstream of pyfli.solver and pyfli.reconstruction: running a trained posterior-sampling model over a decay image (BiPipeline), selecting/aggregating per-pixel posterior-sample parameter combinations against measured decay (ParamSelector), and visualizing a single pixel’s posterior predictive fit (plot_pixel_posterior_fit()).

class BiPipeline(model_type, model_weights, patch_size=(128, 128), num_samples=100, batch_size=1024, custom_objects=None)[source]#

Bases: object

Encapsulates the Bi direct-inference pipeline for FLI decay maps.

MODEL_KEYS: ClassVar[dict[str, list[str]]] = {'bi-exponential': ['tau1', 'tau2', 'alpha1'], 'mono-exponential': ['tau']}#

keys produced per model type

load_model()[source]#

Load the Keras model checkpoint from self.model_weights.

Pass custom_objects at construction time if the checkpoint references custom architecture classes – this module no longer defines any inline.

run_inference(decay, irf, mask)[source]#

Run spatial-patch inference, matching the confirmed-working reference implementation’s input path exactly: each patch is sent to the model as a full, fixed-size batch (same shape/order the model was validated against), never a masked/subsetted batch.

Mask awareness is applied in two places only:
  1. A patch is skipped entirely – no model call at all – if every pixel in it is masked out.

  2. After the model call, masked-out pixels in that patch are zeroed back out in the output maps (post-hoc), so the model itself never sees a reduced or variable-size batch.

This deliberately avoids sending a variable-size, masked-down subset of pixels into model.sample() – an earlier version of this method did that as a compute-saving optimization, but it changed the batch composition/size the model saw per patch and produced near-uniform, wrong per-pixel outputs. Skipping fully-empty patches is safe (it doesn’t change what any processed patch’s batch looks like); subsetting pixels within a partially-masked patch is not.

Parameters:
  • decay (np.ndarray, shape (H, W, N_BINS))

  • irf (np.ndarray, shape (N_BINS,) shared across every pixel, or) – (H, W, N_BINS) per-pixel – same convention as ParamToDecay/DetailedRecon.

  • mask (np.ndarray, shape (H, W), bool-like)

Returns:

  • output_maps (dict[str, np.ndarray] each (H, W) -- per-pixel posterior) – median for each key.

  • output_uncertainties (dict[str, np.ndarray] each (H, W) -- per-pixel) – median absolute deviation (MAD) of the posterior samples, via scipy.stats.median_abs_deviation’s default scale=1.0. This is the raw MAD, not scaled to be std-comparable: for a roughly Gaussian posterior it runs ~0.6745x the equivalent standard deviation (use scale=’normal’ at the call site below if a 1-sigma-comparable value is ever needed instead).

  • output_samples (dict[str, np.ndarray] each (H, W, NUM_SAMPLES) --) – the raw per-pixel posterior draws model.sample() produced, kept around so other statistics can be computed later directly from them instead of re-running inference.

save_outputs(saver, output_maps, output_uncertainties, output_samples=None, tag=None)[source]#

Save output_maps, output_uncertainties, and (if given) output_samples via the provided saver object.

save_detailed(saver, detailed, name=None)[source]#

Save the DetailedRecon.reconstruct dict via the provided saver object.

compute_detailed(output_maps, freq, decay, irf)[source]#

Run DetailedRecon.reconstruct on the output maps for the configured model_type.

Returns the raw results dict (bi_bi or bi_mono equivalent).

static default_cmap()[source]#

Lazily build the default colormap (jet with lowest value pinned to zero).

visualize(saver, output_maps, output_uncertainties, cmap=None)[source]#

Display parameter maps and their uncertainties using DataViewer. Only applicable when model_type == “bi-exponential” (alpha1/tau1/tau2). cmap defaults to default_cmap() (jet_m) when not provided.

No pixel-coordinate/decay-curve argument here: every array plotted (output_maps/output_uncertainties) is a (H, W) 2-D map, and DataViewer.display_data only adds its extra decay-curve panel when a 3-D (H, W, T) array is present in data_list – so a coord would never draw anything, just reserve a blank column.

run(decay, irf, mask, freq, saver=None, save_outputs=True, compute_detailed=True, visualize=False, cmap=None)[source]#

Run the full pipeline: load model -> inference -> (save) -> (detailed) -> (visualize).

cmap is optional; if omitted and visualize=True, defaults to default_cmap() (jet with lowest value pinned to zero, i.e. jet_m).

Returns:

  • dict with keys ("output_maps", "output_uncertainties", "output_samples", "detailed")

  • (``”detailed”`` is None if compute_detailed=False)

class ParamSelector(freq_acq, irf, decay, model_type='bi-exponential', backend='compute_detailed_results', bool_mask=None)[source]#

Bases: object

Evaluate a stack of posterior-sample parameter combinations against measured decay data, and select the best-fitting combination per pixel by a chosen goodness-of-fit metric.

Parameters:
  • freq_acq (float) – Acquisition frequency (e.g. 80 MHz), passed straight through to compute_detailed_results’ freq_acq argument.

  • irf (np.ndarray) – Passed through to compute_detailed_results unchanged for every sample and for the final best-combination re-fit.

  • decay (np.ndarray) – Passed through to compute_detailed_results unchanged for every sample and for the final best-combination re-fit.

  • bool_mask (np.ndarray | None) – Optional (H, W) boolean mask – e.g. the same ROI mask passed to BiPipeline.run_inference. Not used during sample evaluation/selection itself (compute_detailed_results’ own pixel_health_map is derived from decay, not this mask, so a pixel with background counts outside the real ROI can otherwise look “healthy” even though its params are meaningless placeholders). Stored as the default for compute_best_model_fit_result’s own bool_mask argument, which uses it to NaN-out excluded pixels in the final result’s output maps.

  • model_type (str) – “bi-exponential” or “mono-exponential”.

  • backend (str) – Which implementation reconstructs each sample’s fit/residual/ goodness-of-fit maps: "compute_detailed_results" (default, the rescale-to-measured-totals implementation in pyfli.reconstruction.detailed_results) or "reconstructor" (pyfli.reconstruction.ParamToDecay’s vectorized path). Both return the same {"name", "method", "results": {"maps", "error_maps", "TR_maps"}} shape, so switching backends doesn’t change any downstream code.

METRICS: ClassVar[dict[str, tuple[str, str]]] = {'R2': ('r2_stack', 'max'), 'RMSE': ('rmse_stack', 'min'), 'chi2': ('chi2_stack', 'min'), 'reduced_chi2': ('reduced_chi2_stack', 'min')}#

(stack_key, “min” | “max”)} – whether the metric should be minimized (chi2/reduced_chi2/RMSE) or maximized (R2) to find the best-fitting sample per pixel.

Type:

Registry of {metric

BACKENDS: tuple[str, ...] = ('compute_detailed_results', 'reconstructor')#

Valid values for the backend constructor argument.

REDUCERS: ClassVar[dict[str, Callable]] = {'mean': <function mean>, 'median': <function median>}#

numpy function} for compute_aggregate_model_fit_result’s non-“best” methods – collapses the NUM_SAMPLES axis of each output_combination array to a single per-pixel value. Add an entry here to support another reduction (e.g. “mode”) without touching the method itself.

Type:

Registry of {reducer

evaluate_all_samples(output_combination, keep_per_sample_results=False, progress=True)[source]#

Run compute_detailed_results once per posterior sample.

Parameters:
  • output_combination (dict[str, np.ndarray]) – e.g. {'tau1': (H,W,NUM_SAMPLES), 'tau2': (H,W,NUM_SAMPLES), 'alpha1': (H,W,NUM_SAMPLES)} for bi-exponential, or {'tau': (H,W,NUM_SAMPLES)} for mono-exponential.

  • keep_per_sample_results (bool) – If True, also returns the full compute_detailed_results() dict for every sample (memory-heavy for large images/NUM_SAMPLES). Default False – only the scalar metric stacks are kept.

  • progress (bool) – Show a tqdm progress bar over the NUM_SAMPLES loop. Default True; pass False for quiet use (e.g. a 1x1-pixel crop). The per-sample backend log line is suppressed regardless, so it never competes with the bar.

Returns:

Dictionary with keys:

  • chi2_stack, reduced_chi2_stack, rmse_stack, r2_stack : (H, W, NUM_SAMPLES)

  • per_sample_results : list[dict] or None

Return type:

dict

select_best_combination(output_combination, stacks, metric='RMSE')[source]#

For each pixel independently, pick the sample index that optimizes the chosen metric (minimizes chi2/reduced_chi2/RMSE, maximizes R2), and build the corresponding best-per-pixel parameter maps.

Parameters:
  • output_combination (dict[str, np.ndarray]) – Same dict passed to evaluate_all_samples (each array (H, W, NUM_SAMPLES)).

  • stacks (dict) – Output of evaluate_all_samples.

  • metric (str) – One of “chi2”, “reduced_chi2”, “RMSE”, “R2”.

Returns:

Dictionary with keys:

  • best_params : dict[str, np.ndarray] – best (H, W) map per key in output_combination (e.g. tau1/tau2/alpha1)

  • best_sample_idx : (H, W) int array – which sample index won per pixel

  • best_score : (H, W) float array – the winning metric value per pixel

Return type:

dict

compute_best_model_fit_result(output_combination, metric='RMSE', data_name='BI_MODEL_bi_best', bool_mask=None, stacks=None)[source]#

Full pipeline: evaluate every sample, pick the best per-pixel combination by metric, then re-run compute_detailed_results once more on that best combination.

Parameters:
  • output_combination (dict[str, np.ndarray]) – Same dict passed to evaluate_all_samples.

  • metric (str) – One of “chi2”, “reduced_chi2”, “RMSE”, “R2”.

  • data_name (str) – Dataset name recorded in the returned result dict.

  • stacks (dict | None) – Precomputed evaluate_all_samples() output for this exact output_combination. When given, the per-sample evaluation loop (NUM_SAMPLES reconstructions) is skipped and these stacks are used directly – so a caller that already ran evaluate_all_samples (e.g. to inspect the raw stacks or try several metrics) doesn’t pay for it twice. When None (default) it is computed internally. The stacks must come from the same output_combination; passing mismatched stacks yields wrong selections.

  • bool_mask (np.ndarray | None) – Optional (H, W) boolean mask; defaults to self.bool_mask (set at construction) when not given. Pixels where the mask is False are NaN’d out in every array under result[‘results’][‘maps’], [‘TR_maps’], and [‘error_maps’] – so excluded pixels (e.g. outside the real ROI) can’t be mistaken for ordinary fits downstream, since compute_detailed_results’ own pixel_health_map is derived from decay, not this mask. No masking is applied if both this argument and self.bool_mask are None.

Notes

The return value is exactly the dict compute_detailed_results itself returns – {"name", "method", "results": {"maps", "error_maps", "TR_maps"}} – so it plugs directly into the rest of the workflow (e.g. result['results']['maps'].keys(), result['results']['TR_maps']['fit_map'], saver.save_npy(...), DataViewer, etc.) exactly like any other compute_detailed_results output.

Two extra maps are folded into result['results']['maps']:

  • best_sample_idx_map : (H, W) – which posterior sample (0..NUM_SAMPLES-1) won at each pixel

  • <metric>_selection_map : (H, W) – the winning metric value at each pixel (e.g. reduced_chi2_selection_map)

Per-sample diagnostics (the full metric stacks across all samples, and the chosen metric name) are attached as an additional top-level key, sample_selection, without disturbing the primary "name"/"method"/"results" structure.

Returns:

Same shape as compute_detailed_results()’s return value, plus a top-level sample_selection key holding the raw per-sample stacks and best_params used to produce the final maps.

Return type:

dict

compute_aggregate_model_fit_result(output_combination, method='best', metric='RMSE', data_name='BI_MODEL_bi_aggregate', bool_mask=None, stacks=None)[source]#

Collapse the stack of posterior-sample parameter combinations down to a single per-pixel parameter map via method, then run the configured backend once on that combination.

Parameters:
  • output_combination (dict[str, np.ndarray]) – Same dict passed to evaluate_all_samples/compute_best_model_fit_result (each array (H, W, NUM_SAMPLES)).

  • method (str) –

    How to reduce the NUM_SAMPLES axis to a single per-pixel value:
    • ”best”per-pixel sample that optimizes metric

      (delegates to compute_best_model_fit_result; see that method for the extra ‘best_sample_idx_map’/’sample_selection’ diagnostics only this option adds).

    • any key in REDUCERS (default “mean”, “median”) :

      per-pixel reduction across samples, e.g. the per-pixel mean or median tau/alpha1.

  • metric (str) – Only used when method=”best”; one of “chi2”, “reduced_chi2”, “RMSE”, “R2” (see METRICS).

  • data_name (str) – Dataset name recorded in the returned result dict.

  • bool_mask (np.ndarray | None) – Optional (H, W) boolean mask; defaults to self.bool_mask (set at construction) when not given. Forwarded to compute_best_model_fit_result for method=”best”; for the reducer methods, applied the same way directly here – see compute_best_model_fit_result’s bool_mask parameter for what it does.

  • stacks (dict | None) – Only used when method=”best”; forwarded to compute_best_model_fit_result to skip the per-sample evaluation loop when the caller already has its output. Ignored by the reducer methods, which never score individual samples.

Returns:

dict – compute_best_model_fit_result for the exact structure notes). For method=”best” this is exactly compute_best_model_fit_result’s return value, including its extra ‘sample_selection’ key; the other methods have no per-sample selection to report, so that key is simply absent.

Return type:

same shape as compute_detailed_results()'s return value (see

plot_pixel_posterior_fit(output_combination, decay, irf, freq_acq, pixel, model_type='bi-exponential', center='median', metric='reduced_chi2', ci_levels=(92, 68), decay_color=_DECAY_COLOR, fit_color=_FIT_COLOR, band_alpha=0.88, fit_alpha=1.0, decay_alpha=0.8, title=None, ax=None)[source]#

Plot one pixel’s posterior-sample decay reconstructions as nested credible-interval bands, a chosen central curve, and the measured decay.

Parameters:
  • output_combination (dict[str, np.ndarray]) – Posterior-sample parameter maps, e.g. {'tau1': (H,W,NUM_SAMPLES), 'tau2': (H,W,NUM_SAMPLES), 'alpha1': (H,W,NUM_SAMPLES)} for bi-exponential, or {'tau': (H,W,NUM_SAMPLES)} for mono-exponential – same shape convention as ParamSelector.

  • decay (np.ndarray) – Measured decay, (H, W, T).

  • irf (np.ndarray) – IRF, (T,) (shared) or (H, W, T) (per-pixel).

  • freq_acq (float) – Acquisition frequency (MHz), i.e. freq[1].

  • pixel (tuple[int, int]) – (x, y) pixel to plot.

  • model_type (str) – "bi-exponential" or "mono-exponential".

  • center (str) – Which curve to draw as the central line: "median" or "mean" across posterior samples, or "best" (the single sample that optimizes metric at this pixel, via ParamSelector.select_best_combination()).

  • metric (str) – Only used when center="best"; one of ParamSelector.METRICS ("chi2", "reduced_chi2", "RMSE", "R2").

  • ci_levels (tuple[int, ]) – Nested credible-interval widths to shade, e.g. (92, 68) shades a 92% and a 68% band (percentiles (4, 96) and (16, 84) of the per-bin sample distribution).

  • decay_color (str) – Colour of the measured-decay line (defaults to _DECAY_COLOR).

  • fit_color (str) – Colour of the central fit curve and the credible bands (which are tinted-toward-white shades of it); defaults to _FIT_COLOR.

  • band_alpha (float or tuple[float, ]) – Opacity of the credible bands. A scalar applies to every band; a sequence sets them per band, matched positionally to ci_levels.

  • fit_alpha (float) – Opacity of the central fit curve.

  • decay_alpha (float) – Opacity of the measured-decay line.

  • title (str | None) – Axes title; defaults to f"Pixel ({x}, {y})".

  • ax (matplotlib.axes.Axes | None) – Axes to draw into. If omitted, a new figure/axes is created and shown.

Return type:

tuple[matplotlib.figure.Figure, matplotlib.axes.Axes]

Modules

inference

model_inference.py

param_combinations

posterior_pixel_plot

Plot a single pixel's posterior-sample decay reconstructions against its measured decay, in the style of a posterior-predictive check: shaded credible- interval bands plus a chosen central curve (best-fitting sample, median, or mean), overlaid on the actual measured decay.