multilingual-embedding-base / code /emb /build_release.py
Horizon-Labs's picture
Multilingual embedding base v1.0 (bge-m3 compatible)
1b929be verified
Raw History Blame Contribute Delete
4.86 kB
"""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)