File size: 1,255 Bytes
6383c88
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
"""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")