Initial commit: Boltzmann LLM proxy with compression + classifier middleware
- llm-proxy.py v4: aggregator with classify/compress endpoints, gemma4 backend - classifier-server.py: Llama 3.2 1B query complexity classifier - compress-server.py: token-budget compression middleware - start-proxy.sh / start.sh: launcher scripts
This commit is contained in:
@@ -0,0 +1,176 @@
|
||||
#!/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()
|
||||
Reference in New Issue
Block a user