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