File size: 1,849 Bytes
1a9eb8a 704fa80 1a9eb8a 704fa80 1a9eb8a | 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 | """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()
|