File size: 2,005 Bytes
7dd7a5b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
34
35
36
37
38
39
40
41
42
43
44
45
46
"""Make onnx-community/embeddinggemma-2-ONNX text-only for Vespa's hugging-face-embedder.

The upstream graph requires image_features/video_features/audio_features inputs
(each [num_tokens, 512]) that are concatenated after the text embeddings.
For text-only inference these are empty, so we replace them with constant
[0, 512] initializers. The weights are untouched: each variant is written as
<out>/<variant>/model.onnx referencing model.onnx_data (hard-linked to the
upstream data file), matching how modelhub/Vespa store external data.
"""

import os
import sys
from pathlib import Path

import numpy as np
import onnx
from onnx import numpy_helper

SRC = Path(sys.argv[1])  # onnx-community snapshot dir
OUT = Path(sys.argv[2])
VARIANTS = {"fp32": "model", "int8": "model_quantized", "q4": "model_q4"}
MEDIA_INPUTS = ["image_features", "video_features", "audio_features"]

for variant, name in VARIANTS.items():
    dst = OUT / variant
    dst.mkdir(parents=True, exist_ok=True)
    m = onnx.load(SRC / "onnx" / f"{name}.onnx", load_external_data=False)
    for inp in [i for i in m.graph.input if i.name in MEDIA_INPUTS]:
        dtype = onnx.helper.tensor_dtype_to_np_dtype(inp.type.tensor_type.elem_type)
        width = inp.type.tensor_type.shape.dim[1].dim_value
        m.graph.input.remove(inp)
        m.graph.initializer.append(
            numpy_helper.from_array(np.zeros((0, width), dtype=dtype), inp.name)
        )
    for t in m.graph.initializer:
        for e in t.external_data:
            if e.key == "location":
                assert e.value == f"{name}.onnx_data", e.value
                e.value = "model.onnx_data"
    onnx.save_model(m, dst / "model.onnx")
    data = dst / "model.onnx_data"
    data.unlink(missing_ok=True)
    os.link(os.path.realpath(SRC / "onnx" / f"{name}.onnx_data"), data)
    # onnx.checker rejects hard-linked data files; we validate by loading in ORT instead.
    print(variant, [i.name for i in m.graph.input], [o.name for o in m.graph.output])