Image-Text-to-Text
PEFT
Safetensors
vision-language
multimodal
llava
lora
siglip2
n-atlas
nigerian-languages
Instructions to use Modularcomputing/AtlasVision with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- PEFT
How to use Modularcomputing/AtlasVision with PEFT:
Task type is invalid.
- Notebooks
- Google Colab
- Kaggle
File size: 9,993 Bytes
2a2540a ed78840 2a2540a ed78840 2a2540a ed78840 2a2540a ed78840 2a2540a ed78840 2a2540a ed78840 2a2540a ed78840 2a2540a ed78840 2a2540a ed78840 2a2540a ed78840 2a2540a ed78840 2a2540a ed78840 2a2540a ed78840 2a2540a ed78840 2a2540a ed78840 2a2540a ed78840 2a2540a ed78840 2a2540a ed78840 2a2540a ed78840 2a2540a ed78840 2a2540a ed78840 2a2540a ed78840 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 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 | #!/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
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?"
python chat.py --image photo.jpg --question "Kedu ihe di na foto a?" --lang ig
python chat.py --image photo.jpg --stage 1 # stage-1 captioner
python chat.py --image photo.jpg --load-in-4bit # ~8 GB GPU
python chat.py --image photo.jpg --interactive
Stages: "2b" (default, de-biased), "2" (kept for reproducibility), "1" (captioner).
Needs ~18 GB of GPU memory in bf16; --load-in-4bit fits ~8 GB GPUs such as a Colab T4.
"""
import argparse
import contextlib
import io
import os
import sys
import torch
import torch.nn as nn
from PIL import Image
REPO = "Modularcomputing/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|>"
LANGUAGES = {"en": "English", "ig": "Igbo", "yo": "Yoruba", "ha": "Hausa"}
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 isinstance(src, Image.Image):
return src.convert("RGB")
if isinstance(src, str) and 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")
def lang_name(lang):
return "English" if lang is None else LANGUAGES.get(str(lang).lower(), str(lang).title())
class AtlasVision:
def __init__(self, stage="2b", 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.stage = str(stage)
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 self.stage != "1":
from peft import PeftModel
sub = f"stage{self.stage}/lora_adapter"
if os.path.isdir(repo):
self.llm = PeftModel.from_pretrained(self.llm, os.path.join(repo, sub))
else:
self.llm = PeftModel.from_pretrained(self.llm, repo, subfolder=sub)
self.llm.eval()
self.projector = ProjectionMLP(vision_dim, text_dim)
self.projector.load_state_dict(load_file(fetch(repo, f"stage{self.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 = []
def _gen(self, max_new_tokens, temperature, **inputs):
kw = dict(max_new_tokens=max_new_tokens, repetition_penalty=1.1,
eos_token_id=self.eot_id, pad_token_id=self.pad_id)
kw.update(dict(do_sample=True, temperature=temperature, top_p=0.9) if temperature > 0 else dict(do_sample=False))
return self.llm.generate(**inputs, **kw)
@torch.no_grad()
def ask(self, image, question, max_new_tokens=256, temperature=0.0):
"""One vision-language turn, in English. Follow-ups reuse self.history."""
pv = self.processor(images=load_image(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)
out = self._gen(max_new_tokens, temperature, inputs_embeds=embeds, attention_mask=mask)
answer = self.tok.decode(out[0], skip_special_tokens=True).strip()
self.history.append((question, answer))
return answer
@torch.no_grad()
def text(self, prompt, max_new_tokens=512, temperature=0.0):
"""Plain N-ATLaS: LoRA switched off, no image. Translation, explanation, any text task."""
ids = self.tok(USER_HEADER + prompt + ASSIST_HEADER, add_special_tokens=True,
return_tensors="pt").input_ids.to(self.device)
off = self.llm.disable_adapter() if hasattr(self.llm, "disable_adapter") else contextlib.nullcontext()
with off:
out = self._gen(max_new_tokens, temperature, input_ids=ids, attention_mask=torch.ones_like(ids))
return self.tok.decode(out[0, ids.shape[1]:], skip_special_tokens=True).strip()
def translate(self, text, target, source="English", max_new_tokens=512):
target, source = lang_name(target), lang_name(source)
if target == source:
return text
return self.text(f"Translate this {source} text to {target}. Reply with only the translation.\n\n{text}",
max_new_tokens=max_new_tokens)
def chat(self, image, question, lang="en", translate_question=True, max_new_tokens=256, temperature=0.0):
"""Ask about an image in English (en), Igbo (ig), Yoruba (yo) or Hausa (ha).
Cascade on one loaded model: question -> English (plain N-ATLaS) -> AtlasVision answers in English
-> answer -> target language (plain N-ATLaS). Returns both so the English draft can be checked.
"""
name = lang_name(lang)
q_en = question if (name == "English" or not translate_question) else self.translate(question, "English", name)
a_en = self.ask(image, q_en, max_new_tokens, temperature)
answer = a_en if name == "English" else self.translate(a_en, name, "English", max_new_tokens=2 * max_new_tokens)
return {"answer": answer, "english_answer": a_en, "english_question": q_en}
def describe(self, image, lang="en", detailed=True, **kw):
q = "Describe this image in detail." if detailed else "Describe this image briefly."
return self.chat(image, q, lang=lang, translate_question=False, **kw)["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", default="2b", choices=["1", "2", "2b"])
ap.add_argument("--lang", default="en", choices=sorted(LANGUAGES))
ap.add_argument("--repo", default=REPO, help="HF repo id, or a local folder with stage1/ stage2/ stage2b/")
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)
r = model.chat(image, question, lang=a.lang, max_new_tokens=a.max_new_tokens, temperature=a.temperature)
print(f"\nQ: {question}\nA: {r['answer']}", flush=True)
if a.lang != "en":
print(f"[English draft: {r['english_answer']}]", flush=True)
while a.interactive:
try:
q = input("\nQ (empty to quit): ").strip()
except EOFError:
break
if not q:
break
r = model.chat(image, q, lang=a.lang, max_new_tokens=a.max_new_tokens, temperature=a.temperature)
print(f"A: {r['answer']}", flush=True)
if __name__ == "__main__":
sys.exit(main())
|