Source code for pyfli.simulator.sim_calibrator

"""
Validate simulated and experimental decay distributions with summary statistics and
plots.

This module belongs to :mod:`pyfli.simulator` and is part of PyFLI synthetic FLI/FLIM
data generation, hardware noise modeling, calibration, and validation tools. Public API
includes classes :class:`FLIValidator`.
"""

from typing import Any

import matplotlib.pyplot as plt

# simulator/sim_calibrator.py
import numpy as np
from scipy import stats
from scipy.spatial.distance import cosine
from scipy.special import rel_entr

from pyfli import logging


[docs] class FLIValidator: """ Compare simulated and experimental decay cubes with distribution summaries and diagnostic plots. It preprocesses cubes, computes similarity metrics, and reports validation statistics. Parameters ---------- method : str Algorithm or model-selection method to use. threshold : int Threshold applied to counts, masks, or statistics. """ def __init__(self, method: str = "analytical", threshold: int = 10) -> None: self.method = method.lower() self.threshold = threshold def _preprocess_cube(self, data_cube: np.ndarray) -> tuple[Any, ...]: """ Run the preprocess cube routine. Parameters ---------- data_cube : np.ndarray Array cube processed by the routine. Returns ------- tuple[Any, ...] Tuple containing processed decay cube, mask, and normalization metadata. """ if data_cube.ndim == 3: H, W, T = data_cube.shape flat_data = data_cube.reshape(-1, T) else: flat_data = data_cube T = flat_data.shape[-1] pixel_intensities = np.sum(flat_data, axis=1) valid_mask = pixel_intensities >= self.threshold filtered_data = flat_data[valid_mask] filtered_intensities = pixel_intensities[valid_mask] return filtered_data, filtered_intensities
[docs] def run_comprehensive_test( self, sim_dataset: np.ndarray, exp_decay_cube: np.ndarray, normalize: bool = False, ) -> Any: """ Args: normalize: If True, scales both intensity distributions to [0, 1] to compare noise morphology rather than absolute scale. """ sim_raw = sim_dataset["raw_data"]["decay"] exp_flat, exp_counts = self._preprocess_cube(exp_decay_cube) sim_flat_all, sim_counts_all = self._preprocess_cube(sim_raw) n_exp = exp_flat.shape[0] n_sim_total = sim_flat_all.shape[0] if n_exp == 0: return None # Fair subsampling if n_sim_total > n_exp: indices = np.random.choice(n_sim_total, size=n_exp, replace=False) sim_flat = sim_flat_all[indices] sim_counts = sim_counts_all[indices] else: sim_flat, sim_counts = sim_flat_all, sim_counts_all # --- PRE-NORMALIZATION TOGGLE --- if normalize: # Scale counts to [0, 1] based on their respective max values exp_counts = exp_counts / (np.max(exp_counts) + 1e-12) sim_counts = sim_counts / (np.max(sim_counts) + 1e-12) # Method 1: Temporal sim_vec = np.mean(sim_flat, axis=0) exp_vec = np.mean(exp_flat, axis=0) cos_sim = 1 - cosine(sim_vec, exp_vec) p = sim_vec / (np.sum(sim_vec) + 1e-12) q = exp_vec / (np.sum(exp_vec) + 1e-12) kl_div = np.sum(rel_entr(p, q)) # Method 2: Intensity Distribution ks_stat, p_value = stats.ks_2samp(sim_counts, exp_counts) bins = np.linspace( min(sim_counts.min(), exp_counts.min()), max(sim_counts.max(), exp_counts.max()), 50, ) hist_sim, _ = np.histogram(sim_counts, bins=bins, density=True) hist_exp, _ = np.histogram(exp_counts, bins=bins, density=True) intersection = np.minimum(hist_sim, hist_exp).sum() * (bins[1] - bins[0]) return { "cosine_similarity": cos_sim, "kl_divergence": kl_div, "ks_p_value": p_value, "hist_intersection": intersection, "sample_size": len(sim_counts), "sim_vec": sim_vec, "exp_vec": exp_vec, "sim_counts": sim_counts, "exp_counts": exp_counts, }
def _print_summary( self, cos_sim: Any, kl_div: np.ndarray, ks_stat: np.ndarray, p_value: np.ndarray, intersection: np.ndarray, n: Any, ) -> None: """ Print summary. Parameters ---------- cos_sim : Any Cosine similarity statistic printed in the validation summary. kl_div : np.ndarray Kullback-Leibler divergence values to plot or summarize. ks_stat : np.ndarray Kolmogorov-Smirnov statistic printed in the validation summary. p_value : np.ndarray P-value printed in the validation summary. intersection : np.ndarray Histogram-intersection score printed in the validation summary. n : Any Number of samples, bins, gates, or plotted items. Returns ------- None No object is returned; the function perform print summary. """ logging.info("\n" + "=" * 60) logging.info(f"STATISTICAL VALIDATION REPORT (N={n} Pixels)") logging.info("=" * 60) logging.info(f"{'Metric':<25} | {'Value':<15} | {'Target'}") logging.info("-" * 60) logging.info(f"{'Cosine Similarity':<25} | {cos_sim:<15.4f} | >0.99") logging.info(f"{'KL Divergence':<25} | {kl_div:<15.6f} | <0.01") logging.info(f"{'KS Statistic':<25} | {ks_stat:<15.4f} | -> 0.0") logging.info(f"{'KS P-Value':<25} | {p_value:<15.4e} | >0.05") logging.info(f"{'Hist Intersection':<25} | {intersection:<15.4f} | -> 1.0") logging.info("=" * 60 + "\n") def _plot_results( self, sim_vec: np.ndarray, exp_vec: np.ndarray, sim_counts: np.ndarray, exp_counts: np.ndarray, ) -> np.ndarray: """ Plot results. Parameters ---------- sim_vec : np.ndarray Simulated vector used by the validation plot. exp_vec : np.ndarray Experimental vector used by the validation plot. sim_counts : np.ndarray Simulated histogram counts used by the validation plot. exp_counts : np.ndarray Experimental histogram counts used by the validation plot. Returns ------- np.ndarray Matplotlib figure or axes containing the simulation results. """ fig, ax = plt.subplots(1, 2, figsize=(14, 5)) # Plot 1: Mean Temporal Decay (Log Scale) ax[0].semilogy(sim_vec, label="Simulated (Mean)", color="tab:blue", lw=2) ax[0].semilogy( exp_vec, "--", label="Experimental (Mean)", color="tab:orange", lw=2 ) ax[0].set_title("Temporal Profile Fidelity", fontweight="bold") ax[0].set_xlabel("Time Bin") ax[0].set_ylabel("Normalized Intensity") ax[0].legend() ax[0].grid(True, which="both", alpha=0.3) # Plot 2: Intensity Probability Density Function ax[1].hist( sim_counts, bins=50, alpha=0.5, label="Simulated", color="tab:blue", density=True, ) ax[1].hist( exp_counts, bins=50, alpha=0.5, label="Experimental", color="tab:orange", density=True, ) ax[1].set_title("Integrated Intensity PDF", fontweight="bold") ax[1].set_xlabel("Photon Counts (Integrated)") ax[1].set_ylabel("Probability Density") ax[1].legend() plt.tight_layout() plt.show() return fig