Download scripts/fix_gemm.py from swiftail/BiRefNet_lite-onnx-coreml: direct link, hf CLI and curl.
- Browser
- Download file 1.26 kB
-
https://huggingface.co/swiftail/BiRefNet_lite-onnx-coreml/resolve/main/scripts/fix_gemm.py
- Command line
-
hf download hf://swiftail/BiRefNet_lite-onnx-coreml/scripts/fix_gemm.py
-
curl -L -o fix_gemm.py https://huggingface.co/swiftail/BiRefNet_lite-onnx-coreml/resolve/main/scripts/fix_gemm.py
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") | |