File size: 21,177 Bytes
8c867d9
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
7eff85e
 
 
 
8c867d9
 
 
 
7eff85e
 
 
8c867d9
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
"""Build the MM-Jev train / eval sets and cache the frozen tower features (run inside the Colab kernel, `jev` loaded).

Every record: dict(task, modality, state=[Seg with cached features], q=question dict, y=gold option index, raw=raw media
for the latency subset or None). Tower features are cached once (the towers are frozen), so training runs the LM only.
"""
import io, json, random, time, urllib.request
import numpy as np, torch
from datasets import load_dataset
from mmjev import Seg, options_of
from media import shapes_image, beeps, moving_video, COLORS, SHAPES

rng = random.Random(0)
nrng = np.random.default_rng(0)
LOG = open("/content/build.log", "a")


def log(*a):
    print(*a, file=LOG, flush=True)


N_RAW = 12  # eval items per task that keep raw media for end-to-end latency


IMG_GRID, FRAME_GRID = 8, 4   # cache grids at the resolution the LM consumes (64 / 16 tokens): 16x / 64x less RAM


def img_feat(images, grid=IMG_GRID):
    g = jev.vision_tower_features(images).float()
    return torch.nn.functional.avg_pool2d(g, g.shape[-1] // grid).half().cpu()


def aud_feat(clips):
    out = []
    for s in range(0, len(clips), 16):
        out += [a.cpu() for a in jev.audio_tower_features(clips[s:s + 16])]
    return out


def finish(recs, kind):
    """Batch-encode media of records whose state holds raw media, keep raw for the first N_RAW eval items."""
    for split in ("train", "eval"):
        rs = [r for r in recs if r["split"] == split]
        for i, r in enumerate(rs):
            r["raw"] = [Seg(s.kind, s.data, s.audio, s.fps) for s in r["state"]] if (split == "eval" and i < N_RAW) else None
    if kind == "image":
        for s in range(0, len(recs), 32):
            chunk = recs[s:s + 32]
            f = img_feat([r["state"][0].data for r in chunk])
            for r, x in zip(chunk, f):
                r["state"][0] = Seg("image", x)
    elif kind == "audio":
        for s in range(0, len(recs), 64):
            chunk = recs[s:s + 64]
            f = aud_feat([r["state"][0].data for r in chunk])
            for r, x in zip(chunk, f):
                r["state"][0] = Seg("audio", x)
    elif kind == "video":
        for r in recs:
            sg = r["state"][0]
            fr = img_feat(sg.data, FRAME_GRID)
            au = aud_feat([sg.audio])[0] if sg.audio is not None else None
            r["state"][0] = Seg("video", fr, au, sg.fps)
    return recs


def rec(task, mod, split, state, q, y, target=None):
    """One state, one or more questions. y: gold option index per question; target: soft distribution per question."""
    qs = q if isinstance(q, list) else [q]
    ys = [int(v) for v in (y if isinstance(y, list) else [y])]
    if target is None:
        target = [[1.0 if i == yy else 0.0 for i in range(len(options_of(qq)[0]))] for qq, yy in zip(qs, ys)]
    return dict(task=task, modality=mod, split=split, state=state, qs=qs, ys=ys, targets=target)


# ------------------------------------------------------------------------------------------ text
def build_text(n_tr=500, n_ev=200):
    out = []
    b = load_dataset("google/boolq")
    for split, ds, n in (("train", b["train"], n_tr), ("eval", b["validation"], n_ev)):
        for ex in ds.shuffle(seed=0).select(range(n)):
            q = ex["question"].strip().capitalize() + "?"
            out.append(rec("boolq", "text", split, [Seg("text", ex["passage"])],
                           {"type": "noul", "instructions": f"Based on the passage: {q}"}, int(ex["answer"])))
    db = load_dataset("fancyzhx/dbpedia_14")
    names = db["train"].features["label"].names
    crit = {n.lower(): "" for n in names}
    for split, ds, n in (("train", db["train"], n_tr), ("eval", db["test"], n_ev)):
        for ex in ds.shuffle(seed=0).select(range(n)):
            out.append(rec("dbpedia14", "text", split, [Seg("text", ex["title"] + ". " + ex["content"])],
                           {"type": "choice", "instructions": "Which category does this entity belong to?",
                            "criteria": crit}, ex["label"]))
    sst = load_dataset("SetFit/sst5")
    levels = ["very negative", "negative", "neutral", "positive", "very positive"]
    for split, ds, n in (("train", sst["train"], n_tr), ("eval", sst["test"], n_ev)):
        for ex in ds.shuffle(seed=0).select(range(n)):
            out.append(rec("sst5", "text", split, [Seg("text", "Review: " + ex["text"])],
                           {"type": "score", "instructions": "How positive is the sentiment of this review?",
                            "criteria": levels}, ex["label"]))
    for r in out:
        r["raw"] = None
    return out


def build_jevbench():
    base = "https://raw.githubusercontent.com/fstandhartinger/jevbench/main/datasets/public/"
    out = []
    for tier in ("easy", "original", "hard"):
        for line in urllib.request.urlopen(base + f"{tier}.jsonl", timeout=60).read().decode().splitlines():
            if not line.strip():
                continue
            t = json.loads(line)
            q, st = t["question"], t["state"]
            st = st if isinstance(st, str) else json.dumps(st, ensure_ascii=False)
            crit = q.get("criteria")
            if q["type"] == "noul":
                crit = crit or {}
                qq = {"type": "noul", "instructions": q["instructions"],
                      "criteria": {"yes": crit.get("true", "yes"), "no": crit.get("false", "no")}}
                exp = str(t["expected"]).lower()
                y = 1 if exp in ("true", "yes", "1") else 0
            elif q["type"] == "score":
                qq = {"type": "score", "instructions": q["instructions"], "criteria": list(crit)}
                y = int(t["expected"])
            else:
                qq = {"type": "choice", "instructions": q["instructions"],
                      "criteria": {k: (v or "") for k, v in (crit or {}).items()} if isinstance(crit, dict)
                      else {k: "" for k in (crit or t["labels"])}}
                y = list(qq["criteria"]).index(str(t["expected"]))
            r = rec(f"jevbench_{tier}", "text", "eval", [Seg("text", st)], qq, y)
            r["family"] = t.get("family"); r["raw"] = None
            out.append(r)
    return out


# ------------------------------------------------------------------------------------------ image
def build_image(n_ok=(800, 300), n_pope=(600, 300), n_cnt=(400, 150)):
    out = []
    ok = load_dataset("HuggingFaceM4/A-OKVQA")
    for split, ds, n in (("train", ok["train"], n_ok[0]), ("eval", ok["validation"], n_ok[1])):
        for ex in ds.shuffle(seed=0).select(range(n)):
            out.append(rec("aokvqa", "image", split, [Seg("image", ex["image"].convert("RGB"))],
                           {"type": "choice", "instructions": ex["question"], "criteria": {c: "" for c in ex["choices"]}},
                           ex["correct_choice_idx"]))
    pope = load_dataset("lmms-lab/POPE", split="test").shuffle(seed=0)
    imgs = sorted(set(pope["image_source"]))
    rng.shuffle(imgs)
    ev_imgs = set(imgs[: len(imgs) // 4])
    ctr = {"train": 0, "eval": 0}
    for ex in pope:
        split = "eval" if ex["image_source"] in ev_imgs else "train"
        if ctr[split] >= (n_pope[1] if split == "eval" else n_pope[0]):
            continue
        ctr[split] += 1
        out.append(rec("pope", "image", split, [Seg("image", ex["image"].convert("RGB"))],
                       {"type": "noul", "instructions": ex["question"]}, int(ex["answer"].strip().lower() == "yes")))
        if all(ctr[s] >= (n_pope[1] if s == "eval" else n_pope[0]) for s in ctr):
            break
    words = ["zero", "one", "two", "three", "four", "five"]
    for split, n in (("train", n_cnt[0]), ("eval", n_cnt[1])):
        for _ in range(n):
            k = rng.randint(0, 5)
            objs, tries = [], 0
            while len(objs) < k and tries < 200:
                tries += 1
                cx, cy, r_ = rng.uniform(0.12, 0.88), rng.uniform(0.12, 0.88), rng.uniform(0.05, 0.1)
                if all((cx - o[2]) ** 2 + (cy - o[3]) ** 2 > (r_ + o[4] + 0.03) ** 2 for o in objs):
                    objs.append((rng.choice(SHAPES), rng.choice(list(COLORS)), cx, cy, r_))
            out.append(rec("count_syn", "image", split, [Seg("image", shapes_image(objs, size=768))],
                           {"type": "score", "instructions": "How many shapes are in the image?", "criteria": words},
                           len(objs)))
    return finish(out, "image")


# ------------------------------------------------------------------------------------------ audio
def build_audio(n_ch=(600, 200), n_no=(500, 200), n_bp=(300, 120)):
    import librosa
    import io, soundfile as sf
    from datasets import Audio
    # decode the raw bytes ourselves: datasets' Audio decoding needs torchcodec, which not every image ships
    esc = load_dataset("ashraq/esc50", split="train").cast_column("audio", Audio(decode=False))
    cats = sorted(set(esc["category"]))
    pretty = lambda c: c.replace("_", " ")
    clips = {"train": [], "eval": []}
    for ex in esc:
        y, sr = sf.read(io.BytesIO(ex["audio"]["bytes"]), dtype="float32", always_2d=False)
        y = y.mean(1) if y.ndim > 1 else y
        y16 = librosa.resample(y, orig_sr=sr, target_sr=16000) if sr != 16000 else y
        clips["eval" if ex["fold"] == 5 else "train"].append((y16, ex["category"]))
    out = []
    for split in ("train", "eval"):
        pool = clips[split]
        idx = rng.sample(range(len(pool)), min(len(pool), n_ch[0] if split == "train" else n_ch[1]))
        for i in idx:
            y, c = pool[i]
            opts = [c] + rng.sample([x for x in cats if x != c], 4)
            rng.shuffle(opts)
            out.append(rec("esc50_choice", "audio", split, [Seg("audio", y)],
                           {"type": "choice", "instructions": "Which sound is in the recording?",
                            "criteria": {pretty(o): "" for o in opts}}, opts.index(c)))
        idx = rng.sample(range(len(pool)), min(len(pool), n_no[0] if split == "train" else n_no[1]))
        for j, i in enumerate(idx):
            y, c = pool[i]
            pos = j % 2 == 0
            ask = c if pos else rng.choice([x for x in cats if x != c])
            out.append(rec("esc50_noul", "audio", split, [Seg("audio", y)],
                           {"type": "noul", "instructions": f"Does the recording contain the sound of {pretty(ask)}?"},
                           int(pos)))
    words = ["zero", "one", "two", "three", "four", "five"]
    for split, n in (("train", n_bp[0]), ("eval", n_bp[1])):
        for _ in range(n):
            k = rng.randint(0, 5)
            out.append(rec("beeps_count", "audio", split,
                           [Seg("audio", beeps(k, freq=rng.choice([440.0, 660.0, 880.0, 1200.0]), rng=nrng))],
                           {"type": "score", "instructions": "How many beeps are in the audio?", "criteria": words}, k))
    return finish(out, "audio")


# ------------------------------------------------------------------------------------------ video
class _VidList(list):
    """Encodes each video record as soon as it is appended; keeps raw frames only for the first N_RAW eval items."""
    def __init__(self):
        super().__init__(); self.n_eval = {}

    def append(self, r):
        k = r["task"]
        keep_raw = r["split"] == "eval" and self.n_eval.get(k, 0) < N_RAW
        if r["split"] == "eval":
            self.n_eval[k] = self.n_eval.get(k, 0) + 1
        sg = r["state"][0]
        r["raw"] = [Seg(sg.kind, sg.data, sg.audio, sg.fps)] if keep_raw else None
        au = aud_feat([sg.audio])[0] if sg.audio is not None else None
        r["state"][0] = Seg("video", img_feat(sg.data, FRAME_GRID), au, sg.fps)
        super().append(r)


def build_video(n_dir=(300, 120), n_chg=(250, 100), n_fl=(250, 100), n_av=(200, 80)):
    out = _VidList()
    dirs = ["left", "right", "up", "down"]
    col = list(COLORS)
    for split, n in (("train", n_dir[0]), ("eval", n_dir[1])):
        for _ in range(n):
            d, c, s = rng.choice(dirs), rng.choice(col), rng.choice(SHAPES)
            fr = moving_video(s, c, d, size=768, rng=nrng)
            out.append(rec("vid_direction", "video", split, [Seg("video", fr, fps=2.0)],
                           {"type": "choice", "instructions": f"In which direction does the {c} {s} move?",
                            "criteria": {x: "" for x in dirs}}, dirs.index(d)))
    for split, n in (("train", n_chg[0]), ("eval", n_chg[1])):
        for j in range(n):
            c, s = rng.choice(col), rng.choice(SHAPES)
            chg = j % 2 == 0
            fr = moving_video(s, c, rng.choice(dirs), size=768, rng=nrng,
                              change_to=rng.choice([x for x in col if x != c]) if chg else None)
            out.append(rec("vid_color_change", "video", split, [Seg("video", fr, fps=2.0)],
                           {"type": "noul", "instructions": f"Does the {s} change its colour during the video?"},
                           int(chg)))
    for split, n in (("train", n_fl[0]), ("eval", n_fl[1])):
        for _ in range(n):
            k = rng.randint(0, 3)
            fr = moving_video(rng.choice(SHAPES), rng.choice(col), rng.choice(dirs), size=768, flashes=k, rng=nrng)
            out.append(rec("vid_flash_count", "video", split, [Seg("video", fr, fps=2.0)],
                           {"type": "score", "instructions": "How many times does the background flash bright yellow?",
                            "criteria": ["never", "once", "twice", "three times"]}, k))
    for split, n in (("train", n_av[0]), ("eval", n_av[1])):
        for j in range(n):
            has = j % 2 == 0
            fr = moving_video(rng.choice(SHAPES), rng.choice(col), rng.choice(dirs), size=768, rng=nrng)
            au = beeps(rng.randint(1, 4), rng=nrng) if has else beeps(0, rng=nrng)
            out.append(rec("vid_av_beep", "video", split, [Seg("video", fr, audio=au, fps=2.0)],
                           {"type": "noul", "instructions": "Can beeping tones be heard in the video's soundtrack?"},
                           int(has)))
    return list(out)


# ------------------------------------------------------------------------------------------ public benchmarks
def build_typed_decisions():
    """LocalLLaMA/typed-decisions: 1,200 train cases (6,000 decisions) and 400 test cases (2,000 decisions), gold =
    teacher distributions. Scored exactly like the Laya notebook."""
    out = []
    for split, hf in (("train", "train"), ("eval", "test")):
        for row in load_dataset("LocalLLaMA/typed-decisions", "all", split=hf):
            state = json.loads(row["state"]); qd = json.loads(row["questions"]); gold = json.loads(row["gold"])
            st = state if isinstance(state, str) else json.dumps(state, ensure_ascii=False)
            qs, ys, tg, names = [], [], [], []
            for qid, q in qd.items():
                g, t, crit = gold[qid], q["type"], q.get("criteria")
                if t == "noul":
                    crit = crit or {}
                    qq = {"type": "noul", "instructions": q["instructions"],
                          "criteria": {"no": crit.get("false", "no"), "yes": crit.get("true", "yes")}}
                    pt = float(g.get("noul", g.get("probabilities", {}).get("true", 0.5)))
                    dist = [1 - pt, pt]
                    y = int(str(g["label"]).lower() == "true")
                elif t == "score":
                    crit = list(crit) if isinstance(crit, list) else [str(i) for i in range(4)]
                    qq = {"type": "score", "instructions": q["instructions"], "criteria": crit}
                    dist = [float(g.get("probabilities", {}).get(str(i), 0.0)) for i in range(len(crit))]
                    y = int(g.get("label", round(g.get("score", 0))))
                else:
                    keys = list(crit)
                    qq = {"type": "choice", "instructions": q["instructions"], "criteria": {k: (crit[k] or "") for k in keys}}
                    dist = [float(g.get("probabilities", {}).get(k, 0.0)) for k in keys]
                    y = keys.index(str(g["label"]))
                ssum = sum(dist)
                dist = [d / ssum for d in dist] if ssum > 0 else [1 / len(dist)] * len(dist)
                qs.append(qq); ys.append(y); tg.append(dist); names.append(qid)
            r = rec("typed_decisions", "text", split, [Seg("text", st)], qs, ys, tg)
            r.update(workflow=row["workflow"], qnames=names, gold=gold, qraw=qd, raw=None)
            out.append(r)
    return out


def build_btzsc(seed=20260917, n=100):
    """AbdelStark/jev-benchmarks pilot-v1 protocol: btzsc/btzsc @ fef2a2a, 100 class-balanced test examples per set."""
    import random as _r
    from collections import defaultdict
    out = []
    for off, (name, task) in enumerate((("agnews", "topic"), ("emotiondair", "emotion"), ("banking77", "intent"))):
        rows = load_dataset("btzsc/btzsc", name=name, split="test", revision="fef2a2ac62b69c58670047dddf045c53d7c3cb5e")
        binary = [int(v) for v in rows["labels"]]; texts = [str(v) for v in rows["text"]]
        k = next(i for i in range(1, len(texts)) if texts[i] != texts[0])
        labels = [str(rows[i]["hypothesis"]) for i in range(k)]
        vidx, targets = [], []
        for si in range(len(rows) // k):
            v = binary[si * k:(si + 1) * k]
            if sum(v) == 1:
                vidx.append(si); targets.append(v.index(1))
        by = defaultdict(list)
        for i, t in enumerate(targets):
            by[t].append(i)
        rr = _r.Random(seed + off)
        for v in by.values():
            rr.shuffle(v)
        chosen = []
        while len(chosen) < min(n, len(targets)):
            prog = False
            for c in sorted(by):
                if by[c] and len(chosen) < n:
                    chosen.append(by[c].pop()); prog = True
            if not prog:
                break
        for pos in sorted(chosen):
            si = vidx[pos]
            r = rec(f"btzsc_{name}", "text", "eval", [Seg("text", texts[si * k])],
                    {"type": "choice", "instructions": "Which single label best describes the input text?",
                     "criteria": {l: "" for l in labels}}, targets[pos])
            r["raw"] = None
            out.append(r)
    return out


def build_massive_xnli(n_intent=500, n_xnli=1000):
    out = []
    for lang in ("en", "ko"):
        ds = load_dataset("mteb/amazon_massive_scenario", lang, split="test")
        labels = sorted(set(ds["label_text"]))
        crit = {l.replace("_", " "): "" for l in labels}
        for ex in ds:
            r = rec(f"massive_scenario_{lang}", "text", "eval", [Seg("text", ex["text"])],
                    {"type": "choice", "instructions": "Which scenario (domain) does this user request belong to?",
                     "criteria": crit}, labels.index(ex["label_text"]))
            r["raw"] = None; out.append(r)
    ds = load_dataset("mteb/amazon_massive_intent", "en", split="test").shuffle(seed=0)
    labels = sorted(set(ds["label_text"]))
    crit = {l.replace("_", " "): "" for l in labels}
    for ex in ds.select(range(n_intent)):
        r = rec("massive_intent_en", "text", "eval", [Seg("text", ex["text"])],
                {"type": "choice", "instructions": "Which intent does this user request express?", "criteria": crit},
                labels.index(ex["label_text"]))
        r["raw"] = None; out.append(r)
    x = load_dataset("facebook/xnli", "en", split="test").shuffle(seed=0).select(range(n_xnli))
    names = ["entailment", "neutral", "contradiction"]
    for ex in x:
        r = rec("xnli_en", "text", "eval", [Seg("text", f"Premise: {ex['premise']}\nHypothesis: {ex['hypothesis']}")],
                {"type": "choice", "instructions": "What is the relation between the premise and the hypothesis?",
                 "criteria": {"entailment": "the premise implies the hypothesis",
                              "neutral": "the premise neither implies nor contradicts the hypothesis",
                              "contradiction": "the premise contradicts the hypothesis"}}, ex["label"])
        r["raw"] = None; out.append(r)
    return out


if globals().get("RUN_BUILD_MAIN", True):
    t0 = time.time()
    DATA = globals().get("DATA", {})
    for name, fn in (("text", build_text), ("jevbench", build_jevbench), ("typed", build_typed_decisions),
                     ("btzsc", build_btzsc), ("massive_xnli", build_massive_xnli),
                     ("text2", lambda: (exec(open("/content/build_text2.py").read(), globals()), DATA["text2"])[1]),
                     ("image", build_image), ("audio", build_audio), ("video", build_video)):
        try:
            _ = DATA[name]; log(f"[{name}] cached"); continue
        except (NameError, KeyError):
            pass
        t1 = time.time()
        DATA[name] = fn()
        import gc; gc.collect()
        log(f"[{name}] {len(DATA[name])} records in {time.time() - t1:.0f}s")
    torch.save(DATA, "/content/mmjev_data.pt")
    log(f"DONE {sum(len(v) for v in DATA.values())} records, {time.time() - t0:.0f}s")