first commit
This commit is contained in:
@@ -0,0 +1,80 @@
|
||||
"""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
|
||||
Reference in New Issue
Block a user