Image-Text-to-Text
PEFT
Safetensors
vision-language
multimodal
llava
lora
siglip2
n-atlas
nigerian-languages
AtlasVision / code /eval_stage1.py
UncleanCode's picture
Upload via givemeanode export_data
2a2540a verified
Raw History Blame Contribute Delete
4.1 kB
#!/usr/bin/env python3
"""Stage-1 verification on samples NOT seen in training.
1) checkpoint integrity 2) held-out loss: random projector vs trained vs trained-with-shuffled-images
3) captions on held-out images 4) COCO zip readiness for stage 2"""
import json, os, time, zipfile
os.environ.setdefault("HF_HUB_OFFLINE", "1")
import torch
from torch.utils.data import DataLoader
import train as T
OUT = os.path.expanduser("~/eval"); os.makedirs(OUT, exist_ok=True)
res = {}
t0 = time.time()
# 1. checkpoint integrity
sd = torch.load(os.path.join(T.CKPT_DIR, "projector_final.pt"), map_location="cpu")
res["checkpoint"] = {k: list(v.shape) for k, v in sd.items()}
res["checkpoint_all_finite"] = all(torch.isfinite(v).all().item() for v in sd.values())
print("checkpoint:", res["checkpoint"], "finite:", res["checkpoint_all_finite"], flush=True)
model, tok, pad_id, proc = T.build_model()
ann = json.load(open(T.JSON_PATH))
perm = torch.randperm(len(ann), generator=torch.Generator().manual_seed(T.SEED)).tolist()
held = [ann[i] for i in perm[T.NUM_SAMPLES:T.NUM_SAMPLES + 2000]] # never trained on
prefix = T.find_zip_prefix(held, T.ZIP_PATH)
ds = T.LLaVAPretrainDataset(held, T.ZIP_PATH, prefix, proc, tok, T.MAX_TEXT_LEN)
dl = DataLoader(ds, batch_size=32, num_workers=8, collate_fn=T.make_collate(pad_id))
@torch.no_grad()
def val_loss(shuffle_images=False):
model.eval(); tot, n = 0.0, 0
for b in dl:
pv = b["pixel_values"].cuda()
if shuffle_images:
pv = pv.roll(1, dims=0) # each caption paired with a different image
with torch.autocast("cuda", dtype=T.DTYPE):
l = model(pv, b["input_ids"].cuda(), b["attention_mask"].cuda(), b["labels"].cuda()).float()
k = int((b["labels"] != -100).sum()); tot += l.item() * k; n += k
return round(tot / n, 4)
# 2. held-out loss
torch.manual_seed(0)
model.projector = T.ProjectionMLP(model.projector.net[0].in_features, model.projector.net[2].out_features).cuda()
res["heldout_loss_random_projector"] = val_loss()
model.projector.load_state_dict(sd)
res["heldout_loss_trained"] = val_loss()
res["heldout_loss_trained_shuffled_images"] = val_loss(shuffle_images=True)
print("held-out loss:", {k: v for k, v in res.items() if k.startswith("heldout")}, flush=True)
# 3. captions on held-out images
eot = tok.convert_tokens_to_ids(T.EOT)
caps = []
for item in held[:8]:
img, ok = None, True
with zipfile.ZipFile(T.ZIP_PATH) as zf, zf.open(prefix + item["image"]) as f:
from PIL import Image
img = Image.open(f).convert("RGB")
pv = proc(images=img, return_tensors="pt").pixel_values.cuda()
ids = torch.tensor([tok("Describe this image briefly." + T.ASSIST_HEADER, add_special_tokens=False).input_ids], device="cuda")
with torch.no_grad(), torch.autocast("cuda", dtype=T.DTYPE):
e, m, _ = model.build_inputs(pv, ids, torch.ones_like(ids))
out = model.llm.generate(inputs_embeds=e, attention_mask=m, max_new_tokens=50, do_sample=False,
eos_token_id=eot, pad_token_id=pad_id)
caps.append({"image": item["image"], "reference": item["conversations"][1]["value"],
"model": tok.decode(out[0], skip_special_tokens=True).strip()})
res["captions"] = caps
for c in caps:
print(f"\n[{c['image']}]\n ref: {c['reference']}\n model: {c['model']}", flush=True)
# 4. stage-2 data readiness
with zipfile.ZipFile(os.path.expanduser("~/data/coco/train2017.zip")) as zf:
names = set(zf.namelist())
inst = json.load(open(os.path.expanduser("~/data/llava_instruct/llava_instruct_150k.json")))
hits = sum(("train2017/" + a["image"]) in names for a in inst[:2000])
res["coco_zip_files"] = len(names)
res["instruct_conversations"] = len(inst)
res["instruct_images_found_of_2000"] = hits
print("\nstage-2 data:", res["coco_zip_files"], "files in COCO zip;", len(inst), "conversations;",
hits, "/2000 images found", flush=True)
res["eval_minutes"] = round((time.time() - t0) / 60, 1)
json.dump(res, open(os.path.join(OUT, "stage1_eval.json"), "w"), indent=1)
print("EVAL_DONE", flush=True)