File size: 4,857 Bytes
1b929be
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
"""Package an embedding student as a sentence-transformers model + ONNX (sentence_embedding output) + int8 check.
python emb/build_release.py SRC_DIR BASE_HUB_ID OUT_DIR
SRC_DIR: emb/train_emb.py output (encoder + projection.pt). OUT_DIR: Transformer weights at the root, 1_Pooling (mean),
2_Dense (hidden -> 1024, no bias), 3_Normalize; onnx/model.onnx + onnx/model_quantized.onnx (inputs input_ids,
attention_mask; output sentence_embedding = normalised mean-pooled projection) and onnx_sweep.json (cosine fp32 vs int8).
"""
import json, os, shutil, sys
import numpy as np, torch
from huggingface_hub import hf_hub_download
from safetensors.torch import save_file
from transformers import AutoModel, AutoTokenizer

src, base, out = sys.argv[1:4]
os.makedirs(out, exist_ok=True)
for f in ["config.json", "model.safetensors", "tokenizer.json", "val_metrics.json", "train_log.json"]:
    if os.path.exists(f"{src}/{f}"):
        shutil.copy(f"{src}/{f}", f"{out}/{f}")
for f in ["tokenizer_config.json", "special_tokens_map.json"]:
    shutil.copy(hf_hub_download(base, f), f"{out}/{f}")
P = torch.load(f"{src}/projection.pt", map_location="cpu").float()   # (hidden, 1024)
hid = P.shape[0]
json.dump([{"idx": 0, "name": "0", "path": "", "type": "sentence_transformers.models.Transformer"},
           {"idx": 1, "name": "1", "path": "1_Pooling", "type": "sentence_transformers.models.Pooling"},
           {"idx": 2, "name": "2", "path": "2_Dense", "type": "sentence_transformers.models.Dense"},
           {"idx": 3, "name": "3", "path": "3_Normalize", "type": "sentence_transformers.models.Normalize"}], open(f"{out}/modules.json", "w"), indent=1)
json.dump({"max_seq_length": 512, "do_lower_case": False}, open(f"{out}/sentence_bert_config.json", "w"))
json.dump({"similarity_fn_name": "cosine", "prompts": {}, "default_prompt_name": None}, open(f"{out}/config_sentence_transformers.json", "w"))
os.makedirs(f"{out}/1_Pooling", exist_ok=True); os.makedirs(f"{out}/2_Dense", exist_ok=True); os.makedirs(f"{out}/3_Normalize", exist_ok=True)
json.dump({"word_embedding_dimension": hid, "pooling_mode_cls_token": False, "pooling_mode_mean_tokens": True, "pooling_mode_max_tokens": False,
           "pooling_mode_mean_sqrt_len_tokens": False, "pooling_mode_weightedmean_tokens": False, "pooling_mode_lasttoken": False,
           "include_prompt": True}, open(f"{out}/1_Pooling/config.json", "w"), indent=1)
json.dump({"in_features": hid, "out_features": 1024, "bias": False, "activation_function": "torch.nn.modules.linear.Identity"},
          open(f"{out}/2_Dense/config.json", "w"), indent=1)
save_file({"linear.weight": P.T.contiguous()}, f"{out}/2_Dense/model.safetensors")

tok = AutoTokenizer.from_pretrained(out); enc_m = AutoModel.from_pretrained(out, attn_implementation="eager").eval()


class Wrap(torch.nn.Module):
    def __init__(self, m, P):
        super().__init__(); self.m = m; self.P = torch.nn.Parameter(P, requires_grad=False)

    def forward(self, input_ids, attention_mask):
        h = self.m(input_ids=input_ids, attention_mask=attention_mask).last_hidden_state
        a = attention_mask.unsqueeze(-1).to(h.dtype); v = (h * a).sum(1) / a.sum(1)
        return torch.nn.functional.normalize(v @ self.P, dim=-1)


texts = ["How tall is the Eiffel Tower?", "La tour Eiffel mesure 330 mètres.", "电动汽车电池能用多久", "Der Eiffelturm ist mit 330 Metern das höchste Bauwerk von Paris. " * 8]
e = tok(texts, return_tensors="pt", padding=True)
os.makedirs(f"{out}/onnx", exist_ok=True); w = Wrap(enc_m, P).eval()
with torch.no_grad():
    ref = w(e["input_ids"], e["attention_mask"]).numpy()
    torch.onnx.export(w, (e["input_ids"], e["attention_mask"]), f"{out}/onnx/model.onnx", input_names=["input_ids", "attention_mask"],
                      output_names=["sentence_embedding"], dynamic_axes={"input_ids": {0: "batch", 1: "seq"}, "attention_mask": {0: "batch", 1: "seq"},
                                                                          "sentence_embedding": {0: "batch"}}, opset_version=17, dynamo=False)
import onnxruntime as ort
from onnxruntime.quantization import quantize_dynamic, QuantType
quantize_dynamic(f"{out}/onnx/model.onnx", f"{out}/onnx/model_quantized.onnx", weight_type=QuantType.QInt8, op_types_to_quantize=["Gather"])
feeds = {k: e[k].numpy() for k in ["input_ids", "attention_mask"]}
r = {}
for f in ["model.onnx", "model_quantized.onnx"]:
    o = ort.InferenceSession(f"{out}/onnx/{f}", providers=["CPUExecutionProvider"]).run(None, feeds)[0]
    r[f] = dict(cos_to_torch=float((o * ref).sum(-1).min()), mb=os.path.getsize(f"{out}/onnx/{f}") / 1e6)
json.dump(dict(gather_only=dict(min_cos=r["model_quantized.onnx"]["cos_to_torch"], mb=r["model_quantized.onnx"]["mb"]), fp32=r["model.onnx"], n=len(texts)),
          open(f"{out}/onnx_sweep.json", "w"), indent=1)
print(r, flush=True)