#!/usr/bin/env python3 import json, os, sys, numpy as np import onnxruntime as ort from http.server import BaseHTTPRequestHandler, HTTPServer # Load model at startup session = ort.InferenceSession("/app/model.onnx") # Load vocab from vocab.txt vocab = {} with open("/app/vocab.txt", "r", encoding="utf-8") as f: for idx, token in enumerate(f): token = token.strip() vocab[token] = idx # Token IDs from tokenizer_config.json with open("/app/tokenizer_config.json", "r") as f: tokenizer_config = json.load(f) # Map token strings to IDs using vocab cls_token_id = vocab.get(tokenizer_config["cls_token"], 0) sep_token_id = vocab.get(tokenizer_config["sep_token"], 0) pad_token_id = vocab.get(tokenizer_config["pad_token"], 0) unk_token_id = vocab.get(tokenizer_config["unk_token"], 0) def wordpiece_tokenize(text): text = text.lower() tokens = [] buffer = "" for char in text: if char.isspace(): if buffer: tokens.append(buffer if buffer in vocab else "[UNK]") buffer = "" else: buffer += char if buffer: tokens.append(buffer if buffer in vocab else "[UNK]") return tokens def encode(text, max_len=128): token_ids = [vocab.get(t, unk_token_id) for t in wordpiece_tokenize(text)] input_ids = [cls_token_id] + token_ids + [sep_token_id] if len(input_ids) > max_len: input_ids = input_ids[:max_len] else: input_ids = input_ids + [pad_token_id] * (max_len - len(input_ids)) attention_mask = [1] * len(input_ids) token_type_ids = [0] * len(input_ids) return input_ids, attention_mask, token_type_ids print("Model loaded. Starting server on port 8080...") API_SECRET = os.getenv("API_SECRET") class VectorHandler(BaseHTTPRequestHandler): def do_POST(self): if self.path != "/vector": self.send_error(404) return # Check API secret if API_SECRET: auth_header = self.headers.get("Authorization", "") if auth_header != f"Bearer {API_SECRET}": self.send_error(401, "Unauthorized - Invalid or missing API key") return content_length = int(self.headers.get("Content-Length", 0)) body = self.rfile.read(content_length) try: req = json.loads(body) text = req.get("text", "") if not text: self.send_error(400, "Text is required") return input_ids, attention_mask, token_type_ids = encode(text) input_ids_np = np.array([input_ids], dtype=np.int64) attention_mask_np = np.array([attention_mask], dtype=np.int64) token_type_ids_np = np.array([token_type_ids], dtype=np.int64) outputs = session.run( ["last_hidden_state"], {"input_ids": input_ids_np, "attention_mask": attention_mask_np, "token_type_ids": token_type_ids_np} ) token_embeddings = outputs[0][0] input_mask_expanded = np.expand_dims(attention_mask_np[0], axis=-1) sum_embeddings = np.sum(token_embeddings * input_mask_expanded, axis=0) sum_mask = np.maximum(np.sum(input_mask_expanded, axis=0), 1e-9) sentence_embedding = (sum_embeddings / sum_mask).tolist() resp = json.dumps({"vector": sentence_embedding}) self.send_response(200) self.send_header("Content-Type", "application/json") self.end_headers() self.wfile.write(resp.encode()) except Exception as e: self.send_error(500, str(e)) def log_message(self, format, *args): pass # Suppress logs HTTPServer(("0.0.0.0", int(os.getenv("PORT", 8080))), VectorHandler).serve_forever()