File size: 9,894 Bytes
2d8be88
14bf8c2
2d8be88
 
14bf8c2
2d8be88
 
 
 
 
 
 
 
 
 
 
 
 
14bf8c2
2d8be88
 
 
14bf8c2
 
 
 
 
 
 
2d8be88
 
14bf8c2
2d8be88
14bf8c2
 
 
 
2d8be88
14bf8c2
 
2d8be88
 
14bf8c2
 
 
 
 
 
 
 
 
 
2d8be88
 
14bf8c2
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
2d8be88
14bf8c2
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
2d8be88
 
14bf8c2
2d8be88
14bf8c2
 
 
2d8be88
14bf8c2
 
2d8be88
14bf8c2
 
2d8be88
 
 
14bf8c2
 
 
2d8be88
 
14bf8c2
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
2d8be88
 
14bf8c2
 
 
 
 
 
 
 
 
 
2d8be88
 
14bf8c2
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
from __future__ import annotations
import hashlib, json, random
from dataclasses import dataclass, field
from pathlib import Path
from typing import Any, Optional, Sequence
import torch
from torch.utils.data import Dataset

TASK_TYPES = {"choice", "noul", "score"}

@dataclass
class DecisionExample:
    context: str
    options: list[str]
    label: int
    task_type: str
    meta: dict[str, Any] = field(default_factory=dict)
    @property
    def gold_text(self) -> str: return self.options[self.label]
    @property
    def source_family(self) -> str:
        fam = self.meta.get("source_family")
        if fam: return str(fam)
        raw = json.dumps([self.context, self.options, self.label], sort_keys=True, ensure_ascii=True)
        return "singleton::" + hashlib.sha256(raw.encode()).hexdigest()[:20]
    @property
    def pair_id(self) -> Optional[str]:
        v = self.meta.get("pair_id")
        return None if v in (None, "") else str(v)
    @property
    def hard_negatives(self) -> list[str]:
        return [str(x) for x in (self.meta.get("hard_negatives") or [])]

def parse_record(obj: dict[str, Any], noul_true="true", noul_false="false") -> DecisionExample:
    if "context" not in obj or "label" not in obj: raise ValueError("record requires context and label")
    task = str(obj.get("task_type", "choice"))
    if task not in TASK_TYPES: raise ValueError(f"bad task_type {task}")
    options = obj.get("options")
    if not options and task == "noul": options = [noul_true, noul_false]
    if not options: raise ValueError("record requires options")
    options = [str(x) for x in options]
    label = int(obj["label"])
    if not 0 <= label < len(options): raise ValueError("label out of range")
    return DecisionExample(str(obj["context"]), options, label, task, dict(obj.get("meta") or {}))

def load_jsonl(path: str | Path, noul_true="true", noul_false="false") -> list[DecisionExample]:
    rows=[]
    with Path(path).open(encoding="utf-8") as f:
        for i,line in enumerate(f,1):
            if not line.strip(): continue
            try: rows.append(parse_record(json.loads(line), noul_true, noul_false))
            except Exception as e: raise ValueError(f"{path}:{i}: {e}") from e
    return rows

def canonical_row_hash(ex: DecisionExample) -> str:
    obj={"context":ex.context,"options":ex.options,"label":ex.label,"task_type":ex.task_type}
    return hashlib.sha256(json.dumps(obj,sort_keys=True,ensure_ascii=False,separators=(",",":")).encode()).hexdigest()

def overlap_report(a: Sequence[DecisionExample], b: Sequence[DecisionExample]) -> dict[str, Any]:
    ah={canonical_row_hash(x) for x in a}; bh={canonical_row_hash(x) for x in b}; inter=ah & bh
    return {"a_unique":len(ah),"b_unique":len(bh),"overlap_unique":len(inter),"a_fraction":len(inter)/max(1,len(ah)),"b_fraction":len(inter)/max(1,len(bh))}

def _group_id(ex: DecisionExample, index: int) -> str:
    # Exact duplicate rows must stay in one partition too; otherwise duplicate
    # copies can leak from optimizer training into validation/holdout.
    del index
    return f"pair::{ex.pair_id}" if ex.pair_id else f"row::{canonical_row_hash(ex)}"

def stratified_group_split(
    examples: Sequence[DecisionExample], *, seed: int, ratios: Sequence[float]=(0.8,0.1,0.1)
) -> tuple[list[DecisionExample],list[DecisionExample],list[DecisionExample]]:
    """Split within each source family while never breaking pair_id groups."""
    if len(ratios)!=3 or abs(sum(ratios)-1)>1e-6: raise ValueError("ratios must sum to 1")
    fams: dict[str,list[tuple[int,DecisionExample]]] = {}
    for i,ex in enumerate(examples):
        fam = ex.source_family
        stratum = "__untagged__" if fam.startswith("singleton::") else fam
        fams.setdefault(stratum,[]).append((i,ex))
    outs=[[],[],[]]
    for fam, rows in sorted(fams.items()):
        groups: dict[str,list[DecisionExample]]={}
        for i,ex in rows: groups.setdefault(_group_id(ex,i),[]).append(ex)
        items=list(groups.items())
        rng=random.Random(seed ^ int(hashlib.sha256(fam.encode()).hexdigest()[:8],16))
        rng.shuffle(items)
        if len(items)==1:
            alloc=[0]
        else:
            targets=[len(rows)*r for r in ratios]; counts=[0,0,0]; alloc=[]
            seeded=[0] + ([1,2] if len(items)>=3 else [])
            for j,(_,g) in enumerate(items):
                if j < len(seeded): s=seeded[j]
                else:
                    deficits=[targets[x]-counts[x] for x in range(3)]
                    s=max(range(3), key=lambda x:deficits[x])
                alloc.append(s); counts[s]+=len(g)
        for (_,g),s in zip(items,alloc): outs[s].extend(g)
    if not outs[0] or not outs[1] or not outs[2]:
        raise ValueError(f"split produced empty partition: {[len(x) for x in outs]}")
    return tuple(outs)  # type: ignore

def family_counts(examples: Sequence[DecisionExample]) -> dict[str,int]:
    out={}
    for ex in examples: out[ex.source_family]=out.get(ex.source_family,0)+1
    return out

def compute_family_sample_weights(examples: Sequence[DecisionExample], caps: dict[str,float]|None) -> list[float]:
    n=len(examples); w=[1.0]*n
    if not caps or not n: return w
    counts=family_counts(examples)
    for fam,cap in sorted(((str(k),float(v)) for k,v in caps.items()), key=lambda x:x[1]):
        nf=counts.get(fam,0)
        if not nf or nf/n <= cap: continue
        outside=sum(w[i] for i,e in enumerate(examples) if e.source_family!=fam)
        wf=(cap*outside)/(nf*(1-cap))
        for i,e in enumerate(examples):
            if e.source_family==fam: w[i]=wf
    return w

def expected_family_mass(weights, examples, family):
    t=float(sum(weights))
    return 0.0 if t<=0 else sum(w for w,e in zip(weights,examples) if e.source_family==family)/t

def build_option_pool(examples: Sequence[DecisionExample]) -> list[str]:
    seen=set(); out=[]
    for ex in examples:
        for x in ex.options + ex.hard_negatives:
            if x not in seen: seen.add(x); out.append(x)
    return out or ["true","false","unknown","not applicable"]

def build_pair_gold_index(examples: Sequence[DecisionExample]) -> dict[str,list[str]]:
    out={}
    for ex in examples:
        if ex.pair_id: out.setdefault(ex.pair_id,[]).append(ex.gold_text)
    return out

@dataclass
class AugmentedSet:
    options:list[str]
    label:int
    k:int

class SetAugmenter:
    def __init__(self, *, k_min:int,k_max:int,pool:Sequence[str],pair_golds:Optional[dict[str,list[str]]]=None):
        if k_min<2 or k_max<k_min: raise ValueError("bad k range")
        self.k_min,self.k_max,self.pool,self.pair_golds=k_min,k_max,list(pool),pair_golds or {}
    def build(self, ex:DecisionExample, rng:random.Random, k:Optional[int]=None)->AugmentedSet:
        k=max(self.k_min,min(self.k_max,int(k if k is not None else rng.randint(self.k_min,self.k_max))))
        gold=ex.gold_text; chosen=[gold]; seen={gold}
        hard=list(ex.hard_negatives)
        if ex.pair_id: hard += self.pair_golds.get(ex.pair_id,[])
        for x in hard:
            if len(chosen)>=k: break
            if x not in seen: chosen.append(x); seen.add(x)
        rest=[x for x in self.pool if x not in seen]; rng.shuffle(rest)
        for x in rest:
            if len(chosen)>=k: break
            chosen.append(x); seen.add(x)
        j=0
        while len(chosen)<k:
            x=f"[distractor {j}]"; j+=1
            if x not in seen: chosen.append(x); seen.add(x)
        perm=list(range(k)); rng.shuffle(perm)
        return AugmentedSet([chosen[i] for i in perm], perm.index(0), k)

class DecisionDataset(Dataset):
    def __init__(self, examples, augmenter, *, seed:int, augment=True, fixed_k=None, epoch=0):
        self.examples=list(examples); self.augmenter=augmenter; self.seed=int(seed); self.augment=augment; self.fixed_k=fixed_k; self.epoch=epoch
    def set_epoch(self,e): self.epoch=int(e)
    def __len__(self): return len(self.examples)
    def __getitem__(self,i):
        ex=self.examples[i]; rng=random.Random(self.seed+1_000_003*(self.epoch+1)+i)
        if self.augment or self.fixed_k is not None:
            a=self.augmenter.build(ex,rng,k=self.fixed_k); opts,label=a.options,a.label
        else: opts,label=list(ex.options),ex.label
        return {"context":ex.context,"options":opts,"label":label,"k":len(opts)}

class DecisionCollator:
    def __init__(self, tokenizer, *, k_max:int,max_context_tokens:int,max_option_tokens:int,pad_context_to_max:bool=True):
        self.tokenizer=tokenizer; self.k_max=k_max; self.max_context_tokens=max_context_tokens; self.max_option_tokens=max_option_tokens
        self.pad_context_to_max=bool(pad_context_to_max)
    def _tok(self,texts,max_length,*,pad_to_max=True):
        return self.tokenizer(texts,padding="max_length" if pad_to_max else True,truncation=True,max_length=max_length,return_tensors="pt")
    def __call__(self,batch):
        ctx=self._tok([x["context"] for x in batch],self.max_context_tokens,pad_to_max=self.pad_context_to_max); b=len(batch); kmax=self.k_max
        pad=int(getattr(self.tokenizer,"pad_token_id",0) or 0)
        oi=torch.full((b,kmax,self.max_option_tokens),pad,dtype=torch.long); oa=torch.zeros_like(oi); om=torch.zeros((b,kmax),dtype=torch.bool); y=torch.zeros(b,dtype=torch.long)
        texts=[]; coords=[]
        for bi,row in enumerate(batch):
            if len(row["options"])>kmax: raise ValueError("K exceeds k_max")
            y[bi]=int(row["label"])
            for j,t in enumerate(row["options"]): om[bi,j]=True; texts.append(t); coords.append((bi,j))
        if texts:
            tok=self._tok(texts,self.max_option_tokens)
            for n,(bi,j) in enumerate(coords): oi[bi,j]=tok["input_ids"][n]; oa[bi,j]=tok["attention_mask"][n]
        return {"ctx_ids":ctx["input_ids"],"ctx_mask":ctx["attention_mask"],"opt_ids":oi,"opt_mask":om,"opt_attn":oa,"y":y}