swiftail's picture
CoreML-friendly ONNX re-export of BiRefNet_lite (ZhengPeng7)
6383c88 verified
Raw History Blame Contribute Delete
1.26 kB
"""Rewrite Gemm(transB=0, B=initializer) -> Gemm(transB=1, B=B^T) so ORT's CoreML EP can put the
weight in weight.bin instead of inlining a transposed copy as hex text in model.mil.
usage: fix_gemm.py in.onnx out.onnx"""
import sys, onnx
from onnx import numpy_helper
m = onnx.load(sys.argv[1])
inits = {t.name: t for t in m.graph.initializer}
n_fixed = 0
done = set()
for n in m.graph.node:
if n.op_type != "Gemm" or n.input[1] not in inits:
continue
attrs = {a.name: a for a in n.attribute}
if "transB" in attrs and attrs["transB"].i == 1:
continue
name = n.input[1] + "_T"
if name not in done: # weights shared by several Gemms (backbone runs twice) -> one copy
w = numpy_helper.to_array(inits[n.input[1]]).T.copy()
m.graph.initializer.append(numpy_helper.from_array(w, name)); done.add(name)
n.input[1] = name
if "transB" in attrs:
attrs["transB"].i = 1
else:
n.attribute.append(onnx.helper.make_attribute("transB", 1))
n_fixed += 1
used = {i for n in m.graph.node for i in n.input}
keep = [t for t in m.graph.initializer if t.name in used]
del m.graph.initializer[:]
m.graph.initializer.extend(keep)
onnx.save(m, sys.argv[2])
print("fixed", n_fixed, "Gemm nodes")