mini-unet-colorizer / export_onnx.py
User-2468's picture
Update validated colorizer weights, spatial decoder and deployment artifacts
704fa80 verified
Raw History Blame
2.66 kB
"""Export the colorizer AND luminance-guided chroma decoder as one ONNX graph.
Input: N,1,H,W normalized Lab luminance (L*/50-1), minimum8px per side.
Output: N,2,H,W Lab a,b values. Resize/preserve original L* in the app.
"""
import argparse,json
from pathlib import Path
import numpy as np
import torch
from model import load_model
from spatial import guided_chroma
class ChromaPipeline(torch.nn.Module):
def __init__(self,model,radius=8,temperature=.38):
super().__init__();self.model=model;self.radius=radius;self.temperature=temperature
def forward(self,L):
return guided_chroma(L,self.model.decode(self.model(L),self.temperature),self.radius)
def export(a):
import onnx,onnxruntime as ort
torch.set_num_threads(2);torch.manual_seed(2026)
model=ChromaPipeline(load_model(a.model),a.radius).eval()
out=Path(a.output);out.parent.mkdir(parents=True,exist_ok=True)
with torch.inference_mode():
torch.onnx.export(model,torch.zeros(1,1,256,256),str(out),opset_version=17,
input_names=['luminance'],output_names=['chroma'],dynamo=False,
dynamic_axes={'luminance':{0:'batch',2:'height',3:'width'},'chroma':{0:'batch',2:'height',3:'width'}})
onnx.checker.check_model(str(out))
opts=ort.SessionOptions();opts.intra_op_num_threads=2;opts.inter_op_num_threads=1
session=ort.InferenceSession(str(out),sess_options=opts,providers=['CPUExecutionProvider'])
checks=[]
for shape in [(1,1,256,256),(1,1,173,241),(2,1,64,80),(1,1,8,9)]:
x=torch.rand(shape)*2-1
with torch.inference_mode():expected=model(x).numpy()
actual=session.run(None,{'luminance':x.numpy()})[0]
error=np.abs(expected-actual)
assert actual.shape==expected.shape and np.isfinite(actual).all()
np.testing.assert_allclose(actual,expected,atol=.01,rtol=.001)
checks.append({'shape':list(shape),'max_ab_difference':float(error.max()),'mean_ab_difference':float(error.mean())})
result={'input':'N,1,H,W Lab L*/50-1; H,W>=8','output':'N,2,H,W Lab chroma a,b',
'guided_radius':a.radius,'guided_epsilon':.001,'temperature':.38,
'opset':17,'checks':checks,'parameters':sum(p.numel() for p in model.parameters()),
'weights_source':a.model,'runtime':ort.__version__,'bytes':out.stat().st_size}
out.with_suffix('.json').write_text(json.dumps(result,indent=2));print(json.dumps(result,indent=2))
if __name__=='__main__':
p=argparse.ArgumentParser(__doc__);p.add_argument('--model',required=True);p.add_argument('--output',required=True);p.add_argument('--radius',type=int,default=8)
export(p.parse_args())