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