"""织忆 MemoryWeave — bge-m3 ONNX 嵌入服务器 OpenAI /v1/embeddings 兼容接口,Go 代码零改动切换。 使用 ONNX Runtime CPU 推理,RTX 3050 4GB 无压力。 启动: python3 bge_embed_server.py 端口: 8000 模型: /home/muc/models/bge-m3/onnx/ """ import json import logging import math import os from http.server import HTTPServer, BaseHTTPRequestHandler import numpy as np import onnxruntime as ort from transformers import AutoTokenizer MODEL_PATH = os.environ.get("BGE_MODEL_PATH", "/home/muc/models/bge-m3/onnx") MODEL_FILE = os.environ.get("BGE_MODEL_FILE", "model.onnx") PORT = int(os.environ.get("BGE_PORT", "8000")) MAX_BATCH = int(os.environ.get("BGE_MAX_BATCH", "32")) logging.basicConfig(level=logging.INFO, format="[bge-embed] %(message)s") log = logging.getLogger(__name__) # ─── 初始化 ────────────────────────────────────────── log.info("加载 tokenizer: %s", MODEL_PATH) tokenizer = AutoTokenizer.from_pretrained(MODEL_PATH) log.info("加载 ONNX 模型: %s/%s", MODEL_PATH, MODEL_FILE) sess_options = ort.SessionOptions() sess_options.intra_op_num_threads = 4 sess_options.inter_op_num_threads = 2 sess_options.enable_cpu_mem_arena = False # 禁用 arena 分配器,防止内存逐渐扩大 sess_options.enable_mem_pattern = False # 禁用内存模式,避免碎片累积 session = ort.InferenceSession( os.path.join(MODEL_PATH, MODEL_FILE), sess_options=sess_options, providers=["CUDAExecutionProvider", "CPUExecutionProvider"], ) log.info("ONNX 模型就绪 — providers=%s", session.get_providers()) def encode(texts: list[str]) -> list[list[float]]: """批量编码 + mean pooling + L2 归一化""" inputs = tokenizer( texts, padding=True, truncation=True, max_length=8192, return_tensors="np", ) ort_inputs = { "input_ids": inputs["input_ids"], "attention_mask": inputs["attention_mask"], } outputs = session.run(None, ort_inputs) # ONNX 输出: [batch, seq_len, 1024] — token-level embeddings embeddings: np.ndarray = outputs[0] # Mean pooling — 按 attention_mask 加权平均 attention_mask = inputs["attention_mask"].astype(np.float32) mask_expanded = np.expand_dims(attention_mask, -1) # [batch, seq_len, 1] sum_embeddings = np.sum(embeddings * mask_expanded, axis=1) # [batch, 1024] sum_mask = np.clip(np.sum(mask_expanded, axis=1), 1e-9, None) # [batch, 1] embeddings = sum_embeddings / sum_mask # [batch, 1024] # L2 归一化 norms = np.linalg.norm(embeddings, axis=1, keepdims=True) norms = np.maximum(norms, 1e-12) embeddings = embeddings / norms return embeddings.tolist() class EmbedHandler(BaseHTTPRequestHandler): """OpenAI /v1/embeddings 兼容""" def log_message(self, fmt, *args): pass # 安静模式 def _respond(self, code: int, data: dict): body = json.dumps(data, ensure_ascii=False).encode() self.send_response(code) self.send_header("Content-Type", "application/json") self.send_header("Content-Length", str(len(body))) self.end_headers() self.wfile.write(body) def do_GET(self): if self.path == "/health": self._respond(200, {"status": "ok", "model": "bge-m3", "backend": "onnxruntime", "providers": session.get_providers()}) else: self._respond(404, {"error": "not found"}) def do_POST(self): if self.path != "/v1/embeddings": self._respond(404, {"error": "not found"}) return content_len = int(self.headers.get("Content-Length", 0)) body = json.loads(self.rfile.read(content_len)) inputs = body.get("input", []) if isinstance(inputs, str): inputs = [inputs] if not inputs: self._respond(400, {"error": "empty input"}) return if len(inputs) > MAX_BATCH: self._respond(400, {"error": f"batch size {len(inputs)} > max {MAX_BATCH}"}) return try: embeddings = encode(inputs) except Exception as e: log.error("encode error: %s", e) self._respond(500, {"error": str(e)}) return data = [ {"embedding": emb, "index": i, "object": "embedding"} for i, emb in enumerate(embeddings) ] self._respond(200, { "object": "list", "data": data, "model": "bge-m3", "usage": {"prompt_tokens": sum(len(t) for t in inputs), "total_tokens": sum(len(t) for t in inputs)}, }) def main(): server = HTTPServer(("0.0.0.0", PORT), EmbedHandler) log.info("bge-m3 ONNX 嵌入服务器启动 — http://0.0.0.0:%d", PORT) log.info("端点: POST /v1/embeddings GET /health") try: server.serve_forever() except KeyboardInterrupt: log.info("关闭服务器") server.shutdown() if __name__ == "__main__": main()