Image-Text-to-Text
PEFT
Safetensors
vision-language
multimodal
llava
lora
siglip2
n-atlas
nigerian-languages
File size: 4,103 Bytes
2a2540a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
#!/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)