embed_server_npu.py 3.8 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126
  1. #!/usr/bin/env python3
  2. """
  3. BGE-M3 Embedding Server on Ascend NPU
  4. Lightweight HTTP server compatible with Ollama /api/embed format.
  5. Runs on a single NPU, batch processing for high throughput.
  6. Usage:
  7. python embed_server_npu.py [--port 11435] [--device 0]
  8. """
  9. import json
  10. import time
  11. import argparse
  12. from http.server import HTTPServer, BaseHTTPRequestHandler
  13. import torch
  14. try:
  15. import torch_npu
  16. HAS_NPU = True
  17. except ImportError:
  18. HAS_NPU = False
  19. from transformers import AutoTokenizer, AutoModel
  20. MODEL_PATH = None
  21. TOKENIZER = None
  22. MODEL = None
  23. DEVICE = None
  24. def load_model(model_path, device_id=0):
  25. global TOKENIZER, MODEL, DEVICE
  26. print(f"Loading bge-m3 from {model_path}...")
  27. t0 = time.time()
  28. TOKENIZER = AutoTokenizer.from_pretrained(model_path, trust_remote_code=True)
  29. if HAS_NPU:
  30. DEVICE = torch.device(f"npu:{device_id}")
  31. torch.npu.set_device(DEVICE)
  32. MODEL = AutoModel.from_pretrained(model_path, trust_remote_code=True).half().to(DEVICE)
  33. print(f" Loaded on NPU:{device_id} in {time.time()-t0:.1f}s")
  34. else:
  35. DEVICE = torch.device("cpu")
  36. MODEL = AutoModel.from_pretrained(model_path, trust_remote_code=True)
  37. print(f" Loaded on CPU in {time.time()-t0:.1f}s")
  38. MODEL.eval()
  39. def embed_texts(texts, max_length=512):
  40. """Embed a batch of texts, return list of vectors"""
  41. with torch.no_grad():
  42. encoded = TOKENIZER(
  43. texts, padding=True, truncation=True,
  44. max_length=max_length, return_tensors="pt"
  45. )
  46. encoded = {k: v.to(DEVICE) for k, v in encoded.items()}
  47. outputs = MODEL(**encoded)
  48. # Use [CLS] token embedding
  49. embeddings = outputs.last_hidden_state[:, 0, :]
  50. # Normalize
  51. embeddings = torch.nn.functional.normalize(embeddings, p=2, dim=1)
  52. return embeddings.cpu().float().tolist()
  53. class EmbedHandler(BaseHTTPRequestHandler):
  54. def do_POST(self):
  55. if self.path == '/api/embed':
  56. content_length = int(self.headers['Content-Length'])
  57. body = json.loads(self.rfile.read(content_length))
  58. texts = body.get('input', [])
  59. if isinstance(texts, str):
  60. texts = [texts]
  61. t0 = time.time()
  62. embeddings = embed_texts(texts)
  63. elapsed = time.time() - t0
  64. response = {
  65. "embeddings": embeddings,
  66. "model": "bge-m3-npu",
  67. "elapsed": round(elapsed, 3),
  68. }
  69. self.send_response(200)
  70. self.send_header('Content-Type', 'application/json')
  71. self.end_headers()
  72. self.wfile.write(json.dumps(response).encode('utf-8'))
  73. else:
  74. self.send_response(404)
  75. self.end_headers()
  76. def log_message(self, format, *args):
  77. # Suppress default logging for speed
  78. pass
  79. def main():
  80. parser = argparse.ArgumentParser()
  81. parser.add_argument('--model', default='/home/data/hf_models/BAAI/bge-m3')
  82. parser.add_argument('--port', type=int, default=11435)
  83. parser.add_argument('--device', type=int, default=0)
  84. args = parser.parse_args()
  85. load_model(args.model, args.device)
  86. # Warmup
  87. print("Warmup...")
  88. t0 = time.time()
  89. embed_texts(["warmup test"])
  90. print(f" Warmup: {time.time()-t0:.3f}s")
  91. # Benchmark batch
  92. t0 = time.time()
  93. embed_texts(["test1", "test2", "test3", "test4", "test5", "test6", "test7", "test8"])
  94. print(f" Batch 8: {time.time()-t0:.3f}s ({8/(time.time()-t0):.0f} texts/s)")
  95. server = HTTPServer(('0.0.0.0', args.port), EmbedHandler)
  96. print(f"\nEmbedding server on http://0.0.0.0:{args.port}/api/embed")
  97. print(f"Device: {'NPU:' + str(args.device) if HAS_NPU else 'CPU'}")
  98. server.serve_forever()
  99. if __name__ == '__main__':
  100. main()