"""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")