Image-Text-to-Text
PEFT
Safetensors
vision-language
multimodal
llava
lora
siglip2
n-atlas
nigerian-languages
File size: 2,109 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
#!/usr/bin/env python3
"""Caption images with the trained projector.
Usage: python3 infer.py [projector.pt] [image paths...]   (default: 8 held-out images from the zip)"""
import json
import os
import sys
import zipfile

import torch
from PIL import Image

os.environ.setdefault("HF_HUB_OFFLINE", "1")
import train as T  # reuses the exact model layout and prompt format used in training

ckpt = sys.argv[1] if len(sys.argv) > 1 else os.path.join(T.CKPT_DIR, "projector_final.pt")
model, tok, pad_id, processor = T.build_model()
state = torch.load(ckpt, map_location="cpu")
state = state.get("projector_state_dict", state)
model.projector.load_state_dict(state)
model.eval()

if len(sys.argv) > 2:
    images = [(p, Image.open(p).convert("RGB")) for p in sys.argv[2:]]
else:  # images NOT in the training subset
    ann = json.load(open(T.JSON_PATH))
    g = torch.Generator().manual_seed(T.SEED)
    held_out = torch.randperm(len(ann), generator=g)[T.NUM_SAMPLES:T.NUM_SAMPLES + 8].tolist()
    zf = zipfile.ZipFile(T.ZIP_PATH)
    prefix = T.find_zip_prefix([ann[i] for i in held_out], T.ZIP_PATH, n_check=8)
    images = []
    for i in held_out:
        with zf.open(prefix + ann[i]["image"]) as f:
            images.append((f"{ann[i]['image']} | reference: {ann[i]['conversations'][1]['value']}",
                           Image.open(f).convert("RGB")))

prompt = "Describe this image briefly."
for name, img in images:
    pv = processor(images=img, return_tensors="pt").pixel_values.to(T.DEVICE)
    ids = torch.tensor([tok(prompt + T.ASSIST_HEADER, add_special_tokens=False).input_ids], device=T.DEVICE)
    mask = torch.ones_like(ids)
    with torch.no_grad(), torch.autocast(device_type=T.DEVICE.type, dtype=T.DTYPE):
        embeds, full_mask, _ = model.build_inputs(pv, ids, mask)
        out = model.llm.generate(inputs_embeds=embeds, attention_mask=full_mask, max_new_tokens=60,
                                 do_sample=False, eos_token_id=tok.convert_tokens_to_ids(T.EOT), pad_token_id=pad_id)
    print(f"\n{name}\n  -> {tok.decode(out[0], skip_special_tokens=True).strip()}", flush=True)