1- """Embedding sidecar service .
1+ """Embedding sidecar — ONNX Runtime backend .
22
3- Provides a single endpoint:
4- POST /embed {"texts": ["...", "..."]} -> {"embeddings": [[...], [...]]}
5-
6- Uses sentence-transformers all-MiniLM-L6-v2 (384-dim, ~90MB).
7- Falls back to the stub embedder (SHA-256 tiling) if the model is unavailable,
8- so the validator can start without the model downloaded.
3+ Uses all-MiniLM-L6-v2 (quantized, 22MB) via onnxruntime + tokenizers.
4+ No torch, no sentence-transformers, no internet access at runtime.
95
10- Run :
11- pip install -r requirements.txt
12- uvicorn embed.main:app --host 0.0.0.0 --port 8000
6+ Endpoints :
7+ POST /embed {"texts": ["...", "..."]} -> {"embeddings": [[...], [...]]}
8+ GET /health
139"""
1410from __future__ import annotations
1511
1612import hashlib
1713import logging
1814import os
15+ from pathlib import Path
1916from typing import List
2017
2118import numpy as np
2421
2522logger = logging .getLogger (__name__ )
2623
27- app = FastAPI (title = "AgentsProtocol Embed Service" , version = "1.0" )
24+ app = FastAPI (title = "AgentsProtocol Embed Service" , version = "2.0" )
25+
26+ MODEL_DIR = Path (os .getenv ("MODEL_DIR" , "/app/model" ))
2827
29- # ---------------------------------------------------------------------------
30- # Model loading — lazy, with stub fallback
31- # ---------------------------------------------------------------------------
28+ _session = None
29+ _tokenizer = None
3230
33- _model = None
34- _use_stub = False
3531
36- def _load_model () -> None :
37- global _model , _use_stub
38- model_name = os .getenv ("EMBED_MODEL" , "all-MiniLM-L6-v2" )
32+ def _load () -> None :
33+ global _session , _tokenizer
3934 try :
40- from sentence_transformers import SentenceTransformer
41- _model = SentenceTransformer (model_name )
42- logger .info ("Loaded embedding model: %s" , model_name )
35+ import onnxruntime as ort
36+ from tokenizers import Tokenizer
37+
38+ opts = ort .SessionOptions ()
39+ opts .inter_op_num_threads = 1
40+ opts .intra_op_num_threads = 1
41+ _session = ort .InferenceSession (
42+ str (MODEL_DIR / "model.onnx" ),
43+ sess_options = opts ,
44+ providers = ["CPUExecutionProvider" ],
45+ )
46+ _tokenizer = Tokenizer .from_file (str (MODEL_DIR / "tokenizer.json" ))
47+ _tokenizer .enable_padding (pad_id = 0 , pad_token = "[PAD]" , length = 128 )
48+ _tokenizer .enable_truncation (max_length = 128 )
49+ logger .info ("ONNX model loaded from %s" , MODEL_DIR )
4350 except Exception as exc :
44- logger .warning ("Could not load %s (%s) — using stub embedder" , model_name , exc )
45- _use_stub = True
51+ logger .warning ("Could not load ONNX model (%s) — using stub" , exc )
4652
4753
4854@app .on_event ("startup" )
4955async def startup () -> None :
50- _load_model ()
56+ _load ()
5157
5258
53- # ---------------------------------------------------------------------------
54- # Stub embedder (mirrors Rust stub_embed / Python _stub_embed exactly)
55- # ---------------------------------------------------------------------------
59+ def _mean_pool (token_embeddings : np .ndarray , attention_mask : np .ndarray ) -> np .ndarray :
60+ mask = attention_mask [..., np .newaxis ].astype (float )
61+ summed = (token_embeddings * mask ).sum (axis = 1 )
62+ counts = mask .sum (axis = 1 ).clip (min = 1e-9 )
63+ pooled = summed / counts
64+ norms = np .linalg .norm (pooled , axis = 1 , keepdims = True ).clip (min = 1e-9 )
65+ return (pooled / norms ).astype (np .float32 )
66+
5667
5768def _stub_embed (text : str ) -> List [float ]:
5869 DIM = 384
5970 digest = hashlib .sha256 (text .encode ()).digest ()
6071 raw = bytes (digest [i % 32 ] for i in range (DIM ))
6172 vec = np .array ([b - 127.5 for b in raw ], dtype = np .float64 )
6273 norm = np .linalg .norm (vec )
63- if norm > 0 :
64- vec /= norm
65- return vec .tolist ()
66-
74+ return (vec / norm if norm > 0 else vec ).tolist ()
75+
76+
77+ def _onnx_embed (texts : List [str ]) -> List [List [float ]]:
78+ enc = _tokenizer .encode_batch (texts )
79+ input_ids = np .array ([e .ids for e in enc ], dtype = np .int64 )
80+ attention_mask = np .array ([e .attention_mask for e in enc ], dtype = np .int64 )
81+ token_type_ids = np .zeros_like (input_ids )
82+
83+ outputs = _session .run (
84+ None ,
85+ {
86+ "input_ids" : input_ids ,
87+ "attention_mask" : attention_mask ,
88+ "token_type_ids" : token_type_ids ,
89+ },
90+ )
91+ # outputs[0] = last_hidden_state (batch, seq, 384)
92+ pooled = _mean_pool (outputs [0 ], attention_mask )
93+ return pooled .tolist ()
6794
68- # ---------------------------------------------------------------------------
69- # API
70- # ---------------------------------------------------------------------------
7195
7296class EmbedRequest (BaseModel ):
7397 texts : List [str ]
@@ -80,17 +104,17 @@ class EmbedResponse(BaseModel):
80104
81105@app .post ("/embed" , response_model = EmbedResponse )
82106async def embed (req : EmbedRequest ) -> EmbedResponse :
83- if _use_stub or _model is None :
84- embeddings = [ _stub_embed ( t ) for t in req . texts ]
85- return EmbedResponse ( embeddings = embeddings , model = "stub" )
86-
87- vecs = _model . encode ( req . texts , normalize_embeddings = True )
107+ if _session is None or _tokenizer is None :
108+ return EmbedResponse (
109+ embeddings = [ _stub_embed ( t ) for t in req . texts ],
110+ model = "stub" ,
111+ )
88112 return EmbedResponse (
89- embeddings = vecs . tolist ( ),
90- model = _model . get_sentence_embedding_dimension () and "all-MiniLM-L6-v2" ,
113+ embeddings = _onnx_embed ( req . texts ),
114+ model = "all-MiniLM-L6-v2-qint8 " ,
91115 )
92116
93117
94118@app .get ("/health" )
95119async def health () -> dict :
96- return {"status" : "ok" , "stub " : _use_stub }
120+ return {"status" : "ok" , "model_loaded " : _session is not None }
0 commit comments