File size: 23,114 Bytes
dc94f74
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
5b19c77
dc94f74
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
5b19c77
 
 
 
 
 
 
dc94f74
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
5b19c77
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
dc94f74
 
5b19c77
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
dc94f74
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
5b19c77
 
 
dc94f74
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
5b19c77
dc94f74
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
5b19c77
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
dc94f74
 
 
 
 
 
 
 
 
5b19c77
 
dc94f74
 
 
 
 
 
5b19c77
 
 
 
dc94f74
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
5b19c77
dc94f74
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489
490
491
492
493
494
495
496
497
498
499
500
501
502
503
504
505
506
507
508
509
510
511
512
513
514
515
516
517
518
519
520
521
522
523
524
525
526
527
528
529
530
531
532
533
534
535
536
537
538
539
540
541
542
543
544
545
#!/usr/bin/env python3
"""Export a `.pt` checkpoint to the published ONNX artifacts.

Run at **publish time**, from the checkpoint being published, so the
artifacts cannot drift from it. See `docs/onnx.md` for why this exists
and what each precision costs; the short version is that `onnxruntime`
is 27 MB against torch's 336 MB, and the receiving path only ever needs
two convolutional passes.

Six artifacts per checkpoint -- {encoder, decoder} x {fp32, fp16, int8}
-- named after the checkpoint they came from:

    v1.pt  ->  v1-encoder-fp32.onnx   v1-decoder-fp32.onnx
               v1-encoder-fp16.onnx   v1-decoder-fp16.onnx
               v1-encoder-int8.onnx   v1-decoder-int8.onnx

All three precisions are published with every revision because every
precision decodes on every other precision's receiver (measured: fp32
ONNX and torch agree to ~1e-6, roughly 112 dB below the channel noise).
Publishing them saves third parties from rolling their own export, which
is the case that actually risks divergence. **There is one on-air
format, and the precisions are not variants of it.**

Usage:

    scripts/export_onnx.py                      # published checkpoint
    scripts/export_onnx.py --model out/best.pt
    scripts/export_onnx.py --push               # upload to the Hub

Nothing is uploaded without `--push`. Every artifact is verified against
the torch model before it is written, and a tolerance breach fails the
run rather than publishing a bad codec.
"""

import argparse
import hashlib
import json
import shutil
import sys
import tempfile
from pathlib import Path

import numpy as np
import torch

sys.path.insert(0, str(Path(__file__).resolve().parent.parent))

from sstvae import checkpoint  # noqa: E402
from sstvae.codec import load_torch_model  # noqa: E402
from sstvae.config import LATENT_CHANNELS, LATENT_H, LATENT_W  # noqa: E402
from sstvae.images import IMG_H, IMG_W  # noqa: E402

OPSET = 17
PRECISIONS = ("fp32", "fp16", "int8")

# The yardstick for quantisation error is not fp32, it is the channel:
# at the modem's operating point there is this much RMS noise on
# unit-RMS latents. Quantisation noise is just one more small additive
# source on a channel that already carries a much larger one.
CHANNEL_NOISE_RMS = 0.367

# Gates, not targets. Measured values (docs/onnx.md) sit far below these;
# they exist to catch a broken export, not to police the last decimal.
#
#   latent_rms    RMS encoder error on unit-RMS latents. Compare against
#                 CHANNEL_NOISE_RMS -- that is the yardstick, not fp32.
#   quality_db    PSNR *lost* against the source image, running the whole
#                 pipeline at this precision versus running it in torch.
#
# The second one is the metric that means something. PSNR of an ONNX
# reconstruction against the torch reconstruction is a difference with no
# natural scale: int8 scores ~24 dB on it while costing ~0.15 dB of
# actual picture quality, so gating on it would reject a good artifact.
#
# The int8 gates stay loose enough to survive an untuned export failing
# gracefully, but the tuned one is far inside them: latent RMS 7.31e-02,
# costing 0.002 dB on photographs and 0.112 dB off-distribution.
# `quantize_dynamic` turns every Conv into `ConvInteger`, which supports
# only a **per-tensor** weight scale -- `per_channel=True` is silently a
# no-op on this graph, verified -- which is why the fix is per-layer
# exclusion rather than a quantiser setting.
TOLERANCES = {
    "fp32": {"latent_rms": 1e-4, "quality_db": 0.02},
    "fp16": {"latent_rms": 5e-3, "quality_db": 0.05},
    "int8": {"latent_rms": 0.25, "quality_db": 0.50},
}


def sha256(path: Path) -> str:
    h = hashlib.sha256()
    with open(path, "rb") as f:
        for chunk in iter(lambda: f.read(1 << 20), b""):
            h.update(chunk)
    return h.hexdigest()


def probe_images(n: int, seed: int = 0) -> torch.Tensor:
    """Deterministic stand-in images for verification, in [0,1].

    Low-frequency content rather than white noise: a convolutional
    autoencoder trained on photographs behaves quite differently on
    broadband noise, and a verification input that the model has no
    idea what to do with tells you little about whether the export is
    faithful. Pass `--images` for the real thing when it matters.
    """
    g = torch.Generator().manual_seed(seed)
    yy = torch.linspace(0, 1, IMG_H)[:, None]
    xx = torch.linspace(0, 1, IMG_W)[None, :]
    out = []
    for _ in range(n):
        img = torch.zeros(3, IMG_H, IMG_W)
        for c in range(3):
            acc = torch.zeros(IMG_H, IMG_W)
            for _ in range(6):
                fx, fy = torch.rand(2, generator=g) * 6
                ph = torch.rand(1, generator=g) * 6.283
                amp = torch.rand(1, generator=g)
                acc += amp * torch.sin(6.283 * (fx * xx + fy * yy) + ph)
            img[c] = acc
        img = (img - img.amin()) / (img.amax() - img.amin() + 1e-8)
        out.append(img)
    return torch.stack(out)


def load_images(paths: list[Path]) -> torch.Tensor:
    from PIL import Image

    from sstvae.images import fit_image, image_to_tensor

    return torch.stack([image_to_tensor(fit_image(Image.open(p))) for p in paths])


def export_fp32(module: torch.nn.Module, args: tuple, path: Path,
                input_names: list[str], output_names: list[str]) -> None:
    """Export one module at fp32.

    `dynamo=True` is the default and the maintained path; the legacy
    TorchScript exporter produces a numerically identical graph, so
    there is no reason to opt out. `external_data=False` is **not** the
    default and matters: without it the weights land in a `.onnx.data`
    sidecar, and four artifacts that must arrive together are worse than
    two.
    """
    torch.onnx.export(
        module, args, str(path),
        input_names=input_names, output_names=output_names,
        opset_version=OPSET, dynamo=True, external_data=False,
    )


def convert_fp16(src: Path, dst: Path) -> None:
    import onnx
    from onnxconverter_common import float16

    model = onnx.load(str(src))
    # keep_io_types: the caller still hands us fp32 arrays and gets fp32
    # back. The precision is an internal storage choice, not part of the
    # interface -- codec.py should not need to know which file it loaded.
    onnx.save(float16.convert_float_to_float16(model, keep_io_types=True), str(dst))


# int8 tuning: greedily keep the most quantisation-sensitive conv layers
# at fp32, because `ConvInteger` carries one scale per weight *tensor*
# and a single outlier-heavy tensor therefore sets a coarse scale for
# everything in it.
#
# **Tune on off-distribution content too, or this measures nothing.**
# Quantisation sensitivity barely shows on the photographs the model was
# trained on and shows enormously on everything else. Measured on v1's
# decoder: fully quantised costs 0.10 dB on COCO but **1.54 dB** on
# smooth synthetic probes, and excluding one layer takes the latter to
# -0.14 dB. Tuned against COCO alone the search sees 0.10 dB, concludes
# there is nothing to fix, and ships the 1.54 dB. That is not an
# academic case: operators send test cards, charts and screenshots.
INT8_TUNE_SYNTHETIC_PROBES = 2

# Stop once the error is under target -- **absolute, not a
# relative-improvement rule**. At the tail a large *fractional* error
# reduction is worth 0.00 dB of picture, which is how an earlier
# relative-only rule talked itself into three decoder exclusions that
# together bought 0.005 dB. `INT8_MIN_GAIN` remains only as a guard
# against grinding away at a target that cannot be reached.
#
# The encoder is measured in latent RMS against the channel noise, since
# that is the quantity that goes on the air. The decoder is measured in
# dB of picture, since its error reaches nobody else.
INT8_TARGET_DB_UNDER_CHANNEL = 13.0
# A tenth of a dB of your own picture. Tighter than this (0.05 was tried)
# buys a second decoder exclusion worth 0.02 dB for +1.8 MB, which is the
# wrong trade for the precision that exists for constrained devices.
INT8_TARGET_DECODER_DB = 0.10
INT8_MIN_GAIN = 0.15
INT8_MAX_EXCLUDED = 3


def convert_int8(src: Path, dst: Path, exclude: list[str] | None = None) -> None:
    from onnxruntime.quantization import QuantType, quantize_dynamic

    quantize_dynamic(str(src), str(dst), weight_type=QuantType.QInt8,
                     nodes_to_exclude=list(exclude or []))


def _conv_nodes(path: Path) -> list[str]:
    import onnx

    return [n.name for n in onnx.load(str(path)).graph.node
            if n.op_type in ("Conv", "ConvTranspose")]


def tune_int8_exclusions(src: Path, measure, target: float, *,
                         log=print) -> list[str]:
    """Find which conv layers to leave at fp32, by measuring.

    Quantisation error is not spread evenly across layers -- it is
    dominated by a few whose weight distribution has outliers, because
    `quantize_dynamic` emits `ConvInteger`, which carries **one scale
    per tensor**. A single bad tensor therefore sets a coarse scale for
    all of its weights. (This is also why `per_channel=True` does
    nothing here: ConvInteger has nowhere to put per-channel scales.)

    Measured on v1: excluding one 0.59 MB layer took the encoder from
    1.88e-01 to 7.31e-02 RMS -- 5.8 dB to 14.0 dB under the channel
    noise -- for 8% more file. That is worth finding automatically
    rather than hardcoding, because the sensitive layer is a property of
    the trained weights and will move between revisions.

    `measure(path) -> float` returns an error, lower being better, and
    `target` is the value below which the error stops mattering.
    """
    candidates = _conv_nodes(src)
    excluded: list[str] = []
    with tempfile.TemporaryDirectory() as td:
        def err(excl: list[str]) -> float:
            p = Path(td) / f"probe{len(excl)}.onnx"
            convert_int8(src, p, excl)
            return measure(p)

        current = err([])
        log(f"    all quantised: {current:.4e} (target {target:.4e})")
        while current > target and len(excluded) < INT8_MAX_EXCLUDED:
            scored = [(err(excluded + [c]), c)
                      for c in candidates if c not in excluded]
            if not scored:
                break
            best_err, best = min(scored)
            gain = (current - best_err) / current if current else 0.0
            if gain < INT8_MIN_GAIN:
                log(f"    stop: best remaining gain {gain:.1%} < "
                    f"{INT8_MIN_GAIN:.0%}")
                break
            excluded.append(best)
            current = best_err
            log(f"    keep {best} at fp32 -> {current:.4e} ({gain:.0%} better)")
        if current <= target:
            log(f"    at target ({current:.4e} <= {target:.4e})")
    return excluded


def stamp_metadata(path: Path, props: dict) -> None:
    """Record provenance inside the artifact.

    Which checkpoint an `.onnx` came from is exactly the question that
    gets asked when two stations disagree, and an answer that lives in
    the file cannot be separated from it.
    """
    import onnx

    model = onnx.load(str(path))
    for k, v in props.items():
        entry = model.metadata_props.add()
        entry.key, entry.value = str(k), str(v)
    onnx.save(model, str(path))


def run_onnx(path: Path, feeds: dict) -> np.ndarray:
    import onnxruntime as ort

    opts = ort.SessionOptions()
    opts.intra_op_num_threads = 4  # measured best; 1 was ~5x worse
    opts.log_severity_level = 3
    sess = ort.InferenceSession(str(path), opts, providers=["CPUExecutionProvider"])
    name = sess.get_inputs()[0].name
    if len(sess.get_inputs()) == 1:
        return sess.run(None, {name: feeds["primary"]})[0]
    second = sess.get_inputs()[1].name
    return sess.run(None, {name: feeds["primary"], second: feeds["secondary"]})[0]


def psnr(a: np.ndarray, b: np.ndarray) -> float:
    mse = float(np.mean((a.astype(np.float64) - b.astype(np.float64)) ** 2))
    return float("inf") if mse == 0 else 10.0 * float(np.log10(1.0 / mse))


def verify(model, images: torch.Tensor, paths: dict, precision: str) -> dict:
    """Compare one precision's encoder and decoder against torch.

    One image at a time: the graphs are exported at a fixed batch of 1,
    which is what the app actually does and what keeps the export in the
    easy, fully-static regime.
    """
    lat_errs, vs_torch, torch_q, onnx_q = [], [], [], []
    for i in range(images.shape[0]):
        one = images[i : i + 1]
        src = one.numpy().astype(np.float32)
        with torch.no_grad():
            ref_latents = model.encoder(one)
            weights = torch.ones_like(ref_latents)
            ref_image = model.decoder(ref_latents, weights).numpy()

        got_latents = run_onnx(paths["encoder"], {"primary": src})
        lat_errs.append(got_latents - ref_latents.numpy())

        w_np = weights.numpy().astype(np.float32)
        # Decoder alone, from the *torch* latents: isolates decoder error
        # instead of compounding the encoder's.
        vs_torch.append(psnr(
            run_onnx(paths["decoder"],
                     {"primary": ref_latents.numpy().astype(np.float32),
                      "secondary": w_np}),
            ref_image,
        ))
        # Whole pipeline at this precision -- what a station running these
        # artifacts actually gets -- measured against the source image.
        onnx_q.append(psnr(
            run_onnx(paths["decoder"],
                     {"primary": got_latents.astype(np.float32),
                      "secondary": w_np}),
            src,
        ))
        torch_q.append(psnr(ref_image, src))

    lat_err = np.concatenate([e.ravel() for e in lat_errs])
    latent_rms = float(np.sqrt(np.mean(lat_err ** 2)))
    quality_db = float(np.mean(torch_q) - np.mean(onnx_q))

    tol = TOLERANCES[precision]
    return {
        "precision": precision,
        "latent_rms": latent_rms,
        "latent_max": float(np.abs(lat_err).max()),
        "latent_vs_channel_db": (
            float("inf") if latent_rms == 0
            else 20.0 * float(np.log10(CHANNEL_NOISE_RMS / latent_rms))
        ),
        "torch_psnr_db": float(np.mean(torch_q)),
        "onnx_psnr_db": float(np.mean(onnx_q)),
        "quality_lost_db": quality_db,
        "decoder_vs_torch_psnr_db": float(np.mean(vs_torch)),
        "encoder_mb": paths["encoder"].stat().st_size / 1e6,
        "decoder_mb": paths["decoder"].stat().st_size / 1e6,
        "ok": latent_rms <= tol["latent_rms"] and quality_db <= tol["quality_db"],
    }


def main() -> int:
    ap = argparse.ArgumentParser(description=__doc__,
                                 formatter_class=argparse.RawDescriptionHelpFormatter)
    ap.add_argument("--model", default=None,
                    help="checkpoint .pt; defaults to the published one")
    ap.add_argument("--out", default="onnx", type=Path,
                    help="output directory (default: ./onnx)")
    ap.add_argument("--stem", default=None,
                    help="artifact name prefix; defaults to the checkpoint stem "
                         "(v1.pt -> v1-encoder-fp16.onnx)")
    ap.add_argument("--precisions", default=",".join(PRECISIONS),
                    help=f"comma-separated subset of {','.join(PRECISIONS)}")
    ap.add_argument("--images", type=Path, default=None,
                    help="directory of images to verify against; synthetic "
                         "low-frequency probes are used if omitted")
    ap.add_argument("--n-probe", type=int, default=4,
                    help="number of verification images (default 4)")
    ap.add_argument("--no-int8-tuning", action="store_true",
                    help="skip the per-layer int8 sensitivity search (faster, "
                         "but a worse int8 encoder -- see docs/onnx.md)")
    ap.add_argument("--push", action="store_true",
                    help="upload the artifacts to the Hub after verifying")
    ap.add_argument("--repo", default=checkpoint.DEFAULT_REPO,
                    help=f"Hub repo to push to (default {checkpoint.DEFAULT_REPO})")
    args = ap.parse_args()

    precisions = [p.strip() for p in args.precisions.split(",") if p.strip()]
    for p in precisions:
        if p not in PRECISIONS:
            ap.error(f"unknown precision {p!r}; choose from {', '.join(PRECISIONS)}")

    ckpt_path = Path(checkpoint.resolve(args.model))
    stem = args.stem or ckpt_path.stem
    ckpt_sha = sha256(ckpt_path)
    print(f"checkpoint  {ckpt_path.name}  sha256:{ckpt_sha[:16]}...")

    model = load_torch_model(args.model)
    model.eval()

    if args.images:
        files = sorted(
            p for p in args.images.iterdir()
            if p.suffix.lower() in {".jpg", ".jpeg", ".png", ".webp"}
        )[: args.n_probe]
        if not files:
            ap.error(f"no images found in {args.images}")
        images = load_images(files)
        print(f"verifying against {len(files)} image(s) from {args.images}")
    else:
        images = probe_images(args.n_probe)
        print(f"verifying against {args.n_probe} synthetic probe(s)")

    args.out.mkdir(parents=True, exist_ok=True)
    written: list[Path] = []
    results = []

    with tempfile.TemporaryDirectory() as tmp:
        tmp = Path(tmp)
        base = {
            "encoder": tmp / "encoder-fp32.onnx",
            "decoder": tmp / "decoder-fp32.onnx",
        }

        print("exporting fp32 graphs...")
        export_fp32(model.encoder, (images[:1],), base["encoder"],
                    ["image"], ["latents"])
        z0 = torch.zeros(1, LATENT_CHANNELS, LATENT_H, LATENT_W)
        export_fp32(model.decoder, (z0, torch.ones_like(z0)), base["decoder"],
                    ["latents", "weights"], ["image"])

        exclusions = {"encoder": [], "decoder": []}
        if "int8" in precisions and not args.no_int8_tuning:
            # A subset is plenty to *rank* layers -- the full set still
            # verifies the result afterwards -- but it must include
            # off-distribution probes. See INT8_TUNE_SYNTHETIC_PROBES.
            tune_imgs = torch.cat([
                images[:min(3, images.shape[0])],
                probe_images(INT8_TUNE_SYNTHETIC_PROBES, seed=1234),
            ])
            tune_n = tune_imgs.shape[0]
            src = tune_imgs.numpy()
            with torch.no_grad():
                ref_z = torch.cat([model.encoder(tune_imgs[i:i + 1])
                                   for i in range(tune_n)])
                ref_pic = torch.cat([
                    model.decoder(ref_z[i:i + 1], torch.ones_like(ref_z[i:i + 1]))
                    for i in range(tune_n)]).numpy()
            ref_z = ref_z.numpy()
            torch_q = float(np.mean([psnr(ref_pic[i], src[i]) for i in range(tune_n)]))

            def enc_err(path: Path) -> float:
                got = np.concatenate([
                    run_onnx(path, {"primary": src[i:i + 1]})
                    for i in range(tune_n)])
                return float(np.sqrt(np.mean((got - ref_z) ** 2)))

            def dec_err(path: Path) -> float:
                """dB of picture lost, decoding the *torch* latents."""
                got = np.concatenate([
                    run_onnx(path, {"primary": ref_z[i:i + 1],
                                    "secondary": np.ones_like(ref_z[i:i + 1])})
                    for i in range(tune_n)])
                q = float(np.mean([psnr(got[i], src[i]) for i in range(tune_n)]))
                return max(torch_q - q, 0.0)

            plan = {
                "encoder": (enc_err, CHANNEL_NOISE_RMS
                            / 10 ** (INT8_TARGET_DB_UNDER_CHANNEL / 20)),
                "decoder": (dec_err, INT8_TARGET_DECODER_DB),
            }
            for part, (measure, target) in plan.items():
                print(f"  tuning int8 {part} "
                      f"({tune_n} probes, {INT8_TUNE_SYNTHETIC_PROBES} synthetic)...")
                exclusions[part] = tune_int8_exclusions(
                    base[part], measure, target, log=print)

        for precision in precisions:
            paths = {}
            for part in ("encoder", "decoder"):
                dst = args.out / f"{stem}-{part}-{precision}.onnx"
                if precision == "fp32":
                    shutil.copyfile(base[part], dst)
                elif precision == "fp16":
                    convert_fp16(base[part], dst)
                else:
                    convert_int8(base[part], dst, exclusions[part])
                props = {
                    "sstvae.source_checkpoint": ckpt_path.name,
                    "sstvae.source_sha256": ckpt_sha,
                    "sstvae.part": part,
                    "sstvae.precision": precision,
                    "sstvae.opset": OPSET,
                    "sstvae.torch_version": torch.__version__,
                }
                if precision == "int8" and exclusions[part]:
                    props["sstvae.int8_fp32_layers"] = ",".join(exclusions[part])
                stamp_metadata(dst, props)
                paths[part] = dst

            r = verify(model, images, paths, precision)
            results.append(r)
            written.extend(paths.values())
            flag = "ok" if r["ok"] else "FAIL"
            print(
                f"  {precision:>4}  {r['encoder_mb']:5.1f} + {r['decoder_mb']:4.1f} MB  "
                f"latent RMS {r['latent_rms']:.2e} "
                f"({r['latent_vs_channel_db']:5.1f} dB under channel)  "
                f"picture {r['onnx_psnr_db']:6.2f} dB "
                f"({r['quality_lost_db']:+.3f} vs torch)  [{flag}]"
            )

    manifest = args.out / f"{stem}-onnx-manifest.json"
    manifest.write_text(json.dumps({
        "source_checkpoint": ckpt_path.name,
        "source_sha256": ckpt_sha,
        "opset": OPSET,
        "torch_version": torch.__version__,
        "channel_noise_rms": CHANNEL_NOISE_RMS,
        "int8_fp32_layers": exclusions,
        "artifacts": {
            p.name: {"sha256": sha256(p), "bytes": p.stat().st_size} for p in written
        },
        "verification": results,
    }, indent=2) + "\n")
    print(f"wrote {len(written)} artifact(s) + manifest to {args.out}/")

    failed = [r["precision"] for r in results if not r["ok"]]
    if failed:
        print(f"\nFAILED verification: {', '.join(failed)} -- nothing pushed",
              file=sys.stderr)
        return 1

    if args.push:
        from huggingface_hub import HfApi

        api = HfApi()
        api.create_repo(args.repo, exist_ok=True, repo_type="model")
        for p in [*written, manifest]:
            print(f"  uploading {p.name}")
            api.upload_file(path_or_fileobj=str(p), path_in_repo=p.name,
                            repo_id=args.repo,
                            commit_message=f"ONNX artifacts for {ckpt_path.name}")
        print(f"pushed to https://huggingface.co/{args.repo}")
    else:
        print("(not pushed; pass --push to upload)")
    return 0


if __name__ == "__main__":
    raise SystemExit(main())