Download code/pipeline.py from fnruha0921/omnijev-work: direct link, hf CLI and curl.
- Browser
- Download file 17 kB
-
https://huggingface.co/fnruha0921/omnijev-work/resolve/main/code/pipeline.py
- Command line
-
hf download hf://fnruha0921/omnijev-work/code/pipeline.py
-
curl -L -o pipeline.py https://huggingface.co/fnruha0921/omnijev-work/resolve/main/code/pipeline.py
17 kB
| """MM-Jev v3 end-to-end pipeline, run as a detached process on the Colab VM: | |
| nohup python /content/pipeline.py > /content/pipe/stdout.log 2>&1 & | |
| Independent of the local CLI connection. Every stage caches its output under /content/pipe and records completion in | |
| state.json; re-running the script resumes (data parts are cached, training resumes from the last checkpoint). | |
| Progress: /content/pipe/pipeline.log (one line per event), results: /content/pipe/results.json. | |
| """ | |
| import gc, json, os, random, sys, time, traceback | |
| sys.path.insert(0, "/content") | |
| os.environ.setdefault("TORCHINDUCTOR_COMPILE_THREADS", "4") | |
| import torch | |
| P = "/content/pipe"; D = f"{P}/data"; A = f"{P}/assets" | |
| for d in (P, D, A): | |
| os.makedirs(d, exist_ok=True) | |
| LOG, STATE, RES = f"{P}/pipeline.log", f"{P}/state.json", f"{P}/results.json" | |
| MAX_TRAIN = int(os.environ.get("MMJEV_MAX_TRAIN", 36000)) | |
| TEXT_ONLY = bool(os.environ.get("MMJEV_TEXT_ONLY")) # stage 1: text only (validate before multimodal) | |
| WORK_REPO = os.environ.get("MMJEV_WORK_REPO") # e.g. fnruha0921/omnijev-work; None = no hub sync | |
| RUN = os.environ.get("MMJEV_RUN_PREFIX", "run") # hub folder of this stage (run = text stage, run_mm = multimodal) | |
| INIT_FROM_HUB = os.environ.get("MMJEV_INIT_FROM_HUB") # e.g. run/adapter_final.pt: start from the text-stage adapter | |
| _HUB_LOCK = __import__("threading").Lock() | |
| def hub_upload(local, remote, background=True): | |
| """Best effort: a network error must never stop training.""" | |
| if not WORK_REPO or not os.path.exists(local): | |
| return | |
| def _go(): | |
| try: | |
| from huggingface_hub import HfApi | |
| with _HUB_LOCK: | |
| HfApi().upload_file(path_or_fileobj=local, path_in_repo=remote, repo_id=WORK_REPO, | |
| commit_message=f"sync {remote}") | |
| except Exception as e: # noqa: BLE001 | |
| log(f"[hub] upload {remote} failed: {repr(e)[:200]}") | |
| if background: | |
| __import__("threading").Thread(target=_go, daemon=True).start() | |
| else: | |
| _go() | |
| def hub_restore(): | |
| """Pull what a previous (killed) job left behind: checkpoint, final adapter, stage state, results, log.""" | |
| if not WORK_REPO: | |
| return | |
| from huggingface_hub import hf_hub_download | |
| for remote, local in ((f"{RUN}/ckpt/last.pt", f"{P}/ckpt/last.pt"), (f"{RUN}/adapter_final.pt", f"{P}/adapter_final.pt"), | |
| (f"{RUN}/state.json", STATE), (f"{RUN}/results.json", RES), (f"{RUN}/pipeline.log", LOG)): | |
| if os.path.exists(local): | |
| continue | |
| try: | |
| os.makedirs(os.path.dirname(local), exist_ok=True) | |
| path = hf_hub_download(WORK_REPO, remote) | |
| __import__("shutil").copy(path, local) | |
| print("restored", remote, flush=True) | |
| except Exception: | |
| pass | |
| def hub_heartbeat(every=300): | |
| def _loop(): | |
| while True: | |
| time.sleep(every) | |
| for loc, rem in ((LOG, f"{RUN}/pipeline.log"), (RES, f"{RUN}/results.json"), (STATE, f"{RUN}/state.json")): | |
| hub_upload(loc, rem, background=False) | |
| __import__("threading").Thread(target=_loop, daemon=True).start() | |
| def log(*a): | |
| with open(LOG, "a") as f: | |
| print(time.strftime("%m-%d %H:%M:%S"), *a, file=f, flush=True) | |
| def state(): | |
| return json.load(open(STATE)) if os.path.exists(STATE) else {} | |
| def mark(k, v=True): | |
| s = state(); s[k] = v | |
| json.dump(s, open(STATE + ".tmp", "w"), indent=1); os.replace(STATE + ".tmp", STATE) | |
| def results(): | |
| return json.load(open(RES)) if os.path.exists(RES) else {} | |
| def put_result(k, v): | |
| r = results(); r[k] = v | |
| json.dump(r, open(RES + ".tmp", "w"), indent=1, default=float); os.replace(RES + ".tmp", RES) | |
| # ------------------------------------------------------------------------------------------ 1. model | |
| def load_model(): | |
| import mmjev, media | |
| from mmjev import MMJev, load_gemma3n, fp16_safe_vision | |
| torch.backends.cudnn.benchmark = True | |
| t0 = time.time() | |
| if not os.path.exists("/content/gemma3n/model-00004-of-00004.safetensors"): | |
| from huggingface_hub import snapshot_download | |
| snapshot_download("unsloth/gemma-3n-E4B-it", local_dir="/content/gemma3n") | |
| log(f"model downloaded in {time.time() - t0:.0f}s") | |
| base, proc = load_gemma3n("/content/gemma3n") | |
| from datasets import load_dataset | |
| cal = [x["image"].convert("RGB") for x in load_dataset("HuggingFaceM4/A-OKVQA", split="validation[:16]")] | |
| cal += [media.shapes_image([("circle", "red", 0.3, 0.4, 0.1)]), media.moving_video("square", "blue", "left")[0]] | |
| n = fp16_safe_vision(base.model.vision_tower, proc.image_processor(cal, return_tensors="pt")["pixel_values"]) | |
| jev = MMJev(base, proc) | |
| log(f"model loaded in {time.time() - t0:.0f}s, rescaled convs {n}, GPU {torch.cuda.memory_allocated() / 2**30:.1f}G") | |
| return jev | |
| # ------------------------------------------------------------------------------------------ 2. data parts | |
| def part(name, fn): | |
| path = f"{D}/{name}.pt" | |
| if os.path.exists(path): | |
| return torch.load(path, weights_only=False) | |
| t0 = time.time() | |
| recs = fn() | |
| torch.save(recs, path + ".tmp"); os.replace(path + ".tmp", path) | |
| log(f"[data] {name}: {len(recs)} records in {time.time() - t0:.0f}s") | |
| return recs | |
| FAILED = [] | |
| def build_all(jev): | |
| DATA = {} | |
| ns = {"jev": jev, "DATA": DATA, "RUN_BUILD_MAIN": False, "__name__": "build_data"} | |
| exec(open("/content/build_data.py").read(), ns) | |
| if os.path.exists(f"{D}/base.pt"): # saved from an interactive kernel | |
| DATA.update(torch.load(f"{D}/base.pt", weights_only=False)) | |
| else: # rebuild the base parts, each cached | |
| def script(path, key): | |
| def f(): | |
| exec(open(path).read(), ns) | |
| return ns["DATA"][key] | |
| return f | |
| for key, fn in (("text", ns["build_text"]), ("jevbench", ns["build_jevbench"]), ("typed", ns["build_typed_decisions"]), | |
| ("btzsc", ns["build_btzsc"]), ("massive_xnli", ns["build_massive_xnli"]), ("image", ns["build_image"]), | |
| ("text2", script("/content/build_text2.py", "text2")), ("text3", script("/content/build_text3.py", "text3"))): | |
| if TEXT_ONLY and key == "image": | |
| continue | |
| DATA[key] = part(key, fn) | |
| if "audio" not in DATA and not TEXT_ONLY: # ESC-50 (HF); its synthetic beeps task is excluded | |
| DATA["audio"] = part("audio", ns["build_audio"]) | |
| import build_mm_hf as M # every other image / audio / video set: HF datasets | |
| for name, fn in ({} if TEXT_ONLY else M.BUILDERS).items(): | |
| try: | |
| DATA[f"mm_{name}"] = part(f"mm_{name}", lambda fn=fn: fn(jev)) | |
| except Exception: | |
| log(f"[data] mm {name} FAILED: {traceback.format_exc()[-600:]}"); FAILED.append(name) | |
| gc.collect() | |
| import build_text4 as B | |
| for name, fn in B.BUILDERS.items(): | |
| try: | |
| DATA[f"t4_{name}"] = part(f"t4_{name}", fn) | |
| except Exception: | |
| log(f"[data] {name} FAILED: {traceback.format_exc()[-500:]}"); FAILED.append(name) | |
| gc.collect() | |
| import build_text5 as B5 # stage 3: more public typed-decision corpora | |
| for name, fn in B5.BUILDERS.items(): | |
| try: | |
| DATA[f"t5_{name}"] = part(f"t5_{name}", fn) | |
| except Exception: | |
| log(f"[data] {name} FAILED: {traceback.format_exc()[-500:]}"); FAILED.append(name) | |
| gc.collect() | |
| tot = {k: len(v) for k, v in DATA.items()} | |
| log(f"[data] parts {tot}") | |
| return DATA | |
| SYNTHETIC = {"count_syn", "beeps_count", "vid_direction", "vid_color_change", "vid_flash_count", "vid_av_beep"} | |
| CAPS = {"image": int(os.environ.get("MMJEV_CAP_IMAGE", 5000)), "audio": int(os.environ.get("MMJEV_CAP_AUDIO", 3500)), | |
| "video": int(os.environ.get("MMJEV_CAP_VIDEO", 2500)), "text_base": int(os.environ.get("MMJEV_CAP_TEXT_BASE", 6000)), | |
| "text_new": int(os.environ.get("MMJEV_CAP_TEXT_NEW", 10 ** 9))} | |
| # the rest of MAX_TRAIN: public text corpora | |
| def train_split(DATA): | |
| import train as T | |
| T.EXCLUDE_TASKS |= SYNTHETIC | |
| rng = random.Random(0) | |
| buckets = {"image": [], "audio": [], "video": [], "text_base": [], "text_new": [], "text_public": []} | |
| for k, v in DATA.items(): | |
| for r in v: | |
| if r["split"] != "train" or r["task"] in T.EXCLUDE_TASKS: | |
| continue | |
| if any(len(t) < 2 for t in r["targets"]) or not r["qs"]: | |
| continue # degenerate questions (a single option) carry no signal | |
| key = r["modality"] if r["modality"] != "text" else ( | |
| "text_public" if k.startswith("t4_") else "text_new" if k.startswith("t5_") else "text_base") | |
| buckets[key].append(r) | |
| recs = [] | |
| caps = dict(CAPS, text_base=10 ** 9) if TEXT_ONLY else CAPS | |
| for key, v in buckets.items(): | |
| rng.shuffle(v) | |
| if key != "text_public": | |
| recs += v[:caps[key]] | |
| recs += buckets["text_public"][:max(0, MAX_TRAIN - len(recs))] | |
| rng.shuffle(recs) | |
| n_cal = max(300, len(recs) // 40) | |
| return recs[n_cal:], recs[:n_cal] | |
| # ------------------------------------------------------------------------------------------ main | |
| def main(): | |
| hub_restore() | |
| log("=== pipeline start ===") | |
| hub_heartbeat() | |
| jev = load_model() | |
| DATA = build_all(jev) | |
| if FAILED and not os.environ.get("MMJEV_ALLOW_SKIP"): | |
| log(f"PIPELINE FAILED: data parts failed {FAILED} -- fix and re-run (cached parts are reused)") | |
| return | |
| mark("data") | |
| import train as T | |
| import mmjev_eval as E | |
| E.LOGFILE = LOG | |
| T.LOG = LOG | |
| tr, cal = train_split(DATA) | |
| DATA.pop("video", None) # synthetic videos: neither trained on nor evaluated / timed | |
| log(f"train {len(tr)} / calib {len(cal)}; modalities " | |
| f"{ {m: sum(r['modality'] == m for r in tr) for m in ('text', 'image', 'audio', 'video')} }") | |
| if not state().get("trained"): | |
| init = None | |
| if INIT_FROM_HUB and WORK_REPO: | |
| from huggingface_hub import hf_hub_download | |
| init = hf_hub_download(WORK_REPO, INIT_FROM_HUB) | |
| log(f"init adapter from hub {INIT_FROM_HUB}") | |
| T.train(jev, DATA, bs=int(os.environ.get("MMJEV_BS", 16)), recs=(tr, cal), ckpt_dir=f"{P}/ckpt", init=init, | |
| save_path=f"{P}/adapter_final.pt", on_ckpt=lambda path, step: hub_upload(path, f"{RUN}/ckpt/last.pt")) | |
| hub_upload(f"{P}/adapter_final.pt", f"{RUN}/adapter_final.pt", background=False) | |
| mark("trained") | |
| else: | |
| T.add_lora(jev) | |
| T.load_trainable(jev, torch.load(f"{P}/adapter_final.pt", weights_only=False)) | |
| cfgs = {"full": (T.FULL, None), "fast": (T.FAST, T.FAST_WIDTH)} | |
| temps = state().get("temps") or {n: T.fit_temperature_k(jev, cal, fc, w) for n, (fc, w) in cfgs.items()} | |
| mark("temps", temps); log(f"temperatures {temps}") | |
| # accuracy: fastest benchmarks first, one log line per task | |
| order = ["jevbench_easy", "jevbench_original", "btzsc_agnews", "btzsc_emotiondair", "btzsc_banking77", "jevbench_hard", | |
| "typed_decisions", "massive_intent_en", "xnli_en", "massive_scenario_en", "massive_scenario_ko", | |
| "aokvqa", "pope", "scienceqa_img", "ai2d", "vqav2_yesno", "vqav2_count", "vqav2_other", | |
| "esc50_choice", "esc50_noul", "clotho_aqa", "urbansound8k", "speech_commands", | |
| "tempcompass", "mvbench_mini", "ucf101", "kinetics_mini", "boolq", "dbpedia14", "sst5", | |
| "snake_v1_action", "snake_v1_collision"] | |
| if TEXT_ONLY: | |
| order = [t for t in order if t in {r["task"] for v in DATA.values() for r in v if r["modality"] == "text"}] | |
| snake = [dict(r, split="eval") for r in DATA.get("text2", []) if r["task"] == "snake_syn"] | |
| if snake: | |
| DATA["snake_eval"] = snake | |
| order.append("snake_syn") | |
| have = {r["task"] for v in DATA.values() for r in v if r["split"] == "eval"} | |
| order = [t for t in order if t in have] | |
| for n, (fc, w) in cfgs.items(): | |
| done = results().get(f"acc_{n}", {}) | |
| for t in order: | |
| if t in done: | |
| continue | |
| m = E.accuracy_suite(jev, DATA, fc, w, temps[n], tag=f"v3-{n}", tasks=[t]) | |
| done.update(m); put_result(f"acc_{n}", done) | |
| mark("accuracy") | |
| # latency (compiled vision tower + CUDA-graph decoder), question scaling, tree vs naive | |
| vt = jev._core().vision_tower | |
| jev.vision_fn = None if TEXT_ONLY else torch.compile(lambda x: vt(pixel_values=x, do_pooling=False, return_dict=True).last_hidden_state, | |
| mode="max-autotune-no-cudagraphs", dynamic=False) | |
| for n, (fc, w) in cfgs.items(): | |
| if f"latency_{n}" not in results(): | |
| jev.temps_k = temps[n] | |
| put_result(f"latency_{n}", E.latency_suite(jev, DATA, fc, w)) | |
| put_result(f"qscale_{n}", E.questions_scaling(jev, fc, w)) | |
| if "tree_vs_naive" not in results(): | |
| put_result("tree_vs_naive", E.tree_vs_naive(jev, DATA, T.FAST, T.FAST_WIDTH)) | |
| mark("latency") | |
| # demos | |
| import demos | |
| demos.OUT = A | |
| from mmjev import set_ffn_width | |
| set_ffn_width(jev.lm, T.FAST_WIDTH); jev._graphs = {}; jev.temps_k = temps["fast"] | |
| if "snake_tune" not in results(): | |
| thr, sc = demos.tune_snake_native(jev, T.FAST) | |
| put_result("snake_tune", {"thr": thr, "scores": sc}); log(f"[demo] snake threshold {thr} (held-out seeds): {sc}") | |
| demos.COLL_THR = results()["snake_tune"]["thr"] | |
| for seed in (21, 11, 3, 5, 8): | |
| r = demos.snake_demo_native(jev, T.FAST, n=10, seed=seed, fps=18, name=f"snake_seed{seed}") | |
| log(f"[demo] snake seed {seed}: {r}") | |
| put_result(f"snake_seed{seed}", r) | |
| # multimodal cards on held-out raw media (fast config) | |
| ev = [] if TEXT_ONLY else [r for v in DATA.values() for r in v if r["split"] == "eval" and r.get("raw")] | |
| picks = [("image", "aokvqa", 2), ("image", "scienceqa_img", 1), ("audio", "esc50_choice", 2), ("audio", "urbansound8k", 1), | |
| ("audio", "speech_commands", 1), ("video", "tempcompass", 2), ("video", "ucf101", 1), ("video", "mvbench_mini", 1)] | |
| cards = [] | |
| for mod, task, k in picks: | |
| for i, r in enumerate([x for x in ev if x["task"] == task][:k]): | |
| seg = r["raw"][0] | |
| name = f"{task}_{i + 1}" | |
| try: | |
| if mod == "image": | |
| qs = r["qs"] + [{"type": "noul", "instructions": "Is there at least one person in the image?"}, | |
| {"type": "choice", "instructions": "Where was this photo most likely taken?", | |
| "criteria": {"indoors": "inside a building or vehicle", "outdoors": "outside, open air"}}] | |
| res, ms = demos.image_demo(jev, T.FAST, seg.data, qs, name=name, title=f"Image: {task}") | |
| elif mod == "audio": | |
| qs = r["qs"] + [{"type": "noul", "instructions": "Is this sound made by an animal?"}, | |
| {"type": "noul", "instructions": "Is someone speaking in the recording?"}] | |
| res, ms = demos.audio_demo(jev, T.FAST, seg.data, qs, name=name, title="Audio: ESC-50 (held out)") | |
| else: | |
| qs = r["qs"] + [{"type": "noul", "instructions": "Is there a person in the video?"}, | |
| {"type": "score", "instructions": "How much motion is there in the video?", | |
| "criteria": ["almost static", "some motion", "a lot of motion"]}] | |
| res, ms = demos.video_demo(jev, T.FAST, seg.data, qs, audio=seg.audio, fps=seg.fps, name=name, | |
| title=f"Video: {task}") | |
| from mmjev import options_of | |
| gold = options_of(r["qs"][0])[0][r["ys"][0]] | |
| pred = res[0].get("choice", res[0].get("level", res[0].get("noul"))) | |
| cards.append(dict(name=name, gold=gold, pred=pred, latency_ms=round(ms, 1))) | |
| log(f"[demo] {cards[-1]}") | |
| except Exception: | |
| log(f"[demo] {name} failed: {traceback.format_exc()[-400:]}") | |
| put_result("demo_cards", cards) | |
| set_ffn_width(jev.lm, None) | |
| mark("demos") | |
| for f in os.listdir(A): | |
| if f.endswith((".mp4", ".png")): | |
| hub_upload(f"{A}/{f}", f"{RUN}/assets/{f}", background=False) | |
| log("=== PIPELINE DONE ===") | |
| for loc, rem in ((LOG, f"{RUN}/pipeline.log"), (RES, f"{RUN}/results.json"), (STATE, f"{RUN}/state.json")): | |
| hub_upload(loc, rem, background=False) | |
| if __name__ == "__main__": | |
| try: | |
| main() | |
| except BaseException: | |
| log("PIPELINE FAILED\n" + traceback.format_exc()[-3000:]) | |
| raise | |