xiaowei-system/scripts/bge_embed_server.py

176 lines
6.1 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")
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"))
# 2026-09-09: GPU 化 — CUDA 优先OOM/不可用时 ORT 自动或显式降级 CPU
PROVIDERS = os.environ.get("BGE_PROVIDERS", "CUDAExecutionProvider,CPUExecutionProvider").split(",")
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 # 禁用内存模式,避免碎片累积
def _make_session(providers: list[str]):
"""创建 ORT sessionCUDA 缺库/不可用时自动降级 CPUORT 会按列表尝试)"""
return ort.InferenceSession(
os.path.join(MODEL_PATH, MODEL_FILE),
sess_options=sess_options,
providers=providers,
)
_cpu_session = None
try:
session = _make_session(PROVIDERS) # CUDA 缺库时 ORT 自动降级 CPU实测返回 CPU
except Exception as e:
log.warning("首选 providers %s 初始化失败(%s),降级 CPU", PROVIDERS, e)
session = _make_session(["CPUExecutionProvider"])
log.info("ONNX 模型就绪 — providers=%s", session.get_providers())
def _get_cpu_session():
"""懒创建 CPU sessionCUDA 运行期 OOM 降级用)"""
global _cpu_session
if _cpu_session is None:
_cpu_session = _make_session(["CPUExecutionProvider"])
log.info("CPU 降级 session 已创建")
return _cpu_session
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"],
}
try:
outputs = session.run(None, ort_inputs)
except Exception as e:
# CUDA 运行期 OOM / 显存不足 → 降级 CPU 重试(防 4GB 显存撞 llama 崩循环)
log.warning("GPU 推理失败(%s),降级 CPU", e)
outputs = _get_cpu_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()