Files
PromptCR-Lab/backend/app/analysis/charts.py
T
2026-09-19 12:54:45 +08:00

81 lines
2.4 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""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