| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126 |
- #!/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()
|