pyfli.bayes_utils.inference#

model_inference.py

Class-based wrapper around the Bi direct-inference pipeline:
  • loads a trained bi-/mono-exponential Keras model (pass custom_objects at construction time if the checkpoint references custom architecture classes – this module no longer defines any of its own)

  • runs spatial-patch, mask-aware inference: fully masked-out patches are skipped entirely (no model call), and within a partially-masked patch only the masked-in pixels are sent to the model

  • restitches per-pixel medians/MADs into output maps (masked-out pixels are left at 0), and keeps the full per-pixel posterior sample stack (output_samples) so downstream code can compute other statistics from the same posterior draws without re-running inference

  • optionally saves outputs via a saver object

  • optionally runs DetailedRecon.reconstruct on the outputs

  • optionally visualizes results via DataViewer

Usage#

from model_inference import BiPipeline

pipeline = BiPipeline(
    model_type="bi-exponential",
    model_weights="/mnt/e/.../biexpon/model.keras",
    patch_size=(128, 128),
)

results = pipeline.run(
    decay=binned_decay,
    irf=binned_irf,
    mask=b_bool_mask,
    freq=freq[1],
    saver=saver,
    save_outputs=True,
    compute_detailed=True,
    visualize=True,
)

output_maps          = results["output_maps"]
output_uncertainties = results["output_uncertainties"]
output_samples        = results["output_samples"]   # (H, W, NUM_SAMPLES) per key
bi_detailed           = results["detailed"]          # bi_bi or bi_mono

Classes

BiPipeline(model_type, model_weights[, ...])

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

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)