omnijev-work / code /pipeline.py
fnruha0921's picture
stage3: pipeline.py
9a02d71 verified
Raw History Blame Contribute Delete
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