brulee-1's picture
Add files using upload-large-folder tool
9f08d74 verified
Raw
History Blame Contribute Delete
4.49 kB
"""Study plots and report helpers."""
from __future__ import annotations
import csv
from collections.abc import Sequence
from pathlib import Path
import numpy as np
from PIL import Image, ImageDraw
def write_metrics_csv(path: Path, rows: Sequence[dict[str, object]]) -> None:
"""raw rowsからmetrics.csvを生成する。"""
path.parent.mkdir(parents=True, exist_ok=True)
keys = sorted({key for row in rows for key in row})
with path.open("w", newline="", encoding="utf-8") as stream:
writer = csv.DictWriter(stream, fieldnames=keys, extrasaction="ignore")
writer.writeheader()
writer.writerows(rows)
def write_layer_metrics_csv(path: Path, rows: Sequence[dict[str, object]]) -> None:
"""rawのlayers配列を正規化してlayer_metrics.csvへ保存する。"""
flattened: list[dict[str, object]] = []
for row in rows:
for layer in row.get("layers", []):
flattened.append({"condition": row.get("condition"), **layer})
write_metrics_csv(path, flattened)
def make_layer_plot(rows: Sequence[dict[str, object]], path: Path) -> None:
"""layer-wise L2/MSE/cosine/changed plotを生成する。"""
import matplotlib.pyplot as plt
values = [(str(row.get("condition")), layer) for row in rows for layer in row.get("layers", [])]
if not values:
return
figure, axes = plt.subplots(2, 2, figsize=(12, 8), dpi=300)
for axis, key in zip(axes.flat, ("l2", "mse", "cosine", "changed_percent")):
for condition in sorted({name for name, _ in values}):
points = [float(layer[key]) for name, layer in values if name == condition and layer.get(key) is not None and np.isfinite(float(layer[key]))]
if points:
layer_names = [str(layer.get("layer")) for name, layer in values if name == condition and layer.get(key) is not None and np.isfinite(float(layer[key]))]
axis.plot(points, marker=".", label=condition)
axis.set_xticks(range(len(layer_names)), layer_names, rotation=60, ha="right", fontsize=5)
axis.set_title(key)
axis.grid(alpha=0.2)
axes[0, 0].legend(fontsize=6)
figure.tight_layout()
path.parent.mkdir(parents=True, exist_ok=True)
figure.savefig(path, dpi=300)
plt.close(figure)
def make_grid(images: Sequence[tuple[Image.Image, str]], path: Path, columns: int = 4) -> None:
"""画像とラベルから比較 grid を作成する。"""
if not images:
return
cell = max(max(image.width, image.height) for image, _ in images)
label_height = 28
rows = (len(images) + columns - 1) // columns
canvas = Image.new("RGB", (columns * cell, rows * (cell + label_height)), "white")
draw = ImageDraw.Draw(canvas)
for index, (image, label) in enumerate(images):
row, column = divmod(index, columns)
resized = image.convert("RGB").resize((cell, cell))
x, y = column * cell, row * (cell + label_height)
canvas.paste(resized, (x, y))
draw.text((x + 3, y + cell + 4), label[:36], fill="black")
canvas.save(path, dpi=(300, 300))
def make_diff_maps(reference: np.ndarray, candidate: np.ndarray, stem: Path) -> None:
"""signed/absolute pixel difference heatmapを保存する。"""
import matplotlib.pyplot as plt
diff = candidate.astype(float) - reference.astype(float)
for suffix, data, cmap in (("signed", diff.mean(axis=2), "coolwarm"), ("absolute", np.abs(diff).mean(axis=2), "magma")):
figure, axis = plt.subplots(figsize=(5, 5), dpi=300)
axis.imshow(data, cmap=cmap)
axis.set_axis_off()
figure.tight_layout(pad=0)
figure.savefig(stem.with_name(f"{stem.name}_{suffix}.png"), dpi=300, bbox_inches="tight", pad_inches=0)
plt.close(figure)
def make_quality_plot(rows: Sequence[dict[str, object]], path: Path) -> None:
"""PSNR/storageのpublication-quality plotを保存する。"""
import matplotlib.pyplot as plt
valid = [row for row in rows if row.get("success") and row.get("psnr") is not None]
if not valid:
return
figure, axis = plt.subplots(figsize=(8, 5), dpi=300)
axis.scatter([float(row["compression_ratio"]) for row in valid], [float(row["psnr"]) for row in valid], s=24)
axis.set_xlabel("Compression ratio (PNG bytes / candidate bytes)")
axis.set_ylabel("Image PSNR (dB)")
axis.grid(alpha=0.25)
figure.tight_layout()
figure.savefig(path, dpi=300)
plt.close(figure)