#!/usr/bin/env python3 """Prompt compression service using XLM-RoBERTa-large (LLMLingua-2). Scoring approach: runs the classification head on each token, drops tokens with low keep-probability. The compressed text is reconstructed from the kept token spans. Endpoints: POST /compress {"messages": [...], "rate": 0.5} → {"compressed": [...], "stats": {...}} GET /health → {"status": "ok"} """ import http.server import json import os import sys import time import re import math import torch import numpy as np from transformers import AutoTokenizer, AutoModelForTokenClassification MODEL_NAME = os.getenv("COMPRESS_MODEL", "microsoft/llmlingua-2-xlm-roberta-large-meetingbank") LISTEN_PORT = int(os.getenv("COMPRESS_LISTEN_PORT", "8091")) DEVICE = "cpu" DTYPE = torch.float16 def log(msg): sys.stderr.write(f"{time.strftime('%H:%M:%S')} {msg}\n") sys.stderr.flush() # ── model loading ───────────────────────────────────────────────── log(f"loading model {MODEL_NAME} ...") t0 = time.time() tokenizer = AutoTokenizer.from_pretrained(MODEL_NAME) model = AutoModelForTokenClassification.from_pretrained( MODEL_NAME, torch_dtype=DTYPE ).to(DEVICE).eval() log(f"model loaded in {time.time()-t0:.1f}s ({sum(p.numel() for p in model.parameters())/1e6:.0f}M params)") # ── compression core ────────────────────────────────────────────── def _merge_token_spans(text, keep_mask, tokens): """Reconstruct text from kept tokens, merging subwords cleanly.""" kept = [] for i, (tok, keep) in enumerate(zip(tokens, keep_mask)): if not keep: continue piece = tok # Remove prefix space indicator for merged tokens if i > 0 and piece.startswith("##"): piece = piece[2:] elif i > 0 and not piece.startswith(" ") and not re.match(r'^[^\w]', piece): piece = " " + piece kept.append(piece) return "".join(kept).strip() def compress_text(text, rate=0.5): """Compress a single text string, targeting `rate` compression ratio (0-1).""" inputs = tokenizer(text, return_tensors="pt", truncation=True, max_length=2048).to(DEVICE) input_ids = inputs["input_ids"][0] n_tokens = len(input_ids) with torch.no_grad(): outputs = model(**inputs) logits = outputs.logits[0] # (seq_len, num_labels) # For binary classification: label 0 = keep, 1 = drop keep_probs = torch.softmax(logits, dim=-1)[:, 0] # probability of keep probs = keep_probs.cpu().numpy() # Don't drop special tokens ([CLS], [SEP], [PAD]) special_ids = {tokenizer.cls_token_id, tokenizer.sep_token_id, tokenizer.pad_token_id, tokenizer.bos_token_id, tokenizer.eos_token_id, tokenizer.unk_token_id, 0} # padding is_special = [id.item() in special_ids for id in input_ids] # Target: keep rate * (1-rate) tokens (rate=0.5 means keep half) n_to_keep = max(1, int(n_tokens * (1 - rate))) # Mask: force-keep special tokens, then keep top-k by probability keep = np.zeros(n_tokens, dtype=bool) for i in range(n_tokens): if is_special[i]: keep[i] = True n_kept_special = keep.sum() n_remaining = n_to_keep - n_kept_special if n_remaining > 0: # Get indices of non-special tokens sorted by keep probability non_special_idx = [i for i in range(n_tokens) if not is_special[i]] sorted_idx = sorted(non_special_idx, key=lambda i: probs[i], reverse=True) for i in sorted_idx[:n_remaining]: keep[i] = True # Decode kept tokens tokens = tokenizer.convert_ids_to_tokens(input_ids) compressed = _merge_token_spans(text, keep, tokens) return { "compressed": compressed, "original_tokens": int(n_tokens), "compressed_tokens": int(keep.sum()), "ratio": float(keep.sum() / n_tokens), } def compress_messages(messages, rate=0.5): """Compress an array of chat messages. Drops low-info tokens from each.""" compressed = [] for msg in messages: role = msg.get("role", "user") content = msg.get("content", "") if content: result = compress_text(content, rate) compressed.append({"role": role, "content": result["compressed"]}) else: compressed.append(msg) return compressed # ── HTTP handler ───────────────────────────────────────────────── class Handler(http.server.BaseHTTPRequestHandler): server_version = "compress-server/1" def _read_body(self): n = int(self.headers.get("Content-Length", 0)) return json.loads(self.rfile.read(n)) if n else {} def _json_response(self, code, data): payload = json.dumps(data, default=str).encode() self.send_response(code) self.send_header("Content-Type", "application/json") self.send_header("Content-Length", str(len(payload))) self.end_headers() self.wfile.write(payload) def do_GET(self): if self.path == "/health": return self._json_response(200, {"status": "ok"}) self._json_response(404, {"error": "not found"}) def do_POST(self): if self.path == "/compress": body = self._read_body() messages = body.get("messages", []) rate = float(body.get("rate", 0.5)) if not messages: return self._json_response(400, {"error": "missing messages"}) t0 = time.time() compressed = compress_messages(messages, rate) elapsed = time.time() - t0 log(f"compress {len(messages)} msgs rate={rate} in {elapsed:.2f}s") return self._json_response(200, { "compressed": compressed, "rate": rate, "elapsed_s": round(elapsed, 2), }) self._json_response(404, {"error": "not found"}) def log_message(self, fmt, *args): return if __name__ == "__main__": log(f"compress-server on :{LISTEN_PORT}") http.server.ThreadingHTTPServer(("0.0.0.0", LISTEN_PORT), Handler).serve_forever()