File size: 6,032 Bytes
7e9cfd1
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
"""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