File size: 7,160 Bytes
3c3ef46 | 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 | #!/usr/bin/env python3
"""Compute paired RGB metrics for publication candidates against Original BF16."""
from __future__ import annotations
import argparse
import hashlib
import json
import math
import os
from datetime import datetime, timezone
from pathlib import Path
from typing import Any
import numpy as np
from PIL import Image
REFERENCE_ID = "original_bf16"
EXPECTED_GRID_COUNT = 20
CHUNK_BYTES = 4 * 1024 * 1024
class MetricError(RuntimeError):
pass
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--grid-manifest", type=Path, required=True)
parser.add_argument("--output", type=Path, required=True)
return parser.parse_args()
def sha256_file(path: Path) -> str:
digest = hashlib.sha256()
with path.open("rb") as handle:
for chunk in iter(lambda: handle.read(CHUNK_BYTES), b""):
digest.update(chunk)
return digest.hexdigest()
def psnr(mse: float) -> float | str:
if mse == 0.0:
return "infinity"
return 10.0 * math.log10((255.0**2) / mse)
def load_rgb(record: dict[str, Any]) -> np.ndarray:
path = Path(str(record["path"])).resolve(strict=True)
if path.stat().st_size != int(record["bytes"]):
raise MetricError(f"byte drift: {path}")
if sha256_file(path) != str(record["sha256"]):
raise MetricError(f"hash drift: {path}")
with Image.open(path) as image:
if image.size != (1024, 1024):
raise MetricError(f"dimension drift: {path}: {image.size}")
return np.asarray(image.convert("RGB"), dtype=np.float64)
def atomic_json(path: Path, value: dict[str, Any]) -> None:
path.parent.mkdir(parents=True, exist_ok=True)
temporary = path.with_name(f".{path.name}.{os.getpid()}.tmp")
temporary.write_text(
json.dumps(value, indent=2, sort_keys=True) + "\n", encoding="utf-8"
)
os.replace(temporary, path)
def main() -> int:
args = parse_args()
manifest_path = args.grid_manifest.expanduser().resolve(strict=True)
manifest = json.loads(manifest_path.read_text(encoding="utf-8"))
grids = manifest.get("grids")
if not isinstance(grids, list) or len(grids) != EXPECTED_GRID_COUNT:
raise MetricError(f"expected {EXPECTED_GRID_COUNT} verified grids")
if manifest.get("status") != "verified" or manifest.get("profile") != "publication":
raise MetricError("grid manifest is not a verified publication profile")
expected_models: list[tuple[str, str]] | None = None
accumulators: dict[str, dict[str, Any]] = {}
prompt_ids: set[str] = set()
for grid in grids:
prompt_id = str(grid["prompt_id"])
if prompt_id in prompt_ids:
raise MetricError(f"duplicate prompt: {prompt_id}")
prompt_ids.add(prompt_id)
columns = grid.get("columns")
if not isinstance(columns, list):
raise MetricError(f"missing columns for {prompt_id}")
by_id = {str(column["model_id"]): column for column in columns}
if len(by_id) != len(columns) or REFERENCE_ID not in by_id:
raise MetricError(f"invalid model inventory for {prompt_id}")
models = [(str(column["model_id"]), str(column["label"])) for column in columns]
if expected_models is None:
expected_models = models
for model_id, label in models:
if model_id != REFERENCE_ID:
accumulators[model_id] = {
"label": label,
"sum_abs_error": 0.0,
"sum_squared_error": 0.0,
"value_count": 0,
"pairs": [],
}
elif models != expected_models:
raise MetricError(f"column order drift for {prompt_id}")
reference = load_rgb(by_id[REFERENCE_ID])
for model_id, _label in models:
if model_id == REFERENCE_ID:
continue
candidate = load_rgb(by_id[model_id])
difference = candidate - reference
absolute_error = float(np.abs(difference).sum(dtype=np.float64))
squared_error = float(np.square(difference).sum(dtype=np.float64))
value_count = int(difference.size)
pair_mse = squared_error / value_count
accumulator = accumulators[model_id]
accumulator["sum_abs_error"] += absolute_error
accumulator["sum_squared_error"] += squared_error
accumulator["value_count"] += value_count
accumulator["pairs"].append(
{
"prompt_id": prompt_id,
"seed": int(grid["seed"]),
"mae_rgb_8bit": absolute_error / value_count,
"mse_rgb_8bit": pair_mse,
"psnr_db": psnr(pair_mse),
}
)
if expected_models is None:
raise MetricError("no models found")
results: list[dict[str, Any]] = []
for model_id, label in expected_models:
if model_id == REFERENCE_ID:
continue
accumulator = accumulators[model_id]
value_count = int(accumulator["value_count"])
global_mse = float(accumulator["sum_squared_error"]) / value_count
global_mae = float(accumulator["sum_abs_error"]) / value_count
pair_scores = [float(pair["psnr_db"]) for pair in accumulator["pairs"]]
results.append(
{
"model_id": model_id,
"label": label,
"reference_model_id": REFERENCE_ID,
"pair_count": len(pair_scores),
"global_mae_rgb_8bit": global_mae,
"global_mse_rgb_8bit": global_mse,
"global_psnr_db": psnr(global_mse),
"mean_pair_psnr_db": float(np.mean(pair_scores)),
"median_pair_psnr_db": float(np.median(pair_scores)),
"min_pair_psnr_db": min(pair_scores),
"max_pair_psnr_db": max(pair_scores),
"pairs": accumulator["pairs"],
}
)
output = {
"schema_version": 1,
"status": "verified_complete",
"created_at": datetime.now(timezone.utc).isoformat(timespec="seconds"),
"source_grid_manifest_sha256": sha256_file(manifest_path),
"reference": {"model_id": REFERENCE_ID, "label": "Original BF16"},
"contract": {
"prompt_count": EXPECTED_GRID_COUNT,
"width": 1024,
"height": 1024,
"channels": "RGB",
"pixel_range": "8-bit [0,255]",
"same_prompt_seed_sampler_and_settings": True,
},
"results": results,
"limitations": [
"PSNR measures paired pixel similarity to the BF16 reference, not aesthetics or prompt adherence.",
"LPIPS and human-preference metrics were not computed.",
],
}
atomic_json(args.output.expanduser().resolve(), output)
print(json.dumps({"status": output["status"], "results": results}, indent=2))
return 0
if __name__ == "__main__":
raise SystemExit(main())
|