from __future__ import annotations
import os
import re
from typing import Callable, Dict, Iterable, Optional, Tuple
import matplotlib as mpl
import matplotlib.pyplot as plt
import mosaicperm as mp
import numpy as np
import pandas as pd
from matplotlib.lines import Line2D
from scipy.stats import uniform
from .base import BaseExperiment # noqa: F401 (użyte pośrednio przez API)
# Paper-friendly, serif typography
mpl.rcParams.update(
{
"font.family": "serif",
"mathtext.fontset": "cm",
"axes.titlesize": 11,
"axes.labelsize": 10,
"xtick.labelsize": 9,
"ytick.labelsize": 9,
"legend.fontsize": 9,
}
)
_LABEL_RE = re.compile(r"^(?P<gen>\w+)_(?P<size>\w+)_(?P<sym>sym|asym)_(?P<method>\w+)$")
_LABEL_MAP = {"perm": "Permutation", "sign": "Sign-flipping", "ridge": "Ridge", "ols": "OLS"}
[docs]
def parse_label(label: str) -> Dict[str, str]:
"""Parse an experiment label into components.
Args:
label: Label string of the form
``"<generator>_<size>_<sym|asym>_<method>"``.
Returns:
Dict[str, str]: Parsed components with keys:
* ``"gen"``,
* ``"size"``,
* ``"sym"``,
* ``"method"``.
"""
m = _LABEL_RE.match(label)
if not m:
parts = label.split("_")
return {
"gen": parts[0] if len(parts) > 0 else "unknown",
"size": parts[1] if len(parts) > 1 else "unknown",
"sym": parts[2] if len(parts) > 2 else "sym",
"method": parts[-1] if parts else "unknown",
}
return m.groupdict()
[docs]
def enrich_flatten(df_flat: pd.DataFrame) -> pd.DataFrame:
"""Add parsed label components as columns to a flattened DataFrame.
Args:
df_flat: Long-form DataFrame from :meth:`BaseExperiment.flatten`,
containing at least a ``"label"`` column.
Returns:
pandas.DataFrame: Copy of ``df_flat`` with extra columns added:
* ``"gen"``,
* ``"size"``,
* ``"sym"``,
* ``"method"``.
"""
parsed = df_flat["label"].apply(parse_label).apply(pd.Series)
cols_to_add = [c for c in ["gen", "size", "sym", "method"] if c not in df_flat.columns]
parsed = parsed[cols_to_add] if cols_to_add else parsed.iloc[:, 0:0]
return pd.concat([df_flat.reset_index(drop=True), parsed.reset_index(drop=True)], axis=1)
[docs]
def compute_power_table(df_flat: pd.DataFrame, alpha: float = 0.05) -> pd.DataFrame:
"""Aggregate empirical power across methods and configurations.
Args:
df_flat: Long-form DataFrame from :meth:`BaseExperiment.flatten`.
alpha: Significance level used to compute power (default 0.05).
Returns:
pandas.DataFrame: DataFrame with columns:
* ``"gen"``,
* ``"size"``,
* ``"sym"``,
* ``"method"``,
* ``"violation_strength"``,
* ``"power"``.
"""
df = enrich_flatten(df_flat).copy()
df["reject"] = df["p_value"] < alpha
grp = df.groupby(["gen", "size", "sym", "method", "violation_strength"], dropna=False)["reject"]
return grp.mean().reset_index().rename(columns={"reject": "power"})
[docs]
def plot_qq_grid_all_sizes(
df_flat: pd.DataFrame,
generators: Optional[Iterable[str]] = None,
sizes: Optional[Iterable[str]] = None,
v: float = 0.0,
symmetry_options: Iterable[str] = ("sym", "asym"),
method_linestyles: Optional[Dict[str, str]] = None,
figscale: Tuple[int, int] = (12, 16),
output_path: Optional[str] = None,
title_fontsize: int = 17,
label_fontsize: int = 16,
tick_fontsize: int = 15,
legend_fontsize: int = 14,
legend_ncol: Optional[int] = None,
legend_offset: float = -0.012,
tight_rect: Tuple[float, float, float, float] = (0.02, 0.05, 0.98, 0.96),
):
"""Create a QQ-grid by generator (rows) and symmetry (columns).
Each panel shows QQ-plots of p-values at a fixed violation level ``v``,
colored by size and styled by method.
Args:
df_flat: Long-form DataFrame from :meth:`BaseExperiment.flatten`.
generators: Generators to display. If ``None``, all available
generators are used.
sizes: Size categories to display. If ``None``, all available sizes
are used.
v: Violation strength to slice on.
symmetry_options: Iterable of symmetry flags (e.g. ``("sym", "asym")``).
method_linestyles: Mapping from method to line style, e.g.
``{"perm": "-", "sign": "--"}``. If ``None``, a default mapping
is used.
figscale: Figure size as ``(width, height)`` in inches.
output_path: Optional path to save the figure as an image.
title_fontsize: Font size for panel titles.
label_fontsize: Font size for axis labels.
tick_fontsize: Font size for tick labels.
legend_fontsize: Font size for legend labels.
legend_ncol: Number of columns in the combined legend. If ``None``,
a heuristic is used.
legend_offset: Vertical offset for the combined legend in figure
coordinates.
tight_rect: Bounding rectangle for :func:`matplotlib.pyplot.tight_layout`.
Returns:
matplotlib.figure.Figure: The created figure.
Raises:
ValueError: If the required label structure or requested generators /
sizes are not present in the data.
"""
df = enrich_flatten(df_flat).copy()
if method_linestyles is None:
method_linestyles = {"perm": "-", "sign": "--", "ridge": ":", "ols": "-."}
gens_avail = sorted(df["gen"].unique()) if "gen" in df.columns else []
if not gens_avail:
raise ValueError(
"No 'gen' column detected. Ensure labels follow "
"'<gen>_<size>_<sym|asym>_<method>'."
)
if generators is None:
generators = gens_avail
else:
generators = [g for g in generators if g in gens_avail]
if not generators:
raise ValueError("Requested generators not found in data.")
symmetry_options = list(symmetry_options)
n_rows = len(generators)
n_cols = len(symmetry_options)
fig, axes = plt.subplots(n_rows, n_cols, figsize=figscale, sharex=True, sharey=True)
axes = np.atleast_2d(axes)
sizes_avail = sorted(df["size"].unique()) if "size" in df.columns else []
if sizes is None:
sizes_used = sizes_avail
else:
sizes_used = [s for s in sizes if s in sizes_avail]
if not sizes_used:
raise ValueError("Requested sizes not found in data.")
cmap = plt.get_cmap("tab10", max(1, len(sizes_used)))
size_color_map = {size: cmap(i) for i, size in enumerate(sizes_used)}
for r, gen in enumerate(generators):
for c, sym in enumerate(symmetry_options):
ax = axes[r, c]
pane = df[
(df["gen"] == gen)
& (df["sym"] == sym)
& (df["violation_strength"] == v)
]
if pane.empty:
ax.set_visible(False)
continue
pane_sizes = [s for s in sorted(pane["size"].unique()) if s in sizes_used]
methods_present = sorted(pane["method"].unique())
for size in pane_sizes:
for method in methods_present:
sub = pane[(pane["size"] == size) & (pane["method"] == method)]
pvals = np.sort(sub["p_value"].to_numpy())
n = len(pvals)
if n == 0:
continue
uq = uniform.ppf((np.arange(1, n + 1) - 0.5) / n)
ax.plot(
uq,
pvals,
linestyle=method_linestyles.get(method, "-"),
color=size_color_map.get(size, "C0"),
linewidth=2,
)
ax.plot([0, 1], [0, 1], "k--", linewidth=1)
ax.set_xlim(0, 1)
ax.set_ylim(0, 1)
ax.set_aspect("equal", "box")
ax.grid(True, linestyle=":", linewidth=0.6)
if r == n_rows - 1:
ax.set_xlabel("Theoretical quantiles", fontsize=label_fontsize)
if c == 0:
ax.set_ylabel("Empirical quantiles", fontsize=label_fontsize)
ax.set_title(
f"{gen.capitalize()} — {'Sym' if sym == 'sym' else 'Asym'}",
fontsize=title_fontsize,
)
ax.tick_params(axis="both", labelsize=tick_fontsize)
methods_global = sorted(df["method"].unique())
handles, labels = [], []
for size in sizes_used:
for method in methods_global:
handles.append(
Line2D(
[0],
[0],
color=size_color_map[size],
lw=3,
linestyle=method_linestyles.get(method, "-"),
)
)
labels.append(f"{size.capitalize()} ({_LABEL_MAP.get(method, method)})")
if handles:
if legend_ncol is None:
legend_ncol = min(4, max(2, len(methods_global)))
fig.legend(
handles,
labels,
loc="lower center",
bbox_to_anchor=(0.5, legend_offset),
ncol=legend_ncol,
fontsize=legend_fontsize,
title_fontsize=legend_fontsize,
frameon=False,
title="Size (Method)",
)
plt.tight_layout(rect=tight_rect)
if output_path:
os.makedirs(os.path.dirname(output_path), exist_ok=True)
fig.savefig(output_path, dpi=300, bbox_inches="tight")
return fig
[docs]
def plot_qq_grid_generators_by_alpha(
df_flat: pd.DataFrame,
generators: Iterable[str],
alphas: Iterable[float],
sizes: Optional[Iterable[str]] = None,
symmetry: str = "asym",
method_linestyles: Optional[Dict[str, str]] = None,
figscale_base: Tuple[float, float] = (5.8, 4.8),
output_path: Optional[str] = None,
title_fontsize: int = 17,
label_fontsize: int = 16,
tick_fontsize: int = 15,
legend_fontsize: int = 14,
legend_ncol: int = 3,
legend_offset: float = -0.012,
tight_rect: Tuple[float, float, float, float] = (0.02, 0.05, 0.98, 0.94),
):
"""Create a QQ-grid with rows = generators and columns = violation strengths.
Args:
df_flat: Long-form DataFrame from :meth:`BaseExperiment.flatten`.
generators: Generators to include as rows.
alphas: Violation strength values (x-axis conditions) to include as
columns.
sizes: Size categories to display. If ``None``, all available sizes
are used.
symmetry: Symmetry flag to filter on (e.g. ``"asym"``).
method_linestyles: Mapping from method to line style. If ``None``,
a default mapping is used.
figscale_base: Base figure size for a single panel. Total figure size
scales with the grid dimensions.
output_path: Optional path to save the figure as an image.
title_fontsize: Font size for panel titles.
label_fontsize: Font size for axis labels.
tick_fontsize: Font size for tick labels.
legend_fontsize: Font size for legend labels.
legend_ncol: Number of columns in the combined legend.
legend_offset: Vertical offset for the combined legend.
tight_rect: Bounding rectangle for :func:`matplotlib.pyplot.tight_layout`.
Returns:
matplotlib.figure.Figure: The created figure.
Raises:
ValueError: If there is no data after filtering or requested sizes are
not present.
"""
df = enrich_flatten(df_flat).copy()
if method_linestyles is None:
method_linestyles = {"perm": "-", "sign": "--", "ridge": ":", "ols": "-."}
generators = list(generators)
alphas = list(alphas)
df = df[df["sym"] == symmetry]
if df.empty:
raise ValueError(f"No data after filtering for symmetry='{symmetry}'.")
sizes_avail = sorted(df["size"].unique()) if "size" in df.columns else []
if sizes is None:
sizes_used = sizes_avail
else:
sizes_used = [s for s in sizes if s in sizes_avail]
if not sizes_used:
raise ValueError("Requested sizes not found in data.")
n_rows, n_cols = len(generators), len(alphas)
fig_w = figscale_base[0] * n_cols
fig_h = figscale_base[1] * n_rows
fig, axes = plt.subplots(n_rows, n_cols, figsize=(fig_w, fig_h), sharex=True, sharey=True)
axes = np.atleast_2d(axes)
cmap = plt.get_cmap("tab10", max(1, len(sizes_used)))
size_color_map = {size: cmap(i) for i, size in enumerate(sizes_used)}
for r, gen in enumerate(generators):
for c, v in enumerate(alphas):
ax = axes[r, c]
pane = df[(df["gen"] == gen) & (df["violation_strength"] == v)]
if pane.empty:
ax.set_visible(False)
continue
for size in sizes_used:
for method in sorted(pane["method"].unique()):
sub = pane[(pane["size"] == size) & (pane["method"] == method)]
pvals = np.sort(sub["p_value"].to_numpy())
n = len(pvals)
if n == 0:
continue
uq = uniform.ppf((np.arange(1, n + 1) - 0.5) / n)
ax.plot(
uq,
pvals,
linestyle=method_linestyles.get(method, "-"),
color=size_color_map.get(size, "C0"),
linewidth=2,
)
ax.plot([0, 1], [0, 1], "k--", linewidth=1)
ax.set_xlim(0, 1)
ax.set_ylim(0, 1)
ax.set_aspect("equal", "box")
ax.grid(True, linestyle=":", linewidth=0.6)
if c == 0:
ax.set_ylabel("Empirical quantiles", fontsize=label_fontsize)
if r == n_rows - 1:
ax.set_xlabel("Theoretical quantiles", fontsize=label_fontsize)
ax.set_title(f"{gen.capitalize()}, α = {v:g}", fontsize=title_fontsize)
ax.tick_params(axis="both", labelsize=tick_fontsize)
methods_global = sorted(df["method"].unique())
handles, labels = [], []
for size in sizes_used:
for method in methods_global:
handles.append(
Line2D(
[0],
[0],
color=size_color_map[size],
lw=3,
linestyle=method_linestyles.get(method, "-"),
)
)
labels.append(f"{size.capitalize()} ({_LABEL_MAP.get(method, method)})")
if handles:
fig.legend(
handles,
labels,
loc="lower center",
bbox_to_anchor=(0.5, legend_offset),
ncol=legend_ncol,
fontsize=legend_fontsize,
title_fontsize=legend_fontsize,
frameon=False,
title="Size (Method)",
)
plt.tight_layout(rect=tight_rect)
plt.subplots_adjust(wspace=0, hspace=0.15)
if output_path:
os.makedirs(os.path.dirname(output_path), exist_ok=True)
fig.savefig(output_path, dpi=300, bbox_inches="tight")
return fig