"""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