File size: 4,491 Bytes
9f08d74
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
"""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)