litert-models / conversion /package /scripts /convert_onnx.py
unicorn who dev
Complete FireViewer Android learning SDK, source and verification records
330e909 verified
Raw History Blame Contribute Delete
913 Bytes
"""Fix batch/spatial dimensions explicitly before the direct FlatBuffer export."""
import json,subprocess,sys
from pathlib import Path
import onnx
source,out=map(Path,sys.argv[1:3])
spatial_size=int(sys.argv[3])
model=onnx.load(str(source),load_external_data=False)
shapes=[]
for inp in model.graph.input:
dims=[]
for axis,dim in enumerate(inp.type.tensor_type.shape.dim):
if dim.dim_value>0: dims.append(dim.dim_value)
elif axis==0:dims.append(1)
elif len(inp.type.tensor_type.shape.dim)==4 and axis in {2,3}:dims.append(spatial_size)
else:raise ValueError(f'Unresolved input dimension: {inp.name} axis={axis}')
shapes.append(inp.name+':'+','.join(map(str,dims)))
out.mkdir(parents=True,exist_ok=True)
(out/'input-shapes.json').write_text(json.dumps(shapes,indent=2)+'\n')
subprocess.run([str(Path(sys.executable).parent/'onnx2tf'),'-i',str(source),'-o',str(out),'-b','1','-ois',*shapes],check=True)