File size: 8,694 Bytes
eae424a da5fb15 eae424a | 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 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 | #!/usr/bin/env python3
"""Aggregate decoder training/evaluation JSON without manual transcription."""
from __future__ import annotations
import argparse
import hashlib
import json
import statistics
from collections import defaultdict
from pathlib import Path
def mean_sd(values: list[float]) -> dict[str, float | None]:
if not values:
return {"mean": None, "sample_stddev": None}
return {
"mean": statistics.fmean(values),
"sample_stddev": statistics.stdev(values) if len(values) > 1 else None,
}
def display(value: dict[str, float | None], digits: int) -> str:
mean = value["mean"]
stddev = value["sample_stddev"]
assert mean is not None
return (
f"{mean:.{digits}f}"
if stddev is None
else f"{mean:.{digits}f} ± {stddev:.{digits}f}"
)
def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument("--run-dir", action="append", required=True, type=Path)
parser.add_argument("--quality-name", default="quality-imagenette-validation.json")
parser.add_argument("--output-json", required=True, type=Path)
parser.add_argument("--output-markdown", required=True, type=Path)
parser.add_argument("--minimum-runs", type=int, default=3)
args = parser.parse_args()
if args.minimum_runs < 1:
raise SystemExit("--minimum-runs must be positive")
runs: list[dict[str, object]] = []
grouped: dict[tuple[object, ...], list[dict[str, object]]] = defaultdict(list)
for directory in args.run_dir:
training = json.loads((directory / "training.json").read_text(encoding="utf-8"))
quality = json.loads((directory / args.quality_name).read_text(encoding="utf-8"))
notes_path = directory / "run-notes.json"
notes = json.loads(notes_path.read_text(encoding="utf-8")) if notes_path.exists() else {}
checks = {
"decoder_hash": training["decoder_sha256"] == quality["decoder_sha256"],
"model_hash": training["model_sha256"] == quality["model_sha256"],
"manifest_hash": training["dataset_manifest_sha256"]
== quality["manifest_sha256"],
"layers": training["encoder_layers"] == quality["encoder_layers"],
"image_size": training["image_size"] == quality["image_size"],
}
if not all(checks.values()):
raise SystemExit(f"inconsistent run {directory}: {checks}")
run = {
"directory": directory.name,
"seed": training["seed"],
"encoder_layers": training["encoder_layers"],
"image_size": training["image_size"],
"training_images": training["training_images"],
"evaluation_images": quality["summary"]["count"],
"dataset": quality["dataset_name"],
"split": quality["split"],
"manifest_sha256": quality["manifest_sha256"],
"model_sha256": quality["model_sha256"],
"decoder_sha256": quality["decoder_sha256"],
"training_record_sha256": hashlib.sha256(
(directory / "training.json").read_bytes()
).hexdigest(),
"quality_record_sha256": hashlib.sha256(
(directory / args.quality_name).read_bytes()
).hexdigest(),
"training_seconds": training["training_seconds"],
# Interactive development runs are valid training/quality records,
# but their wall time is not a performance result. Timing is
# admitted only when the run notes opt in after documenting a
# controlled host state; absence of notes must fail closed.
"training_timing_valid": notes.get("training_timing_valid", False),
"final_training_l1": training["final_l1"],
"global_psnr_db": quality["summary"]["global_psnr_db"],
"median_image_psnr_db": quality["summary"]["psnr_db"]["median"],
"median_image_ssim": quality["summary"]["ssim"]["median"],
"median_image_mae": quality["summary"]["mae"]["median"],
}
runs.append(run)
key = (
run["encoder_layers"],
run["image_size"],
run["dataset"],
run["split"],
run["manifest_sha256"],
run["model_sha256"],
)
grouped[key].append(run)
aggregates: list[dict[str, object]] = []
for (layers, size, dataset, split, manifest_sha, model_sha), cell_runs in sorted(
grouped.items()
):
cell_runs.sort(key=lambda run: int(run["seed"]))
seeds = [int(run["seed"]) for run in cell_runs]
if len(cell_runs) < args.minimum_runs:
raise SystemExit(
f"{layers}L/{size}/{dataset}/{split} has {len(cell_runs)} runs; "
f"require at least {args.minimum_runs}"
)
if len(set(seeds)) != len(seeds):
raise SystemExit(f"duplicate seed in {layers}L/{size}/{dataset}/{split}: {seeds}")
training_counts = {int(run["training_images"]) for run in cell_runs}
evaluation_counts = {int(run["evaluation_images"]) for run in cell_runs}
if len(training_counts) != 1 or len(evaluation_counts) != 1:
raise SystemExit(
f"inconsistent sample counts in {layers}L/{size}/{dataset}/{split}"
)
valid_training_times = [
float(run["training_seconds"])
for run in cell_runs
if run["training_timing_valid"]
]
aggregate = {
"encoder_layers": layers,
"image_size": size,
"dataset": dataset,
"split": split,
"manifest_sha256": manifest_sha,
"model_sha256": model_sha,
"seeds": seeds,
"runs": len(cell_runs),
"training_images_per_run": cell_runs[0]["training_images"],
"evaluation_images_per_run": cell_runs[0]["evaluation_images"],
"global_psnr_db": mean_sd(
[float(run["global_psnr_db"]) for run in cell_runs]
),
"median_image_psnr_db": mean_sd(
[float(run["median_image_psnr_db"]) for run in cell_runs]
),
"median_image_ssim": mean_sd(
[float(run["median_image_ssim"]) for run in cell_runs]
),
"median_image_mae": mean_sd(
[float(run["median_image_mae"]) for run in cell_runs]
),
"valid_training_timing_runs": len(valid_training_times),
"training_seconds": mean_sd(valid_training_times),
}
aggregates.append(aggregate)
artifact = {"schema_version": 1, "runs": runs, "aggregates": aggregates}
args.output_json.parent.mkdir(parents=True, exist_ok=True)
args.output_markdown.parent.mkdir(parents=True, exist_ok=True)
args.output_json.write_text(json.dumps(artifact, indent=2) + "\n", encoding="utf-8")
lines = [
"# Decoder quality summary",
"",
"Generated from immutable training and per-image evaluation JSON.",
"",
"| Encoder | Size | Seeds | Validation images/run | Global PSNR (dB) | Median-image PSNR (dB) | Median RGB SSIM | Median MAE |",
"|---:|---:|---:|---:|---:|---:|---:|---:|",
]
for cell in aggregates:
lines.append(
f"| {cell['encoder_layers']}L | {cell['image_size']} | {cell['runs']} "
f"| {cell['evaluation_images_per_run']} "
f"| {display(cell['global_psnr_db'], 2)} "
f"| {display(cell['median_image_psnr_db'], 2)} "
f"| {display(cell['median_image_ssim'], 4)} "
f"| {display(cell['median_image_mae'], 4)} |"
)
lines.extend(
[
"",
"| Seed | Decoder SHA-256 | Training (s) | Final-batch train L1 | Global PSNR (dB) | Median RGB SSIM |",
"|---:|---|---:|---:|---:|---:|",
]
)
for run in sorted(runs, key=lambda item: (item["encoder_layers"], item["image_size"], item["seed"])):
training_seconds = (
f"{float(run['training_seconds']):.1f}"
if run["training_timing_valid"]
else "excluded"
)
lines.append(
f"| {run['seed']} | `{run['decoder_sha256']}` "
f"| {training_seconds} "
f"| {float(run['final_training_l1']):.5f} "
f"| {float(run['global_psnr_db']):.2f} "
f"| {float(run['median_image_ssim']):.4f} |"
)
args.output_markdown.write_text("\n".join(lines) + "\n", encoding="utf-8")
print(f"wrote {args.output_json} and {args.output_markdown}")
if __name__ == "__main__":
main()
|