148 lines
5.0 KiB
Python
148 lines
5.0 KiB
Python
"""织忆 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=["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()
|