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: 7,255 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 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 | #!/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())
|