"""
Render FLI maps, fit summaries, and selected-pixel traces.
This module belongs to :mod:`pyfli.data_vnp` and is part of PyFLI visualization,
normalization, plotting, and mono-versus-bi-exponential comparison tools. Public API
includes classes :class:`DataViewer`.
"""
import os
from typing import Any
import matplotlib.pyplot as plt
import numpy as np
from matplotlib import gridspec
from ..plot_style import legend_outside
[docs]
class DataViewer:
"""
Render fitted parameter maps, raw decay traces, IRF traces, and selected-pixel
summaries. The class is aimed at interactive inspection and publication-style
diagnostic figures.
Parameters
----------
save_path : str | None
Output path used when saving generated masks, figures, or data.
fig_name : str | None
Figure name or output stem used by plotting helpers.
"""
def __init__(
self, save_path: str | None = None, fig_name: str | None = None
) -> None:
self.fig_name = fig_name
self.save_path = save_path
if save_path and not os.path.exists(save_path):
os.makedirs(save_path)
def _apply_marker(self, ax: Any, coord: Any) -> None:
"""
Apply marker.
Parameters
----------
ax : Any
Matplotlib axes object on which the plot is drawn.
coord : Any
Pixel coordinate highlighted or summarized by the plot.
Returns
-------
None
No object is returned; the function apply marker.
"""
if coord is not None:
x, y = coord
ax.scatter(
y, x, color="red", s=40, marker="x", edgecolors="white", linewidths=1.5
)
[docs]
def display_data(
self,
data_list: np.ndarray,
structure: tuple[int, ...] = (1, 1),
coord: Any | None = None,
data_names: np.ndarray | None = None,
cmaps: np.ndarray | None = None,
v_ranges: np.ndarray | None = None,
figsize: np.ndarray | None = None,
normalize: bool = False,
yscale: str = "linear",
) -> tuple[Any, ...]:
"""
Display data.
Parameters
----------
data_list : np.ndarray
List of data arrays displayed by the viewer.
structure : tuple[int, ...]
Layout structure that controls how data arrays are displayed.
coord : Any | None
Pixel coordinate highlighted or summarized by the plot.
data_names : np.ndarray | None
Labels assigned to displayed data arrays.
cmaps : np.ndarray | None
Matplotlib colormap used by the visualization.
v_ranges : np.ndarray | None
Per-map display ranges used by the viewer.
figsize : np.ndarray | None
Figure size passed to Matplotlib.
normalize : bool
Whether data are normalized before comparison or plotting.
yscale : str
Scale used for the y-axis.
Returns
-------
tuple[Any, ...]
Tuple containing displayed data handles and summary values.
"""
num_plots = len(data_list)
r, c = structure
names = data_names or [f"Data {i + 1}" for i in range(num_plots)]
first_3d_idx = next((i for i, d in enumerate(data_list) if d.ndim == 3), None)
show_decay = coord is not None and first_3d_idx is not None
n_cols = c + (1 if show_decay else 0)
# FIX: Added layout="constrained" to figure initialization
fig = plt.figure(figsize=figsize or (n_cols * 5, r * 4), layout="constrained")
# Main grid
gs = gridspec.GridSpec(r, n_cols, figure=fig, wspace=0.25, hspace=0.25)
# Reference normalization
ref_max = None
if normalize and show_decay:
x, y = coord
ref_max = np.max(data_list[first_3d_idx][x, y, :])
# IMAGE PANELS
img_axes = []
for i in range(num_plots):
row, col = divmod(i, c)
# Subgrid: [main axis | colorbar axis]
subgs = gs[row, col].subgridspec(1, 2, width_ratios=[20, 1], wspace=0.05)
ax = fig.add_subplot(subgs[0])
cax = fig.add_subplot(subgs[1])
img_axes.append(ax)
data = data_list[i]
img = np.sum(data, axis=2) if data.ndim == 3 else data
if ref_max is not None:
img = img / (ref_max + 1e-9)
elif normalize is True:
img = (img - np.nanmin(img)) / (np.nanmax(img) - np.nanmin(img) + 1e-9)
cmap = cmaps[i] if cmaps and i < len(cmaps) else "viridis"
vr = v_ranges[i] if v_ranges and i < len(v_ranges) else (None, None)
im = ax.imshow(img, cmap=cmap, vmin=vr[0], vmax=vr[1])
self._apply_marker(ax, coord)
ax.set_title(names[i])
plt.colorbar(im, cax=cax)
# Plot Panel (forcing to same size)
ax_decay = None
if show_decay:
x, y = coord
subgs = gs[:, -1].subgridspec(1, 2, width_ratios=[20, 1], wspace=0.05)
ax_decay = fig.add_subplot(subgs[0])
cax_dummy = fig.add_subplot(subgs[1]) # placeholder
for i, data in enumerate(data_list):
if data.ndim == 3:
ax_decay.plot(data[x, y, :], label=names[i], lw=1.5)
ax_decay.set(
yscale=yscale,
title=f"Plots @ ({x},{y})",
# xlabel="Time Bin",
# ylabel="Counts"
)
legend_outside(ax_decay, fontsize="small")
ax_decay.grid(True, which="both", alpha=0.3)
cax_dummy.axis("off")
if self.save_path:
if self.fig_name:
plt.savefig(
os.path.join(self.save_path, self.fig_name + ".png"),
dpi=300,
bbox_inches="tight",
)
else:
plt.savefig(
os.path.join(self.save_path, "combined_display.png"),
dpi=300,
bbox_inches="tight",
)
plt.show()
return fig, img_axes, ax_decay
[docs]
def plot_pyfli_fit_summary(
self,
data: np.ndarray,
pixel: np.ndarray | None = None,
title: str = "FLI Fit Summary",
mode: tuple[str, ...] = ("decay", "irf", "fit", "residuals"),
esp: float = 1e0,
) -> Any:
# since simulator/pyfli img processing output is in specific disctionary
# best to use in the simulator
"""
Plot pyfli fit summary.
Parameters
----------
data : np.ndarray
Data array or mapping processed by the routine.
pixel : np.ndarray | None
Selected pixel coordinate.
title : str
Title displayed on the generated plot.
mode : tuple[str, ...]
Mode selector used by the fitting, loading, or plotting routine.
esp : float
Small epsilon used to stabilize divisions and comparisons.
Returns
-------
Any
Object produced by plot pyfli fit summary.
"""
mode = set(mode) # faster lookup
is_3d = pixel is not None
decay = data["raw_data"]["decay"]
irf = data["raw_data"]["irf"]
fit = data["results"]["TR_maps"]["fit_map"]
residuals = data["results"]["TR_maps"]["residual_map"]
maps = data["results"]["maps"]
if is_3d:
x, y = pixel
decay_1d = decay[x, y, :]
irf_1d = irf[x, y, :]
fit_1d = fit[x, y, :]
residuals_1d = residuals[x, y, :]
else:
decay_1d = decay
irf_1d = irf
fit_1d = fit
residuals_1d = residuals
eps = esp
decay_log = np.clip(decay_1d, eps, None)
fit_log = np.clip(fit_1d, eps, None)
irf_scaled = (irf_1d / np.max(irf_1d)) * np.max(decay_1d)
irf_log = np.clip(irf_scaled, eps, None)
def fmt(v: np.ndarray) -> Any:
"""
Run the fmt routine.
Parameters
----------
v : np.ndarray
Vector or matrix evaluated by the simplex projection.
Returns
-------
Any
Object produced by fmt.
"""
try:
return f"{float(v):.4f}"
except Exception:
return "NA"
lines = []
for k, v in maps.items():
try:
if is_3d and isinstance(v, np.ndarray):
val = v[x, y]
else:
val = v
except Exception:
val = v
lines.append(f"{k}: {fmt(val)}")
text_str = "\n".join(lines)
if is_3d:
fig = plt.figure(figsize=(20, 5))
gs = gridspec.GridSpec(1, 4, width_ratios=[1, 1, 0.6, 1], wspace=0.3)
ax1 = fig.add_subplot(gs[0])
ax2 = fig.add_subplot(gs[1])
ax_text = fig.add_subplot(gs[2])
ax3 = fig.add_subplot(gs[3])
else:
fig = plt.figure(figsize=(14, 5))
gs = gridspec.GridSpec(1, 3, width_ratios=[1, 1, 0.8], wspace=0.3)
ax1 = fig.add_subplot(gs[0])
ax2 = fig.add_subplot(gs[1])
ax_text = fig.add_subplot(gs[2])
ax3 = None
x_axis = np.arange(len(decay_1d))
# (1,2) Log --------
if "decay" in mode:
ax1.scatter(
x_axis,
decay_log,
s=20,
color="black",
marker="*",
alpha=0.7,
label="Decay (scatter)",
)
ax1.plot(
x_axis,
decay_log,
color="#1c4f8c",
lw=1.2,
alpha=0.6,
label="Decay (line)",
)
if "irf" in mode:
ax1.plot(
x_axis, irf_log, linestyle="--", color="#8c5f00", lw=1.5, label="IRF"
)
if "fit" in mode:
ax1.plot(x_axis, fit_log, color="green", lw=1.5, label="Fit")
if "residuals" in mode:
ax1.plot(
x_axis,
np.clip(residuals_1d, eps, None),
color="#8c2e24",
lw=1.2,
label="Residuals",
)
ax1.set_yscale("log")
ax1.set_ylim(eps, np.max(decay_log) * 1.2)
ax1.set_title("Log Scale")
ax1.set_xlabel("Time/Bins")
ax1.set_ylabel("Intensity")
ax1.grid(True, alpha=0.3)
legend_outside(ax1, fontsize=8)
# (1,2) LINEAR --------
if "decay" in mode:
ax2.scatter(
x_axis,
decay_1d,
s=20,
color="black",
marker="*",
alpha=0.7,
label="Decay (scatter)",
)
ax2.plot(
x_axis,
decay_1d,
color="#1c4f8c",
lw=1.2,
alpha=0.6,
label="Decay (line)",
)
if "irf" in mode:
ax2.plot(
x_axis, irf_scaled, linestyle="--", color="#8c5f00", lw=1.5, label="IRF"
)
if "fit" in mode:
ax2.plot(x_axis, fit_1d, color="green", lw=1.5, label="Fit")
if "residuals" in mode:
ax2.plot(x_axis, residuals_1d, color="#8c2e24", lw=1.2, label="Residuals")
ax2.set_title("Linear Scale")
ax2.set_xlabel("Time/Bins")
ax2.set_ylabel("Intensity")
ax2.grid(True, alpha=0.3)
legend_outside(ax2, fontsize=8)
# -------- TEXT PANEL --------
ax_text.axis("off")
ax_text.text(0.0, 1.0, text_str, fontsize=12, va="top", family="monospace")
ax_text.set_title("Fit Summary")
# -------- (1,3) IMAGE --------
if is_3d:
img = np.sum(decay, axis=2)
im = ax3.imshow(img)
ax3.plot(y, x, "rx", markersize=8, mew=2)
ax3.set_title(f"Intensity Map\nPixel ({x},{y})")
plt.colorbar(im, ax=ax3, fraction=0.046, pad=0.02)
fig.suptitle(title, fontsize=14)
plt.show()
return fig, (ax1, ax2, ax_text, ax3)
# Fixed, CVD-validated categorical order — assigned by series identity
# (never cycled/re-ranked), so "decay" is always the same hue across calls.
_SERIES_COLORS = ("#1c4f8c", "#8c3e1f", "#168c62", "#8c5f00", "#3e318c")
_INK = "#0b0b0b"
_MUTED = "#898781"
_AXIS = "#c3c2b7"
_GRID = "#e1e0d9"
[docs]
def plot_fli_px(
self,
data_list: np.ndarray,
pixel: np.ndarray | None = None, # pixel: (x, y) → enables time-series plotting
title: str = "FLI Data Viewer",
mode: str | None = None, # index in data_list which are to be added in the plot
mode2: Any | None = None, # index in the data_list which has to be displayed
names: Any | None = None,
esp: float = 1e0,
cmap: str = "viridis",
) -> tuple[Any, ...]:
"""
Plot FLI px.
Parameters
----------
data_list : np.ndarray
List of data arrays displayed by the viewer.
pixel : np.ndarray | None
Selected pixel coordinate.
title : str
Title displayed on the generated plot.
mode : str | None
Mode selector used by the fitting, loading, or plotting routine.
mode2 : Any | None
Secondary display mode used by the pixel viewer.
names : Any | None
Dataset names used in summaries and plots.
esp : float
Small epsilon used to stabilize divisions and comparisons.
cmap : str
Colormap used for the image panels.
Returns
-------
tuple[Any, ...]
Tuple containing the pixel plot handles and extracted pixel values.
"""
n_total = len(data_list)
if mode is None:
mode = list(range(n_total))
if mode2 is None:
mode2 = mode
selected_plot = [data_list[i] for i in mode]
selected_img = [data_list[i] for i in mode2]
labels_plot = names if names else [f"Data {i}" for i in mode]
labels_img = names if names else [f"Data {i}" for i in mode2]
is_pixel = pixel is not None
n_imgs = len(selected_img)
rc = {
"axes.edgecolor": self._AXIS,
"axes.labelcolor": self._INK,
"axes.linewidth": 1.0,
"axes.titlecolor": self._INK,
"axes.titleweight": "normal",
"axes.titlesize": 11,
"axes.labelsize": 10,
"xtick.color": self._MUTED,
"ytick.color": self._MUTED,
"xtick.labelsize": 9,
"ytick.labelsize": 9,
"grid.color": self._GRID,
"legend.frameon": False,
"legend.fontsize": 9,
"figure.facecolor": "white",
"savefig.facecolor": "white",
}
with plt.rc_context(rc):
if is_pixel:
fig = plt.figure(figsize=(5 * (n_imgs + 2), 4), layout="constrained")
gs = gridspec.GridSpec(1, n_imgs + 2, figure=fig)
ax_log = fig.add_subplot(gs[0])
ax_lin = fig.add_subplot(gs[1])
img_axes = [fig.add_subplot(gs[i + 2]) for i in range(n_imgs)]
else:
fig = plt.figure(figsize=(5 * n_imgs, 4), layout="constrained")
gs = gridspec.GridSpec(1, n_imgs, figure=fig)
img_axes = [fig.add_subplot(gs[i]) for i in range(n_imgs)]
eps = esp
if is_pixel:
x, y = pixel
for i, data in enumerate(selected_plot):
label = labels_plot[i]
color = self._SERIES_COLORS[i % len(self._SERIES_COLORS)]
if data.ndim == 3:
signal = data[x, y, :]
else:
signal = data
t = np.arange(len(signal))
ax_log.plot(t, np.clip(signal, eps, None), lw=2, color=color)
ax_lin.plot(t, signal, lw=2, color=color, label=label)
for ax, scale_title in (
(ax_log, "Log Scale"),
(ax_lin, "Linear Scale"),
):
ax.set_title(scale_title)
ax.set_xlabel("Time Bin")
ax.spines["top"].set_visible(False)
ax.spines["right"].set_visible(False)
ax.grid(True, alpha=0.6, lw=0.6)
ax.set_axisbelow(True)
ax_log.set_yscale("log")
ax_log.set_ylabel("Counts (log)")
ax_lin.set_ylabel("Counts")
if len(selected_plot) > 1:
legend_outside(ax_lin)
for i, data in enumerate(selected_img):
label = labels_img[i]
ax = img_axes[i]
if data.ndim == 3:
img = np.sum(data, axis=2)
im = ax.imshow(img, cmap=cmap)
if is_pixel:
ax.plot(
y,
x,
marker="x",
markersize=10,
markeredgewidth=2.5,
color="white",
zorder=5,
)
cbar = plt.colorbar(im, ax=ax, fraction=0.046, pad=0.02)
cbar.set_label("Total Counts", size=9)
cbar.ax.tick_params(labelsize=8)
ax.set_xlabel("X (px)")
ax.set_ylabel("Y (px)")
else:
ax.plot(data, lw=2, color=self._SERIES_COLORS[0])
ax.spines["top"].set_visible(False)
ax.spines["right"].set_visible(False)
ax.grid(True, alpha=0.6, lw=0.6)
ax.set_axisbelow(True)
ax.set_title(label)
fig.suptitle(title, fontsize=14, fontweight="bold", color=self._INK)
if self.save_path:
fname = (self.fig_name if self.fig_name else "fli_px_display") + ".png"
plt.savefig(
os.path.join(self.save_path, fname), dpi=300, bbox_inches="tight"
)
plt.show()
if is_pixel:
return fig, (ax_log, ax_lin, img_axes)
return fig, img_axes