#!/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)