#!/usr/bin/env python3 """ BGE-M3 Embedding Server on Ascend NPU Lightweight HTTP server compatible with Ollama /api/embed format. Runs on a single NPU, batch processing for high throughput. Usage: python embed_server_npu.py [--port 11435] [--device 0] """ import json import time import argparse from http.server import HTTPServer, BaseHTTPRequestHandler import torch try: import torch_npu HAS_NPU = True except ImportError: HAS_NPU = False from transformers import AutoTokenizer, AutoModel MODEL_PATH = None TOKENIZER = None MODEL = None DEVICE = None def load_model(model_path, device_id=0): global TOKENIZER, MODEL, DEVICE print(f"Loading bge-m3 from {model_path}...") t0 = time.time() TOKENIZER = AutoTokenizer.from_pretrained(model_path, trust_remote_code=True) if HAS_NPU: DEVICE = torch.device(f"npu:{device_id}") torch.npu.set_device(DEVICE) MODEL = AutoModel.from_pretrained(model_path, trust_remote_code=True).half().to(DEVICE) print(f" Loaded on NPU:{device_id} in {time.time()-t0:.1f}s") else: DEVICE = torch.device("cpu") MODEL = AutoModel.from_pretrained(model_path, trust_remote_code=True) print(f" Loaded on CPU in {time.time()-t0:.1f}s") MODEL.eval() def embed_texts(texts, max_length=512): """Embed a batch of texts, return list of vectors""" with torch.no_grad(): encoded = TOKENIZER( texts, padding=True, truncation=True, max_length=max_length, return_tensors="pt" ) encoded = {k: v.to(DEVICE) for k, v in encoded.items()} outputs = MODEL(**encoded) # Use [CLS] token embedding embeddings = outputs.last_hidden_state[:, 0, :] # Normalize embeddings = torch.nn.functional.normalize(embeddings, p=2, dim=1) return embeddings.cpu().float().tolist() class EmbedHandler(BaseHTTPRequestHandler): def do_POST(self): if self.path == '/api/embed': content_length = int(self.headers['Content-Length']) body = json.loads(self.rfile.read(content_length)) texts = body.get('input', []) if isinstance(texts, str): texts = [texts] t0 = time.time() embeddings = embed_texts(texts) elapsed = time.time() - t0 response = { "embeddings": embeddings, "model": "bge-m3-npu", "elapsed": round(elapsed, 3), } self.send_response(200) self.send_header('Content-Type', 'application/json') self.end_headers() self.wfile.write(json.dumps(response).encode('utf-8')) else: self.send_response(404) self.end_headers() def log_message(self, format, *args): # Suppress default logging for speed pass def main(): parser = argparse.ArgumentParser() parser.add_argument('--model', default='/home/data/hf_models/BAAI/bge-m3') parser.add_argument('--port', type=int, default=11435) parser.add_argument('--device', type=int, default=0) args = parser.parse_args() load_model(args.model, args.device) # Warmup print("Warmup...") t0 = time.time() embed_texts(["warmup test"]) print(f" Warmup: {time.time()-t0:.3f}s") # Benchmark batch t0 = time.time() embed_texts(["test1", "test2", "test3", "test4", "test5", "test6", "test7", "test8"]) print(f" Batch 8: {time.time()-t0:.3f}s ({8/(time.time()-t0):.0f} texts/s)") server = HTTPServer(('0.0.0.0', args.port), EmbedHandler) print(f"\nEmbedding server on http://0.0.0.0:{args.port}/api/embed") print(f"Device: {'NPU:' + str(args.device) if HAS_NPU else 'CPU'}") server.serve_forever() if __name__ == '__main__': main()