stanceeval2026 / code /src /retrieve.py
zaher-m's picture
Add files using upload-large-folder tool
7e9cfd1 verified
Raw
History Blame Contribute Delete
6.03 kB
"""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