Download export_onnx.py from User-2468/mini-unet-colorizer: direct link, hf CLI and curl.
- Browser
- Download file 2.66 kB
-
https://huggingface.co/User-2468/mini-unet-colorizer/resolve/6467bfa9d1bb8071cc0987ad5889728b6181d65b/export_onnx.py
- Command line
-
hf download hf://User-2468/mini-unet-colorizer@6467bfa9d1bb8071cc0987ad5889728b6181d65b/export_onnx.py
-
curl -L -o export_onnx.py https://huggingface.co/User-2468/mini-unet-colorizer/resolve/6467bfa9d1bb8071cc0987ad5889728b6181d65b/export_onnx.py
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()) | |