"""Dense retrieval for in-context examples: embed everything with a transformer encoder (mean-pooled), then for each query pull the nearest neighbors while keeping the three classes balanced. """ import json import urllib.request import numpy as np import torch from transformers import (AutoModel, AutoModelForSequenceClassification, AutoTokenizer) LABELS = ["Against", "Favor", "None"] def embed_texts_endpoint(texts, base_url, model, batch_size=64, instruction=None, timeout=180): """Embed via an OpenAI-compatible /v1/embeddings server. When ``instruction`` is set it is prepended to each text in the ``Instruct: ...\\nQuery: ...`` form expected by Qwen3-Embedding; leave it None for symmetric similarity between texts of the same kind. """ def fmt(t): if instruction: return f"Instruct: {instruction}\nQuery: {t}" return t url = base_url.rstrip("/") + "/embeddings" out = [] for i in range(0, len(texts), batch_size): batch = [fmt(t) for t in texts[i:i + batch_size]] body = json.dumps({"model": model, "input": batch}).encode("utf-8") req = urllib.request.Request( url, data=body, headers={"Content-Type": "application/json"}) with urllib.request.urlopen(req, timeout=timeout) as r: data = json.load(r) rows = sorted(data["data"], key=lambda d: d["index"]) vecs = np.asarray([d["embedding"] for d in rows], dtype=np.float32) out.append(vecs) emb = np.concatenate(out, axis=0) norms = np.linalg.norm(emb, axis=1, keepdims=True) return emb / np.clip(norms, 1e-9, None) class Reranker: """CrossEncoder that rescores (query, doc) pairs by relevance.""" def __init__(self, model_name, device): self.tok = AutoTokenizer.from_pretrained(model_name) self.model = AutoModelForSequenceClassification.from_pretrained( model_name ).to(device).eval() self.device = device @torch.no_grad() def score(self, query, docs, max_len=256, batch_size=32): scores = [] for i in range(0, len(docs), batch_size): batch = docs[i:i + batch_size] enc = self.tok([query] * len(batch), batch, truncation=True, padding=True, max_length=max_len, return_tensors="pt").to(self.device) logits = self.model(**enc).logits.float() s = (logits.squeeze(-1) if logits.shape[-1] == 1 else logits[:, -1]) scores.extend(s.cpu().numpy().tolist()) return scores @torch.no_grad() def embed_texts(model_name, texts, device, batch_size=64, max_len=128): tok = AutoTokenizer.from_pretrained(model_name, trust_remote_code=True) model = AutoModel.from_pretrained( model_name, trust_remote_code=True ).to(device).eval() out = [] for i in range(0, len(texts), batch_size): batch = texts[i:i + batch_size] enc = tok(batch, truncation=True, padding=True, max_length=max_len, return_tensors="pt").to(device) hidden = model(**enc).last_hidden_state mask = enc["attention_mask"].unsqueeze(-1).float() pooled = (hidden * mask).sum(1) / mask.sum(1).clamp(min=1e-9) pooled = torch.nn.functional.normalize(pooled, dim=-1) out.append(pooled.cpu().numpy()) return np.concatenate(out, axis=0) class Retriever: def __init__(self, train_df, model_name, device, embed_url=None, embed_api_model=None, instruction=None): self.texts = train_df["text"].tolist() self.labels = train_df["stance"].tolist() self.embed_url = embed_url self.embed_api_model = embed_api_model self.instruction = instruction if embed_url: self.emb = embed_texts_endpoint( self.texts, embed_url, embed_api_model) else: self.emb = embed_texts(model_name, self.texts, device) self.by_class = {lb: np.array( [i for i, x in enumerate(self.labels) if x == lb] ) for lb in LABELS} def embed_queries(self, texts, model_name, device): if self.embed_url: return embed_texts_endpoint( texts, self.embed_url, self.embed_api_model, instruction=self.instruction) return embed_texts(model_name, texts, device) def balanced_shots(self, q_emb, k, query_text=None, reranker=None, pool_m=10, exclude_text=None, sample_m=0, rng=None): sims = self.emb @ q_emb order = {lb: idx[np.argsort(-sims[idx])] for lb, idx in self.by_class.items() if len(idx)} if exclude_text is not None: order = {lb: idx[[self.texts[i] != exclude_text for i in idx]] for lb, idx in order.items()} order = {lb: idx for lb, idx in order.items() if len(idx)} if sample_m and rng is not None: shuffled = {} for lb, idx in order.items(): top = idx[:sample_m].copy() rng.shuffle(top) shuffled[lb] = top order = shuffled if reranker is not None and query_text is not None: reordered = {} for lb, idx in order.items(): cand = idx[:pool_m] rs = reranker.score(query_text, [self.texts[i] for i in cand]) reordered[lb] = cand[np.argsort(-np.asarray(rs))] order = reordered pos = {lb: 0 for lb in order} shots = [] while len(shots) < k and any( pos[lb] < len(order[lb]) for lb in order ): for lb in LABELS: if lb in order and pos[lb] < len(order[lb]) and len(shots) < k: i = order[lb][pos[lb]] pos[lb] += 1 shots.append((self.texts[i], lb)) return shots