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:
objectEncapsulates 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:
A patch is skipped entirely – no model call at all – if every pixel in it is masked out.
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:
objectEvaluate 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 inpyfli.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
backendconstructor 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:
- 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:
- 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 pixelbest_score: (H, W) float array – the winning metric value per pixel
- Return type:
- 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) – Precomputedevaluate_all_samples()output for this exactoutput_combination. When given, the per-sample evaluation loop (NUM_SAMPLES reconstructions) is skipped and these stacks are used directly – so a caller that already ranevaluate_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 sameoutput_combination; passing mismatched stacks yields wrong selections.bool_mask (
np.ndarray | None) – Optional (H, W) boolean mask; defaults toself.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 andself.bool_maskare 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_selectionkey holding the raw per-sample stacks and best_params used to produce the final maps.- Return type:
- 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).
- ”best”per-pixel sample that optimizes
- any key in
REDUCERS(default “mean”, “median”) : per-pixel reduction across samples, e.g. the per-pixel mean or median tau/alpha1.
- any key in
metric (
str) – Only used when method=”best”; one of “chi2”, “reduced_chi2”, “RMSE”, “R2” (seeMETRICS).data_name (
str) – Dataset name recorded in the returned result dict.bool_mask (
np.ndarray | None) – Optional (H, W) boolean mask; defaults toself.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 asParamSelector.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 optimizesmetricat this pixel, viaParamSelector.select_best_combination()).metric (
str) – Only used whencenter="best"; one ofParamSelector.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 (
floatortuple[float,]) – Opacity of the credible bands. A scalar applies to every band; a sequence sets them per band, matched positionally toci_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 tof"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
model_inference.py |
|
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. |