Image-Text-to-Text
PEFT
Safetensors
vision-language
multimodal
llava
lora
siglip2
n-atlas
nigerian-languages
AtlasVision / chat.py
UncleanCode's picture
Upload via givemeanode export_data
2a2540a verified
Raw History Blame
7.26 kB
#!/usr/bin/env python3
"""AtlasVision inference: SigLIP2 vision encoder + N-ATLaS (Llama-3 8B) via a trained MLP projector.
pip install -U torch transformers peft safetensors pillow accelerate huggingface_hub
# optional for 4-bit on small GPUs (e.g. Colab T4): pip install bitsandbytes
export HF_TOKEN=hf_... # needs accepted access to the gated NCAIR1/N-ATLaS
python chat.py --image photo.jpg --question "What is happening in this picture?" # stage 2
python chat.py --stage 1 --image photo.jpg # stage 1 captioner
python chat.py --image https://example.com/cat.jpg --question "Kedu ihe dị na foto a?"
python chat.py --image photo.jpg --load-in-4bit # ~8 GB GPU
python chat.py --image photo.jpg --interactive # several questions
Needs ~18 GB of GPU memory in bf16 (A100, L4, RTX 4090...); --load-in-4bit fits ~8 GB GPUs such as a Colab T4.
"""
import argparse
import io
import os
import sys
import torch
import torch.nn as nn
from PIL import Image
REPO = "FUTO-NIGERIA/AtlasVision"
LLM = "NCAIR1/N-ATLaS"
VISION = "google/siglip2-base-patch16-224"
USER_HEADER = "<|start_header_id|>user<|end_header_id|>\n\n"
ASSIST_HEADER = "<|eot_id|><|start_header_id|>assistant<|end_header_id|>\n\n"
EOT = "<|eot_id|>"
class ProjectionMLP(nn.Module):
def __init__(self, vision_dim, text_dim):
super().__init__()
self.net = nn.Sequential(nn.Linear(vision_dim, text_dim), nn.GELU(), nn.Linear(text_dim, text_dim))
def forward(self, x):
return self.net(x)
def fetch(repo, filename):
if os.path.isdir(repo):
return os.path.join(repo, filename)
from huggingface_hub import hf_hub_download
return hf_hub_download(repo, filename)
def load_image(src):
if src.startswith(("http://", "https://")):
import urllib.request
with urllib.request.urlopen(src) as r:
return Image.open(io.BytesIO(r.read())).convert("RGB")
return Image.open(src).convert("RGB")
class AtlasVision:
def __init__(self, stage=2, repo=REPO, llm=LLM, vision=VISION, load_in_4bit=False, device=None):
from safetensors.torch import load_file
from transformers import AutoImageProcessor, AutoModel, AutoModelForCausalLM, AutoTokenizer
self.device = torch.device(device or ("cuda" if torch.cuda.is_available() else "cpu"))
cuda = self.device.type == "cuda"
self.dtype = torch.bfloat16 if (not cuda or torch.cuda.is_bf16_supported()) else torch.float16
self.tok = AutoTokenizer.from_pretrained(llm)
self.pad_id = self.tok.pad_token_id if self.tok.pad_token_id is not None else self.tok.eos_token_id
self.eot_id = self.tok.convert_tokens_to_ids(EOT)
self.processor = AutoImageProcessor.from_pretrained(vision)
full_vision = AutoModel.from_pretrained(vision, dtype=self.dtype)
self.vision = full_vision.vision_model.to(self.device).eval()
vision_dim = full_vision.config.vision_config.hidden_size
del full_vision
kw = {"dtype": self.dtype}
if load_in_4bit:
from transformers import BitsAndBytesConfig
kw["quantization_config"] = BitsAndBytesConfig(load_in_4bit=True, bnb_4bit_quant_type="nf4",
bnb_4bit_compute_dtype=self.dtype)
kw["device_map"] = {"": self.device.index or 0}
self.llm = AutoModelForCausalLM.from_pretrained(llm, **kw)
if not load_in_4bit:
self.llm.to(self.device)
text_dim = self.llm.config.hidden_size
if stage == 2:
from peft import PeftModel
if os.path.isdir(repo):
self.llm = PeftModel.from_pretrained(self.llm, os.path.join(repo, "stage2/lora_adapter"))
else:
self.llm = PeftModel.from_pretrained(self.llm, repo, subfolder="stage2/lora_adapter")
self.llm.eval()
self.projector = ProjectionMLP(vision_dim, text_dim)
self.projector.load_state_dict(load_file(fetch(repo, f"stage{stage}/projector.safetensors")))
self.projector.to(self.device, dtype=torch.float32).eval()
self.prefix = torch.tensor([self.tok(USER_HEADER, add_special_tokens=True).input_ids], device=self.device)
self.history = [] # (question, answer) turns about the current image
@torch.no_grad()
def ask(self, image, question, max_new_tokens=256, temperature=0.0):
pv = self.processor(images=image, return_tensors="pt").pixel_values.to(self.device, self.dtype)
img = self.projector(self.vision(pixel_values=pv).last_hidden_state.float()).to(self.dtype)
text = ""
for i, (q, a) in enumerate(self.history):
text += (q if i == 0 else USER_HEADER + q) + ASSIST_HEADER + a + EOT
text += (question if not self.history else USER_HEADER + question) + ASSIST_HEADER
ids = torch.tensor([self.tok(text, add_special_tokens=False).input_ids], device=self.device)
emb = self.llm.get_input_embeddings()
embeds = torch.cat([emb(self.prefix).to(self.dtype), img, emb(ids).to(self.dtype)], dim=1)
mask = torch.ones(embeds.shape[:2], dtype=torch.long, device=self.device)
gen = dict(max_new_tokens=max_new_tokens, repetition_penalty=1.1, eos_token_id=self.eot_id, pad_token_id=self.pad_id)
if temperature > 0:
gen.update(do_sample=True, temperature=temperature, top_p=0.9)
else:
gen.update(do_sample=False)
out = self.llm.generate(inputs_embeds=embeds, attention_mask=mask, **gen)
answer = self.tok.decode(out[0], skip_special_tokens=True).strip()
self.history.append((question, answer))
return answer
def main():
ap = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
ap.add_argument("--image", required=True, help="path or URL")
ap.add_argument("--question", default=None)
ap.add_argument("--stage", type=int, default=2, choices=[1, 2])
ap.add_argument("--repo", default=REPO, help="HF repo id, or a local folder containing stage1/ and stage2/")
ap.add_argument("--llm", default=LLM)
ap.add_argument("--vision", default=VISION)
ap.add_argument("--load-in-4bit", action="store_true")
ap.add_argument("--max-new-tokens", type=int, default=256)
ap.add_argument("--temperature", type=float, default=0.0)
ap.add_argument("--interactive", action="store_true")
a = ap.parse_args()
question = a.question or ("Describe this image briefly." if a.stage == 1 else "Describe this image in detail.")
model = AtlasVision(a.stage, a.repo, a.llm, a.vision, a.load_in_4bit)
image = load_image(a.image)
print(f"\nQ: {question}\nA: {model.ask(image, question, a.max_new_tokens, a.temperature)}", flush=True)
while a.interactive:
try:
q = input("\nQ (empty to quit): ").strip()
except EOFError:
break
if not q:
break
print(f"A: {model.ask(image, q, a.max_new_tokens, a.temperature)}", flush=True)
if __name__ == "__main__":
sys.exit(main())