Image-Text-to-Text
PEFT
Safetensors
vision-language
multimodal
llava
lora
siglip2
n-atlas
nigerian-languages
AtlasVision / code /infer.py
UncleanCode's picture
Upload via givemeanode export_data
2a2540a verified
Raw History Blame Contribute Delete
2.11 kB
#!/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)