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())