Download code/build_data.py from fnruha0921/omnijev-work: direct link, hf CLI and curl.
- Browser
- Download file 21.2 kB
-
https://huggingface.co/fnruha0921/omnijev-work/resolve/main/code/build_data.py
- Command line
-
hf download hf://fnruha0921/omnijev-work/code/build_data.py
-
curl -L -o build_data.py https://huggingface.co/fnruha0921/omnijev-work/resolve/main/code/build_data.py
21.2 kB
| """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") | |