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
|
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:
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)