xiaowei-system/scripts/bge_embed_server.py

145 lines
4.7 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""织忆 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")
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/model.onnx", MODEL_PATH)
sess_options = ort.SessionOptions()
sess_options.intra_op_num_threads = 4
sess_options.inter_op_num_threads = 2
session = ort.InferenceSession(
os.path.join(MODEL_PATH, "model.onnx"),
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"})
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()