Skip to content

plotting

tit.plotting

Plotting utilities for TI-Toolbox.

Non-Blender visualization and figure-generation helpers including intensity-vs-focality scatter plots, permutation null distributions, montage distribution plots, and the report slice figures (slices, dti_qc).

Most functions use lazy imports so import tit.plotting does not pull in matplotlib unless a plot function is actually called.

SaveFigOptions dataclass

SaveFigOptions(dpi: int = 600, bbox_inches: str = 'tight', facecolor: str = 'white', edgecolor: str = 'none')

Options forwarded to Figure.savefig.

Attributes

dpi : int Resolution in dots per inch (default 600). bbox_inches : str Bounding-box mode passed to savefig (default "tight"). facecolor : str Background colour of the saved figure (default "white"). edgecolor : str Edge colour of the saved figure (default "none").

ensure_headless_matplotlib_backend

ensure_headless_matplotlib_backend(backend: str = 'Agg') -> None

Best-effort backend setup for headless environments.

Important: - This should be called BEFORE importing matplotlib.pyplot. - If a backend is already active, we do not force-change it.

Source code in tit/plotting/_common.py
def ensure_headless_matplotlib_backend(backend: str = "Agg") -> None:
    """
    Best-effort backend setup for headless environments.

    Important:
    - This should be called BEFORE importing matplotlib.pyplot.
    - If a backend is already active, we do not force-change it.
    """
    import os
    import matplotlib

    os.environ.setdefault("MPLBACKEND", backend)

    # Silence noisy `findfont:` chatter (safe even if pyplot was already imported).
    suppress_matplotlib_findfont_noise()

    current = str(matplotlib.get_backend() or "")
    if current and current.lower() != backend.lower():
        # Backend already selected; don't override.
        return

    matplotlib.use(backend)  # type: ignore[attr-defined]

savefig_close

savefig_close(fig: Any, output_file: str, *, fmt: str | None = None, opts: SaveFigOptions = SaveFigOptions()) -> str

Save a matplotlib Figure and close it.

Uses fig.savefig (not plt.savefig) to avoid relying on global pyplot state.

Source code in tit/plotting/_common.py
def savefig_close(
    fig: Any,
    output_file: str,
    *,
    fmt: str | None = None,
    opts: SaveFigOptions = SaveFigOptions(),
) -> str:
    """
    Save a matplotlib Figure and close it.

    Uses fig.savefig (not plt.savefig) to avoid relying on global pyplot state.
    """
    fig.savefig(
        output_file,
        dpi=opts.dpi,
        bbox_inches=opts.bbox_inches,
        facecolor=opts.facecolor,
        edgecolor=opts.edgecolor,
        format=fmt,
    )
    import matplotlib.pyplot as plt

    plt.close(fig)

    return output_file

plot_cluster_size_mass_correlation

plot_cluster_size_mass_correlation(cluster_sizes: ndarray, cluster_masses: ndarray, output_file: str, *, dpi: int = 300) -> str | None

Plot correlation between cluster size and cluster mass from permutation null distribution.

Source code in tit/plotting/stats.py
def plot_cluster_size_mass_correlation(
    cluster_sizes: np.ndarray,
    cluster_masses: np.ndarray,
    output_file: str,
    *,
    dpi: int = 300,
) -> str | None:
    """
    Plot correlation between cluster size and cluster mass from permutation null distribution.
    """
    from scipy.stats import pearsonr

    ensure_headless_matplotlib_backend()
    import matplotlib.pyplot as plt

    import seaborn as sns

    sns.set_style("whitegrid")
    sns.set_context("notebook", font_scale=1.0)

    # Remove zeros
    mask = (cluster_sizes > 0) & (cluster_masses > 0)
    sizes_nonzero = cluster_sizes[mask]
    masses_nonzero = cluster_masses[mask]
    if len(sizes_nonzero) < 2:
        return None

    r_value, p_value = pearsonr(sizes_nonzero, masses_nonzero)

    fig, ax = plt.subplots(figsize=(10, 8))

    if sns is not None:
        sns.regplot(
            x=sizes_nonzero,
            y=masses_nonzero,
            ax=ax,
            scatter_kws={
                "alpha": 0.6,
                "s": 50,
                "color": "steelblue",
                "edgecolors": "black",
                "linewidths": 0.5,
            },
            line_kws={"color": "red", "linewidth": 2},
        )
    else:
        ax.scatter(
            sizes_nonzero,
            masses_nonzero,
            alpha=0.6,
            s=50,
            c="steelblue",
            edgecolors="black",
            linewidths=0.5,
        )
        z = np.polyfit(sizes_nonzero, masses_nonzero, 1)
        xs = np.linspace(
            float(np.min(sizes_nonzero)), float(np.max(sizes_nonzero)), 100
        )
        ax.plot(xs, z[0] * xs + z[1], color="red", linewidth=2)

    z = np.polyfit(sizes_nonzero, masses_nonzero, 1)
    ax.set_xlabel("Maximum Cluster Size (voxels)", fontsize=12, fontweight="bold")
    ax.set_ylabel(
        "Maximum Cluster Mass (sum of t-statistics)", fontsize=12, fontweight="bold"
    )
    ax.set_title(
        f"Cluster Size vs Cluster Mass Correlation\nPearson r = {r_value:.3f} (p = {p_value:.2e})",
        fontsize=14,
        fontweight="bold",
    )

    textstr = (
        f"n = {len(sizes_nonzero)} permutations\n"
        f"r = {r_value:.3f}\n"
        f"p = {p_value:.2e}\n"
        f"Linear fit: y = {z[0]:.2f}x + {z[1]:.2f}"
    )
    props = dict(boxstyle="round", facecolor="wheat", alpha=0.8)
    ax.text(
        0.05,
        0.95,
        textstr,
        transform=ax.transAxes,
        fontsize=11,
        verticalalignment="top",
        bbox=props,
    )

    ax.grid(True, alpha=0.3)
    fig.tight_layout()

    return savefig_close(fig, output_file, fmt="pdf", opts=SaveFigOptions(dpi=dpi))

plot_permutation_null_distribution

plot_permutation_null_distribution(null_distribution: ndarray, threshold: float, observed_clusters: Sequence[Mapping[str, float]], output_file: str, *, alpha: float = 0.05, cluster_stat: str = 'size', dpi: int = 300) -> str

Plot permutation null distribution with threshold and observed clusters.

Source code in tit/plotting/stats.py
def plot_permutation_null_distribution(
    null_distribution: np.ndarray,
    threshold: float,
    observed_clusters: Sequence[Mapping[str, float]],
    output_file: str,
    *,
    alpha: float = 0.05,
    cluster_stat: str = "size",
    dpi: int = 300,
) -> str:
    """
    Plot permutation null distribution with threshold and observed clusters.
    """
    ensure_headless_matplotlib_backend()
    import matplotlib.pyplot as plt
    import seaborn as sns

    sns.set_style("whitegrid")
    sns.set_context("notebook", font_scale=1.0)

    fig, ax = plt.subplots(figsize=(10, 6))

    # Labels based on cluster statistic
    if cluster_stat == "size":
        x_label = "Maximum Cluster Size (voxels)"
        title = "Permutation Null Distribution of Maximum Cluster Sizes"
        threshold_label = f"Discrete Threshold (p<{alpha}): {threshold:.1f} voxels"
    else:
        x_label = "Maximum Cluster Mass (sum of t-statistics)"
        title = "Permutation Null Distribution of Maximum Cluster Mass"
        threshold_label = f"Discrete Threshold (p<{alpha}): {threshold:.2f}"

    # Histogram
    if sns is not None:
        sns.histplot(
            null_distribution,
            bins=200,
            alpha=0.7,
            color="gray",
            edgecolor="black",
            label="Null Distribution",
            ax=ax,
        )
    else:
        ax.hist(
            null_distribution,
            bins=200,
            alpha=0.7,
            color="gray",
            edgecolor="black",
            label="Null Distribution",
        )

    # Threshold line
    ax.axvline(
        threshold, color="red", linestyle="--", linewidth=2, label=threshold_label
    )

    # Observed clusters
    sig_plotted = False
    nonsig_plotted = False
    for cluster in observed_clusters:
        stat_value = float(cluster["stat_value"])
        p_value = cluster.get("p_value", None)
        if p_value is not None:
            is_significant = float(p_value) < 0.05
        else:
            is_significant = stat_value > threshold

        color = "green" if is_significant else "orange"
        label = None
        if is_significant and not sig_plotted:
            label = "Significant Clusters (p<0.05)"
            sig_plotted = True
        elif (not is_significant) and (not nonsig_plotted):
            label = "Non-significant Clusters (p≥0.05)"
            nonsig_plotted = True

        ax.axvline(
            stat_value, color=color, linestyle="-", linewidth=2, alpha=0.7, label=label
        )

    ax.set_xlabel(x_label, fontsize=12)
    ax.set_ylabel("Frequency", fontsize=12)
    ax.set_title(title, fontsize=14, fontweight="bold")
    ax.legend(loc="upper right", fontsize=10)
    ax.grid(True, alpha=0.3)
    fig.tight_layout()

    return savefig_close(fig, output_file, fmt="pdf", opts=SaveFigOptions(dpi=dpi))

plot_electrode_score_heatmap

plot_electrode_score_heatmap(*, eeg_positions_csv: str, montage_scores: Sequence[dict], output_file: str, top_n: int = 50, dpi: int = 300, title_prefix: str = 'Ex-Search') -> str | None

Plot electrode participation in the top-top_n montages.

Each electrode receives the sum of the composite index of the top montages it appears in (colour) and its frequency (marker size); the 12 highest-scoring electrodes are labelled.

Source code in tit/plotting/ti_metrics.py
def plot_electrode_score_heatmap(
    *,
    eeg_positions_csv: str,
    montage_scores: Sequence[dict],
    output_file: str,
    top_n: int = 50,
    dpi: int = 300,
    title_prefix: str = "Ex-Search",
) -> str | None:
    """Plot electrode participation in the top-``top_n`` montages.

    Each electrode receives the sum of the composite index of the top
    montages it appears in (colour) and its frequency (marker size); the
    12 highest-scoring electrodes are labelled.
    """
    layout = _resolve_layout(eeg_positions_csv)
    positions = layout["positions"]
    if not positions or not montage_scores:
        return None

    ranked = _rank_plottable_montages(
        montage_scores=montage_scores,
        positions=positions,
        metric_key="composite",
        top_n=top_n,
    )
    if not ranked:
        return None

    scores = {label: 0.0 for label in positions}
    counts = {label: 0 for label in positions}
    for item in ranked:
        composite = float(item["composite"])
        for label in _montage_electrodes(item):
            scores[label] += composite
            counts[label] += 1
    active = [label for label, count in counts.items() if count > 0]

    ensure_headless_matplotlib_backend()
    import matplotlib.image as mpimg
    import matplotlib.pyplot as plt

    fig, ax = plt.subplots(figsize=(11, 9))
    label_offset = _draw_layout_background(ax, layout, mpimg)
    if layout["template_path"] is not None:
        inactive = [label for label in positions if label not in active]
        ax.scatter(
            [positions[label][0] for label in inactive],
            [positions[label][1] for label in inactive],
            s=12,
            color="#b0b0b0",
            edgecolors="none",
            zorder=1,
        )

    size_scale = 60 if layout["template_path"] is not None else 22
    sc = ax.scatter(
        [positions[label][0] for label in active],
        [positions[label][1] for label in active],
        c=[scores[label] for label in active],
        s=[60 + size_scale * counts[label] for label in active],
        cmap="inferno",
        edgecolors="black",
        linewidths=0.5,
        zorder=3,
    )
    for label in sorted(active, key=lambda lb: scores[lb], reverse=True)[:12]:
        x, y = positions[label]
        ax.text(
            x + label_offset,
            y - label_offset,
            label,
            fontsize=10,
            weight="bold",
            bbox={"facecolor": "white", "edgecolor": "none", "alpha": 0.65, "pad": 1},
            zorder=4,
        )

    fig.colorbar(sc, ax=ax, fraction=0.035, pad=0.01).set_label(
        "Summed Composite Index Across Top Montages"
    )
    ax.set_title(
        f"{title_prefix} Electrode Contribution (Top {len(ranked)} Montages, "
        f"{layout['eeg_net_name']})",
        fontsize=16,
        pad=12,
    )

    freq = [counts[label] for label in active]
    ax.text(
        0.02,
        0.02,
        "Color = summed composite index; marker size = frequency in the top "
        f"montages ({min(freq)}-{max(freq)} of {len(ranked)}).",
        transform=ax.transAxes,
        fontsize=9,
        color="#444444",
        va="bottom",
    )
    fig.tight_layout()
    return savefig_close(fig, output_file, opts=SaveFigOptions(dpi=dpi))

plot_intensity_vs_focality

plot_intensity_vs_focality(*, intensity: Sequence[float], focality: Sequence[float], composite: Sequence[float] | None, output_file: str, dpi: int = 300) -> str | None

Scatter plot of intensity vs focality, optionally colored by composite index.

Source code in tit/plotting/ti_metrics.py
def plot_intensity_vs_focality(
    *,
    intensity: Sequence[float],
    focality: Sequence[float],
    composite: Sequence[float] | None,
    output_file: str,
    dpi: int = 300,
) -> str | None:
    """
    Scatter plot of intensity vs focality, optionally colored by composite index.
    """
    if (not intensity) or (not focality):
        return None

    ensure_headless_matplotlib_backend()
    import matplotlib.pyplot as plt

    fig, ax = plt.subplots(figsize=(6, 5))
    if composite and any(c is not None for c in composite):
        sc = ax.scatter(
            intensity,
            focality,
            c=composite,
            cmap="viridis",
            s=40,
            edgecolor="black",
            alpha=0.7,
        )
        fig.colorbar(sc, ax=ax).set_label("Composite Index", fontsize=12)
    else:
        ax.scatter(intensity, focality, s=40, edgecolor="black", alpha=0.7)

    ax.set_xlabel("TImean_ROI (V/m)", fontsize=12)
    ax.set_ylabel("Focality", fontsize=12)
    ax.set_title("Intensity vs Focality", fontsize=14, fontweight="bold")
    ax.grid(alpha=0.3)
    fig.tight_layout()
    return savefig_close(fig, output_file, opts=SaveFigOptions(dpi=dpi))

plot_montage_distributions

plot_montage_distributions(*, timax_values: Sequence[float], timean_values: Sequence[float], focality_values: Sequence[float], output_file: str, dpi: int = 300) -> str | None

Create 3 side-by-side histograms for TImax, TImean and Focality distributions.

Source code in tit/plotting/ti_metrics.py
def plot_montage_distributions(
    *,
    timax_values: Sequence[float],
    timean_values: Sequence[float],
    focality_values: Sequence[float],
    output_file: str,
    dpi: int = 300,
) -> str | None:
    """
    Create 3 side-by-side histograms for TImax, TImean and Focality distributions.
    """
    if (not timax_values) and (not timean_values) and (not focality_values):
        return None

    ensure_headless_matplotlib_backend()
    import matplotlib.pyplot as plt

    fig, axes = plt.subplots(1, 3, figsize=(15, 4))
    configs = [
        (timax_values, axes[0], "TImax (V/m)", "TImax Distribution", "#2196F3"),
        (timean_values, axes[1], "TImean (V/m)", "TImean Distribution", "#4CAF50"),
        (focality_values, axes[2], "Focality", "Focality Distribution", "#FF9800"),
    ]

    for values, ax, xlabel, title, color in configs:
        if values:
            ax.hist(values, bins=20, color=color, edgecolor="black", alpha=0.7)
            ax.set_xlabel(xlabel, fontsize=12)
            ax.set_ylabel("Frequency", fontsize=12)
            ax.set_title(title, fontsize=14, fontweight="bold")
            ax.grid(axis="y", alpha=0.3)

    fig.tight_layout()
    return savefig_close(fig, output_file, opts=SaveFigOptions(dpi=dpi))

plot_montage_score_map

plot_montage_score_map(*, eeg_positions_csv: str, montage_scores: Sequence[dict], output_file: str, top_n: int = 50, dpi: int = 300, metric_key: str = 'composite', metric_label: str = 'Composite Index (TImean_ROI x Focality)', title_metric: str = 'Composite Score', cmap_name: str | tuple[str, str] = 'cividis', title_prefix: str = 'Ex-Search') -> str | None

Draw the top-top_n montages as electrode-pair curves on the EEG layout.

Every montage_scores record needs electrodes (4 or 8 labels as consecutive pairs; the legacy e1_plus..e2_minus keys are also accepted) and the requested metric_key. Curves are coloured by the metric; the best montage's electrodes are highlighted and labelled.

Source code in tit/plotting/ti_metrics.py
def plot_montage_score_map(
    *,
    eeg_positions_csv: str,
    montage_scores: Sequence[dict],
    output_file: str,
    top_n: int = 50,
    dpi: int = 300,
    metric_key: str = "composite",
    metric_label: str = "Composite Index (TImean_ROI x Focality)",
    title_metric: str = "Composite Score",
    cmap_name: str | tuple[str, str] = "cividis",
    title_prefix: str = "Ex-Search",
) -> str | None:
    """Draw the top-``top_n`` montages as electrode-pair curves on the EEG layout.

    Every ``montage_scores`` record needs ``electrodes`` (4 or 8 labels as
    consecutive pairs; the legacy ``e1_plus``..``e2_minus`` keys are also
    accepted) and the requested ``metric_key``.  Curves are coloured by the
    metric; the best montage's electrodes are highlighted and labelled.
    """
    layout = _resolve_layout(eeg_positions_csv)
    positions = layout["positions"]
    if not positions or not montage_scores:
        return None
    ranked = _rank_plottable_montages(
        montage_scores=montage_scores,
        positions=positions,
        metric_key=metric_key,
        top_n=top_n,
    )
    if not ranked:
        return None

    ensure_headless_matplotlib_backend()
    import matplotlib as mpl
    import matplotlib.image as mpimg
    import matplotlib.patches  # noqa: F401  (attribute access below)
    import matplotlib.path  # noqa: F401
    import matplotlib.pyplot as plt

    values = [float(item[metric_key]) for item in ranked]
    vmin, vmax = min(values), max(values)
    if vmin == vmax:
        vmax = vmin + 1e-12
    norm = mpl.colors.Normalize(vmin=vmin, vmax=vmax)
    cmap = _get_montage_metric_cmap(plt, mpl, cmap_name)

    fig, ax = plt.subplots(figsize=(11, 9))
    label_offset = _draw_layout_background(ax, layout, mpimg)

    for rank, item in enumerate(reversed(ranked), 1):
        color = cmap(norm(float(item[metric_key])))
        alpha = 0.18 + 0.62 * rank / len(ranked)
        linewidth = 1.0 + 3.0 * rank / len(ranked)
        for pair_idx, (a, b) in enumerate(_montage_pairs(_montage_electrodes(item))):
            _draw_pair_curve(
                ax,
                positions[a],
                positions[b],
                color=color,
                alpha=alpha,
                linewidth=linewidth,
                curve_side=1 if pair_idx % 2 == 0 else -1,
                mpl=mpl,
            )

    best = ranked[0]
    for label in _montage_electrodes(best):
        x, y = positions[label]
        ax.scatter(
            [x], [y], s=430, facecolors="none", edgecolors="#ffea00", linewidths=4, zorder=5
        )
        ax.text(
            x + label_offset,
            y - label_offset,
            label,
            color="black",
            fontsize=11,
            weight="bold",
            bbox={"facecolor": "white", "edgecolor": "none", "alpha": 0.65, "pad": 1},
            zorder=6,
        )

    sm = mpl.cm.ScalarMappable(norm=norm, cmap=cmap)
    sm.set_array([])
    fig.colorbar(sm, ax=ax, fraction=0.035, pad=0.01).set_label(metric_label)
    ax.set_title(
        f"Top {len(ranked)} {title_prefix} Montages by {title_metric} "
        f"({layout['eeg_net_name']})",
        fontsize=16,
        pad=12,
    )
    fig.tight_layout()
    return savefig_close(fig, output_file, opts=SaveFigOptions(dpi=dpi))