"""Export the release network at fixed 256x256, with numerical verification.""" import argparse import importlib.util import tempfile from pathlib import Path import numpy as np import torch import onnx import onnxruntime as ort from semantic_model import load_semantic def main(): p=argparse.ArgumentParser();p.add_argument('--model',default='.');p.add_argument('--output',default='colorizer.onnx');args=p.parse_args() source=Path(__file__).with_name('semantic_model.py').read_text() source=source.replace("x=F.pad(self.neutral_rgb(L),(0,(-w)%32,0,(-h)%32),mode='replicate')","x=self.neutral_rgb(L)") source=source.replace("F.adaptive_avg_pool2d(f,(8,8)).flatten(2).transpose(1,2) for f in projected[1:]","F.avg_pool2d(f,k).flatten(2).transpose(1,2) for f,k in zip(projected[1:],[4,2,1])") with tempfile.TemporaryDirectory() as tmp: path=Path(tmp)/'export_model.py';path.write_text(source) spec=importlib.util.spec_from_file_location('export_model',path);module=importlib.util.module_from_spec(spec);spec.loader.exec_module(module) model=module.load_semantic(args.model,'cpu');reference=load_semantic(args.model,'cpu') x=torch.linspace(-1,1,256*256).reshape(1,1,256,256) with torch.no_grad():expected=reference(x).numpy();assert np.max(np.abs(expected-model(x).numpy()))<1e-5 torch.onnx.export(model,x,args.output,input_names=['L'],output_names=['ab'],opset_version=17,do_constant_folding=True,dynamo=False) onnx.checker.check_model(args.output) session=ort.InferenceSession(args.output,providers=['CPUExecutionProvider']) error=float(np.max(np.abs(session.run(None,{'L':x.numpy()})[0]-expected))) if error>=.005:raise RuntimeError(f'ONNX parity failed: {error}') print(f'Exported {args.output}; maximum Lab error {error:.6f}') if __name__=='__main__':main()