mini-unet-colorizer / export_onnx.py
User-2468's picture
Release v3.0: selected sub-4M semantic colorizer, inference and verified ONNX
1a9eb8a verified
Raw History Blame
1.85 kB
"""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()