"""Matplotlib-based chart generation for paper figures.""" import base64 from io import BytesIO from typing import Dict, List import matplotlib import matplotlib.pyplot as plt import numpy as np matplotlib.rcParams["font.sans-serif"] = ["DejaVu Sans"] matplotlib.rcParams["axes.unicode_minus"] = False def _to_base64(fig: matplotlib.figure.Figure) -> str: buf = BytesIO() fig.savefig(buf, format="png", dpi=150, bbox_inches="tight") buf.seek(0) return base64.b64encode(buf.read()).decode("utf-8") def heatmap(data: Dict[str, Dict[str, float]], title: str = "Heatmap") -> str: """Generate a heatmap from a nested dict (rows × columns).""" rows = list(data.keys()) cols = sorted({c for row in data.values() for c in row.keys()}) matrix = np.array([[data[row].get(col, 0.0) for col in cols] for row in rows]) fig, ax = plt.subplots(figsize=(8, 6)) im = ax.imshow(matrix, cmap="YlOrRd", aspect="auto") ax.set_xticks(np.arange(len(cols))) ax.set_yticks(np.arange(len(rows))) ax.set_xticklabels(cols) ax.set_yticklabels(rows) ax.set_title(title) for i in range(len(rows)): for j in range(len(cols)): text = ax.text(j, i, f"{matrix[i, j]:.2f}", ha="center", va="center", color="black") fig.colorbar(im, ax=ax) encoded = _to_base64(fig) plt.close(fig) return encoded def boxplot(groups: Dict[str, List[float]], title: str = "Boxplot") -> str: fig, ax = plt.subplots(figsize=(8, 6)) labels = list(groups.keys()) values = [groups[label] for label in labels] ax.boxplot(values) ax.set_xticklabels(labels) ax.set_title(title) ax.set_ylabel("Score") encoded = _to_base64(fig) plt.close(fig) return encoded def grouped_bar( data: Dict[str, Dict[str, float]], title: str = "Grouped Bar Chart", ) -> str: fig, ax = plt.subplots(figsize=(10, 6)) categories = list(data.keys()) subcategories = sorted({sc for row in data.values() for sc in row.keys()}) x = np.arange(len(categories)) width = 0.8 / len(subcategories) for idx, subcat in enumerate(subcategories): values = [data[cat].get(subcat, 0.0) for cat in categories] ax.bar(x + idx * width, values, width, label=subcat) ax.set_xticks(x + width * (len(subcategories) - 1) / 2) ax.set_xticklabels(categories) ax.set_ylabel("Score") ax.set_title(title) ax.legend() encoded = _to_base64(fig) plt.close(fig) return encoded