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
Download chat.py from Modularcomputing/AtlasVision: direct link, hf CLI and curl.
- Browser
- Download file 9.99 kB
-
https://huggingface.co/Modularcomputing/AtlasVision/resolve/main/chat.py
- Command line
-
hf download hf://Modularcomputing/AtlasVision/chat.py
-
curl -L -o chat.py https://huggingface.co/Modularcomputing/AtlasVision/resolve/main/chat.py
9.99 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 | |
| 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) | |
| 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 | |
| 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()) | |