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
Upload via givemeanode export_data
Browse files- README.md +52 -12
- chat.py +77 -27
- code/chat.py +206 -0
README.md
CHANGED
|
@@ -15,7 +15,6 @@ language:
|
|
| 15 |
datasets:
|
| 16 |
- liuhaotian/LLaVA-Pretrain
|
| 17 |
- liuhaotian/LLaVA-Instruct-150K
|
| 18 |
-
- lmms-lab/POPE
|
| 19 |
tags:
|
| 20 |
- vision-language
|
| 21 |
- multimodal
|
|
@@ -63,7 +62,7 @@ stage2/projector.safetensors # stage-2 projector (use together with the
|
|
| 63 |
stage2/lora_adapter/ # PEFT LoRA adapter for NCAIR1/N-ATLaS
|
| 64 |
stage2b/projector.safetensors # RECOMMENDED: de-biased projector
|
| 65 |
stage2b/lora_adapter/ # RECOMMENDED: de-biased LoRA adapter
|
| 66 |
-
chat.py # standalone inference script (CLI + Python API)
|
| 67 |
code/ # exact training, evaluation and launch scripts used
|
| 68 |
eval/ # raw evaluation results (JSON) and report
|
| 69 |
logs/ # per-step training logs (loss, grad norm, LR, speed)
|
|
@@ -76,7 +75,7 @@ Accept the conditions on the [N-ATLaS model page](https://huggingface.co/NCAIR1/
|
|
| 76 |
```bash
|
| 77 |
pip install -U torch transformers peft safetensors pillow accelerate huggingface_hub
|
| 78 |
export HF_TOKEN=hf_... # a token from the account that was granted N-ATLaS access
|
| 79 |
-
wget https://huggingface.co/
|
| 80 |
|
| 81 |
python chat.py --image photo.jpg --question "What is happening in this picture?" # add --stage 2b
|
| 82 |
python chat.py --image photo.jpg --question "Kedu ihe dị na foto a?" # Igbo
|
|
@@ -147,7 +146,7 @@ roughly doubles the loss — the language model is relying on the visual content
|
|
| 147 |
|
| 148 |
| Metric | Stage 1 | Stage 2 | Stage 2b (recommended) |
|
| 149 |
|---|---|---|---|
|
| 150 |
-
| Held-out LLaVA-Instruct loss (500 unseen conversations, lower is better) | 2.2521 | 1.1271 |
|
| 151 |
|
| 152 |
**POPE** (object hallucination: yes/no questions about COCO val2014 images, scored from the Yes/No token
|
| 153 |
probabilities; the benchmark is balanced 50% yes / 50% no, so a yes-ratio near 0.5 is ideal):
|
|
@@ -169,10 +168,17 @@ Stage 2b continues stage 2 for one epoch on a 150k mix drawn from the public LLa
|
|
| 169 |
short-answer VQA conversations (**57,804 "yes" vs 59,443 "no"** turns) plus 40k LLaVA-Instruct conversations
|
| 170 |
replayed so detailed description is not lost. Every POPE test image and every held-out conversation was removed
|
| 171 |
from this mix before training, so the numbers above are not contaminated; region/bounding-box tasks were dropped
|
| 172 |
-
as unsupported. Result: the yes-ratio fell to 0.56–0.63
|
| 173 |
-
|
| 174 |
-
|
| 175 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 176 |
|
| 177 |
Recall (0.88) still exceeds precision
|
| 178 |
(0.70) on the adversarial split, so a mild "yes" lean remains on the
|
|
@@ -214,7 +220,7 @@ In addition to the vehicles, there are multiple people walking around the area.
|
|
| 214 |
|
| 215 |
### Questions in Nigerian languages (stage 2b)
|
| 216 |
|
| 217 |
-
|
| 218 |
|
| 219 |
| Language | Question | Answer |
|
| 220 |
|---|---|---|
|
|
@@ -233,6 +239,36 @@ Does the stage-2 LoRA damage N-ATLaS's original text abilities? Same Igbo questi
|
|
| 233 |
|
| 234 |
Both answer fluently in Igbo, so the adapter keeps N-ATLaS's text abilities intact (responses are cut at 120 tokens).
|
| 235 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 236 |
## Limitations
|
| 237 |
|
| 238 |
- **Hallucination.** Long descriptions are fluent but often add plausible details that are not in the image
|
|
@@ -257,7 +293,11 @@ This repository contains weights trained on top of other people's models and dat
|
|
| 257 |
conditions on the [model page](https://huggingface.co/NCAIR1/N-ATLaS).
|
| 258 |
- **SigLIP2**: Apache-2.0.
|
| 259 |
- **LLaVA-Instruct-150K** (stage 2 data): CC BY-NC 4.0, generated with GPT-4 and subject to OpenAI's terms —
|
| 260 |
-
**stage-2 weights are for non-commercial research use.**
|
|
|
|
|
|
|
|
|
|
|
|
|
| 261 |
- **LLaVA-Pretrain** (stage 1 data): captions/images from LAION, Conceptual Captions and SBU under their respective terms.
|
| 262 |
- **COCO images**: Flickr images under their individual Creative Commons licenses.
|
| 263 |
|
|
@@ -273,9 +313,9 @@ This repository contains weights trained on top of other people's models and dat
|
|
| 273 |
```bibtex
|
| 274 |
@misc{atlasvision2026,
|
| 275 |
title = {AtlasVision: a vision-language extension of N-ATLaS},
|
| 276 |
-
author = {
|
| 277 |
year = {2026},
|
| 278 |
-
url = {https://huggingface.co/
|
| 279 |
}
|
| 280 |
@inproceedings{liu2023llava,
|
| 281 |
title = {Visual Instruction Tuning},
|
|
|
|
| 15 |
datasets:
|
| 16 |
- liuhaotian/LLaVA-Pretrain
|
| 17 |
- liuhaotian/LLaVA-Instruct-150K
|
|
|
|
| 18 |
tags:
|
| 19 |
- vision-language
|
| 20 |
- multimodal
|
|
|
|
| 62 |
stage2/lora_adapter/ # PEFT LoRA adapter for NCAIR1/N-ATLaS
|
| 63 |
stage2b/projector.safetensors # RECOMMENDED: de-biased projector
|
| 64 |
stage2b/lora_adapter/ # RECOMMENDED: de-biased LoRA adapter
|
| 65 |
+
chat.py # standalone inference script (CLI + Python API, incl. lang= cascade)
|
| 66 |
code/ # exact training, evaluation and launch scripts used
|
| 67 |
eval/ # raw evaluation results (JSON) and report
|
| 68 |
logs/ # per-step training logs (loss, grad norm, LR, speed)
|
|
|
|
| 75 |
```bash
|
| 76 |
pip install -U torch transformers peft safetensors pillow accelerate huggingface_hub
|
| 77 |
export HF_TOKEN=hf_... # a token from the account that was granted N-ATLaS access
|
| 78 |
+
wget https://huggingface.co/Modularcomputing/AtlasVision/resolve/main/chat.py
|
| 79 |
|
| 80 |
python chat.py --image photo.jpg --question "What is happening in this picture?" # add --stage 2b
|
| 81 |
python chat.py --image photo.jpg --question "Kedu ihe dị na foto a?" # Igbo
|
|
|
|
| 146 |
|
| 147 |
| Metric | Stage 1 | Stage 2 | Stage 2b (recommended) |
|
| 148 |
|---|---|---|---|
|
| 149 |
+
| Held-out LLaVA-Instruct loss (500 unseen conversations, lower is better) | 2.2521 | **1.1271** | 1.1602 |
|
| 150 |
|
| 151 |
**POPE** (object hallucination: yes/no questions about COCO val2014 images, scored from the Yes/No token
|
| 152 |
probabilities; the benchmark is balanced 50% yes / 50% no, so a yes-ratio near 0.5 is ideal):
|
|
|
|
| 168 |
short-answer VQA conversations (**57,804 "yes" vs 59,443 "no"** turns) plus 40k LLaVA-Instruct conversations
|
| 169 |
replayed so detailed description is not lost. Every POPE test image and every held-out conversation was removed
|
| 170 |
from this mix before training, so the numbers above are not contaminated; region/bounding-box tasks were dropped
|
| 171 |
+
as unsupported. Result: the yes-ratio fell to 0.56–0.63 and accuracy rose by **24–28 points**.
|
| 172 |
+
|
| 173 |
+
The held-out conversation loss rose slightly (1.1271 → 1.1602) — a small regression from the changed data mix.
|
| 174 |
+
Description quality looks preserved in the examples below, but we have not measured that quantitatively.
|
| 175 |
+
|
| 176 |
+
Averaged over the three splits, stage 2b scores F1 0.80. We deliberately do not put this head-to-head with
|
| 177 |
+
the number quoted for LLaVA-1.5-7B: that paper reports a single POPE figure without stating how it aggregates the
|
| 178 |
+
splits, it covers the full POPE benchmark (COCO + A-OKVQA + GQA) while we evaluate the COCO split only, and
|
| 179 |
+
independent reproductions of LLaVA-1.5 on COCO with greedy decoding land anywhere from F1 ≈ 0.78 to ≈ 0.87
|
| 180 |
+
depending on setup. Our setup also differs in scale (150k vs 665k samples), resolution (224 px vs 336 px) and base
|
| 181 |
+
model, so treat these as our own baseline, not a ranking.
|
| 182 |
|
| 183 |
Recall (0.88) still exceeds precision
|
| 184 |
(0.70) on the adversarial split, so a mild "yes" lean remains on the
|
|
|
|
| 220 |
|
| 221 |
### Questions in Nigerian languages (stage 2b)
|
| 222 |
|
| 223 |
+
In our tests, stage 2b answered the Igbo and Yoruba questions **in those languages** without any translation step, where stage 2 always replied in English. This is **one example per language**, both short and containing English loanwords ("skier", "snow"), so treat it as an encouraging signal rather than a measured result — a proper multilingual evaluation is still to come. Hausa fell back to English. For reliable answers in all three languages, use the `lang=` cascade described below.
|
| 224 |
|
| 225 |
| Language | Question | Answer |
|
| 226 |
|---|---|---|
|
|
|
|
| 239 |
|
| 240 |
Both answer fluently in Igbo, so the adapter keeps N-ATLaS's text abilities intact (responses are cut at 120 tokens).
|
| 241 |
|
| 242 |
+
## Answers in Igbo, Yoruba and Hausa (cascade)
|
| 243 |
+
|
| 244 |
+
The image training data is English, so AtlasVision reasons about images in English. `chat.py` can route any of the
|
| 245 |
+
four languages through **one** loaded model, switching the LoRA adapter off to recover the original N-ATLaS for the
|
| 246 |
+
translation steps:
|
| 247 |
+
|
| 248 |
+
1. **Question → English** with plain N-ATLaS (adapter off).
|
| 249 |
+
2. **Look and answer in English** with AtlasVision (adapter on, image tokens in).
|
| 250 |
+
3. **Answer → target language** with plain N-ATLaS (adapter off).
|
| 251 |
+
|
| 252 |
+
```python
|
| 253 |
+
from chat import AtlasVision, load_image
|
| 254 |
+
|
| 255 |
+
model = AtlasVision(stage="2b")
|
| 256 |
+
image = load_image("photo.jpg")
|
| 257 |
+
|
| 258 |
+
r = model.chat(image, "Mutane nawa ne a cikin hoton?", lang="ha")
|
| 259 |
+
print(r["answer"]) # Hausa
|
| 260 |
+
print(r["english_answer"]) # English draft, so you can check what was translated
|
| 261 |
+
|
| 262 |
+
print(model.describe(image, lang="ig")) # Igbo description
|
| 263 |
+
print(model.translate("Good morning", "yo")) # plain N-ATLaS translation, no image
|
| 264 |
+
```
|
| 265 |
+
|
| 266 |
+
From the command line: `python chat.py --image photo.jpg --question "Kedu ihe di na foto a?" --lang ig`.
|
| 267 |
+
|
| 268 |
+
Costs and caveats: each answer takes roughly 2–3× longer than English-only; translation quality is N-ATLaS's own; and
|
| 269 |
+
any error in the English answer is carried into the translation, which is why `english_answer` is always returned.
|
| 270 |
+
Pass `translate_question=False` to skip step 1 when the question is already short and simple.
|
| 271 |
+
|
| 272 |
## Limitations
|
| 273 |
|
| 274 |
- **Hallucination.** Long descriptions are fluent but often add plausible details that are not in the image
|
|
|
|
| 293 |
conditions on the [model page](https://huggingface.co/NCAIR1/N-ATLaS).
|
| 294 |
- **SigLIP2**: Apache-2.0.
|
| 295 |
- **LLaVA-Instruct-150K** (stage 2 data): CC BY-NC 4.0, generated with GPT-4 and subject to OpenAI's terms —
|
| 296 |
+
**stage-2 and stage-2b weights are for non-commercial research use.**
|
| 297 |
+
- **LLaVA-1.5 mix / VQAv2, OK-VQA, A-OKVQA** (stage 2b data): drawn from `llava_v1_5_mix665k`, which inherits the
|
| 298 |
+
CC BY-NC 4.0 terms above. The underlying short-answer sets carry their own licenses (VQAv2 annotations are
|
| 299 |
+
CC BY 4.0 over COCO images; OK-VQA and A-OKVQA are research datasets released by their authors) — check each
|
| 300 |
+
dataset's page before any redistribution or commercial use.
|
| 301 |
- **LLaVA-Pretrain** (stage 1 data): captions/images from LAION, Conceptual Captions and SBU under their respective terms.
|
| 302 |
- **COCO images**: Flickr images under their individual Creative Commons licenses.
|
| 303 |
|
|
|
|
| 313 |
```bibtex
|
| 314 |
@misc{atlasvision2026,
|
| 315 |
title = {AtlasVision: a vision-language extension of N-ATLaS},
|
| 316 |
+
author = {Modularcomputing},
|
| 317 |
year = {2026},
|
| 318 |
+
url = {https://huggingface.co/Modularcomputing/AtlasVision}
|
| 319 |
}
|
| 320 |
@inproceedings{liu2023llava,
|
| 321 |
title = {Visual Instruction Tuning},
|
chat.py
CHANGED
|
@@ -2,18 +2,19 @@
|
|
| 2 |
"""AtlasVision inference: SigLIP2 vision encoder + N-ATLaS (Llama-3 8B) via a trained MLP projector.
|
| 3 |
|
| 4 |
pip install -U torch transformers peft safetensors pillow accelerate huggingface_hub
|
| 5 |
-
# optional for 4-bit on small GPUs (e.g. Colab T4): pip install bitsandbytes
|
| 6 |
export HF_TOKEN=hf_... # needs accepted access to the gated NCAIR1/N-ATLaS
|
| 7 |
|
| 8 |
-
python chat.py --image photo.jpg --question "What is happening in this picture?"
|
| 9 |
-
python chat.py --
|
| 10 |
-
python chat.py --image
|
| 11 |
-
python chat.py --image photo.jpg --load-in-4bit
|
| 12 |
-
python chat.py --image photo.jpg --interactive
|
| 13 |
|
| 14 |
-
|
|
|
|
| 15 |
"""
|
| 16 |
import argparse
|
|
|
|
| 17 |
import io
|
| 18 |
import os
|
| 19 |
import sys
|
|
@@ -22,12 +23,13 @@ import torch
|
|
| 22 |
import torch.nn as nn
|
| 23 |
from PIL import Image
|
| 24 |
|
| 25 |
-
REPO = "
|
| 26 |
LLM = "NCAIR1/N-ATLaS"
|
| 27 |
VISION = "google/siglip2-base-patch16-224"
|
| 28 |
USER_HEADER = "<|start_header_id|>user<|end_header_id|>\n\n"
|
| 29 |
ASSIST_HEADER = "<|eot_id|><|start_header_id|>assistant<|end_header_id|>\n\n"
|
| 30 |
EOT = "<|eot_id|>"
|
|
|
|
| 31 |
|
| 32 |
|
| 33 |
class ProjectionMLP(nn.Module):
|
|
@@ -47,18 +49,25 @@ def fetch(repo, filename):
|
|
| 47 |
|
| 48 |
|
| 49 |
def load_image(src):
|
| 50 |
-
if
|
|
|
|
|
|
|
| 51 |
import urllib.request
|
| 52 |
with urllib.request.urlopen(src) as r:
|
| 53 |
return Image.open(io.BytesIO(r.read())).convert("RGB")
|
| 54 |
return Image.open(src).convert("RGB")
|
| 55 |
|
| 56 |
|
|
|
|
|
|
|
|
|
|
|
|
|
| 57 |
class AtlasVision:
|
| 58 |
-
def __init__(self, stage=
|
| 59 |
from safetensors.torch import load_file
|
| 60 |
from transformers import AutoImageProcessor, AutoModel, AutoModelForCausalLM, AutoTokenizer
|
| 61 |
|
|
|
|
| 62 |
self.device = torch.device(device or ("cuda" if torch.cuda.is_available() else "cpu"))
|
| 63 |
cuda = self.device.type == "cuda"
|
| 64 |
self.dtype = torch.bfloat16 if (not cuda or torch.cuda.is_bf16_supported()) else torch.float16
|
|
@@ -84,24 +93,32 @@ class AtlasVision:
|
|
| 84 |
self.llm.to(self.device)
|
| 85 |
text_dim = self.llm.config.hidden_size
|
| 86 |
|
| 87 |
-
if stage =
|
| 88 |
from peft import PeftModel
|
|
|
|
| 89 |
if os.path.isdir(repo):
|
| 90 |
-
self.llm = PeftModel.from_pretrained(self.llm, os.path.join(repo,
|
| 91 |
else:
|
| 92 |
-
self.llm = PeftModel.from_pretrained(self.llm, repo, subfolder=
|
| 93 |
self.llm.eval()
|
| 94 |
|
| 95 |
self.projector = ProjectionMLP(vision_dim, text_dim)
|
| 96 |
-
self.projector.load_state_dict(load_file(fetch(repo, f"stage{stage}/projector.safetensors")))
|
| 97 |
self.projector.to(self.device, dtype=torch.float32).eval()
|
| 98 |
|
| 99 |
self.prefix = torch.tensor([self.tok(USER_HEADER, add_special_tokens=True).input_ids], device=self.device)
|
| 100 |
-
self.history = []
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 101 |
|
| 102 |
@torch.no_grad()
|
| 103 |
def ask(self, image, question, max_new_tokens=256, temperature=0.0):
|
| 104 |
-
|
|
|
|
| 105 |
img = self.projector(self.vision(pixel_values=pv).last_hidden_state.float()).to(self.dtype)
|
| 106 |
|
| 107 |
text = ""
|
|
@@ -113,23 +130,52 @@ class AtlasVision:
|
|
| 113 |
emb = self.llm.get_input_embeddings()
|
| 114 |
embeds = torch.cat([emb(self.prefix).to(self.dtype), img, emb(ids).to(self.dtype)], dim=1)
|
| 115 |
mask = torch.ones(embeds.shape[:2], dtype=torch.long, device=self.device)
|
| 116 |
-
|
| 117 |
-
if temperature > 0:
|
| 118 |
-
gen.update(do_sample=True, temperature=temperature, top_p=0.9)
|
| 119 |
-
else:
|
| 120 |
-
gen.update(do_sample=False)
|
| 121 |
-
out = self.llm.generate(inputs_embeds=embeds, attention_mask=mask, **gen)
|
| 122 |
answer = self.tok.decode(out[0], skip_special_tokens=True).strip()
|
| 123 |
self.history.append((question, answer))
|
| 124 |
return answer
|
| 125 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 126 |
|
| 127 |
def main():
|
| 128 |
ap = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
|
| 129 |
ap.add_argument("--image", required=True, help="path or URL")
|
| 130 |
ap.add_argument("--question", default=None)
|
| 131 |
-
ap.add_argument("--stage",
|
| 132 |
-
ap.add_argument("--
|
|
|
|
| 133 |
ap.add_argument("--llm", default=LLM)
|
| 134 |
ap.add_argument("--vision", default=VISION)
|
| 135 |
ap.add_argument("--load-in-4bit", action="store_true")
|
|
@@ -138,10 +184,13 @@ def main():
|
|
| 138 |
ap.add_argument("--interactive", action="store_true")
|
| 139 |
a = ap.parse_args()
|
| 140 |
|
| 141 |
-
question = a.question or ("Describe this image briefly." if a.stage == 1 else "Describe this image in detail.")
|
| 142 |
model = AtlasVision(a.stage, a.repo, a.llm, a.vision, a.load_in_4bit)
|
| 143 |
image = load_image(a.image)
|
| 144 |
-
|
|
|
|
|
|
|
|
|
|
| 145 |
while a.interactive:
|
| 146 |
try:
|
| 147 |
q = input("\nQ (empty to quit): ").strip()
|
|
@@ -149,7 +198,8 @@ def main():
|
|
| 149 |
break
|
| 150 |
if not q:
|
| 151 |
break
|
| 152 |
-
|
|
|
|
| 153 |
|
| 154 |
|
| 155 |
if __name__ == "__main__":
|
|
|
|
| 2 |
"""AtlasVision inference: SigLIP2 vision encoder + N-ATLaS (Llama-3 8B) via a trained MLP projector.
|
| 3 |
|
| 4 |
pip install -U torch transformers peft safetensors pillow accelerate huggingface_hub
|
|
|
|
| 5 |
export HF_TOKEN=hf_... # needs accepted access to the gated NCAIR1/N-ATLaS
|
| 6 |
|
| 7 |
+
python chat.py --image photo.jpg --question "What is happening in this picture?"
|
| 8 |
+
python chat.py --image photo.jpg --question "Kedu ihe di na foto a?" --lang ig
|
| 9 |
+
python chat.py --image photo.jpg --stage 1 # stage-1 captioner
|
| 10 |
+
python chat.py --image photo.jpg --load-in-4bit # ~8 GB GPU
|
| 11 |
+
python chat.py --image photo.jpg --interactive
|
| 12 |
|
| 13 |
+
Stages: "2b" (default, de-biased), "2" (kept for reproducibility), "1" (captioner).
|
| 14 |
+
Needs ~18 GB of GPU memory in bf16; --load-in-4bit fits ~8 GB GPUs such as a Colab T4.
|
| 15 |
"""
|
| 16 |
import argparse
|
| 17 |
+
import contextlib
|
| 18 |
import io
|
| 19 |
import os
|
| 20 |
import sys
|
|
|
|
| 23 |
import torch.nn as nn
|
| 24 |
from PIL import Image
|
| 25 |
|
| 26 |
+
REPO = "Modularcomputing/AtlasVision"
|
| 27 |
LLM = "NCAIR1/N-ATLaS"
|
| 28 |
VISION = "google/siglip2-base-patch16-224"
|
| 29 |
USER_HEADER = "<|start_header_id|>user<|end_header_id|>\n\n"
|
| 30 |
ASSIST_HEADER = "<|eot_id|><|start_header_id|>assistant<|end_header_id|>\n\n"
|
| 31 |
EOT = "<|eot_id|>"
|
| 32 |
+
LANGUAGES = {"en": "English", "ig": "Igbo", "yo": "Yoruba", "ha": "Hausa"}
|
| 33 |
|
| 34 |
|
| 35 |
class ProjectionMLP(nn.Module):
|
|
|
|
| 49 |
|
| 50 |
|
| 51 |
def load_image(src):
|
| 52 |
+
if isinstance(src, Image.Image):
|
| 53 |
+
return src.convert("RGB")
|
| 54 |
+
if isinstance(src, str) and src.startswith(("http://", "https://")):
|
| 55 |
import urllib.request
|
| 56 |
with urllib.request.urlopen(src) as r:
|
| 57 |
return Image.open(io.BytesIO(r.read())).convert("RGB")
|
| 58 |
return Image.open(src).convert("RGB")
|
| 59 |
|
| 60 |
|
| 61 |
+
def lang_name(lang):
|
| 62 |
+
return "English" if lang is None else LANGUAGES.get(str(lang).lower(), str(lang).title())
|
| 63 |
+
|
| 64 |
+
|
| 65 |
class AtlasVision:
|
| 66 |
+
def __init__(self, stage="2b", repo=REPO, llm=LLM, vision=VISION, load_in_4bit=False, device=None):
|
| 67 |
from safetensors.torch import load_file
|
| 68 |
from transformers import AutoImageProcessor, AutoModel, AutoModelForCausalLM, AutoTokenizer
|
| 69 |
|
| 70 |
+
self.stage = str(stage)
|
| 71 |
self.device = torch.device(device or ("cuda" if torch.cuda.is_available() else "cpu"))
|
| 72 |
cuda = self.device.type == "cuda"
|
| 73 |
self.dtype = torch.bfloat16 if (not cuda or torch.cuda.is_bf16_supported()) else torch.float16
|
|
|
|
| 93 |
self.llm.to(self.device)
|
| 94 |
text_dim = self.llm.config.hidden_size
|
| 95 |
|
| 96 |
+
if self.stage != "1":
|
| 97 |
from peft import PeftModel
|
| 98 |
+
sub = f"stage{self.stage}/lora_adapter"
|
| 99 |
if os.path.isdir(repo):
|
| 100 |
+
self.llm = PeftModel.from_pretrained(self.llm, os.path.join(repo, sub))
|
| 101 |
else:
|
| 102 |
+
self.llm = PeftModel.from_pretrained(self.llm, repo, subfolder=sub)
|
| 103 |
self.llm.eval()
|
| 104 |
|
| 105 |
self.projector = ProjectionMLP(vision_dim, text_dim)
|
| 106 |
+
self.projector.load_state_dict(load_file(fetch(repo, f"stage{self.stage}/projector.safetensors")))
|
| 107 |
self.projector.to(self.device, dtype=torch.float32).eval()
|
| 108 |
|
| 109 |
self.prefix = torch.tensor([self.tok(USER_HEADER, add_special_tokens=True).input_ids], device=self.device)
|
| 110 |
+
self.history = []
|
| 111 |
+
|
| 112 |
+
def _gen(self, max_new_tokens, temperature, **inputs):
|
| 113 |
+
kw = dict(max_new_tokens=max_new_tokens, repetition_penalty=1.1,
|
| 114 |
+
eos_token_id=self.eot_id, pad_token_id=self.pad_id)
|
| 115 |
+
kw.update(dict(do_sample=True, temperature=temperature, top_p=0.9) if temperature > 0 else dict(do_sample=False))
|
| 116 |
+
return self.llm.generate(**inputs, **kw)
|
| 117 |
|
| 118 |
@torch.no_grad()
|
| 119 |
def ask(self, image, question, max_new_tokens=256, temperature=0.0):
|
| 120 |
+
"""One vision-language turn, in English. Follow-ups reuse self.history."""
|
| 121 |
+
pv = self.processor(images=load_image(image), return_tensors="pt").pixel_values.to(self.device, self.dtype)
|
| 122 |
img = self.projector(self.vision(pixel_values=pv).last_hidden_state.float()).to(self.dtype)
|
| 123 |
|
| 124 |
text = ""
|
|
|
|
| 130 |
emb = self.llm.get_input_embeddings()
|
| 131 |
embeds = torch.cat([emb(self.prefix).to(self.dtype), img, emb(ids).to(self.dtype)], dim=1)
|
| 132 |
mask = torch.ones(embeds.shape[:2], dtype=torch.long, device=self.device)
|
| 133 |
+
out = self._gen(max_new_tokens, temperature, inputs_embeds=embeds, attention_mask=mask)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 134 |
answer = self.tok.decode(out[0], skip_special_tokens=True).strip()
|
| 135 |
self.history.append((question, answer))
|
| 136 |
return answer
|
| 137 |
|
| 138 |
+
@torch.no_grad()
|
| 139 |
+
def text(self, prompt, max_new_tokens=512, temperature=0.0):
|
| 140 |
+
"""Plain N-ATLaS: LoRA switched off, no image. Translation, explanation, any text task."""
|
| 141 |
+
ids = self.tok(USER_HEADER + prompt + ASSIST_HEADER, add_special_tokens=True,
|
| 142 |
+
return_tensors="pt").input_ids.to(self.device)
|
| 143 |
+
off = self.llm.disable_adapter() if hasattr(self.llm, "disable_adapter") else contextlib.nullcontext()
|
| 144 |
+
with off:
|
| 145 |
+
out = self._gen(max_new_tokens, temperature, input_ids=ids, attention_mask=torch.ones_like(ids))
|
| 146 |
+
return self.tok.decode(out[0, ids.shape[1]:], skip_special_tokens=True).strip()
|
| 147 |
+
|
| 148 |
+
def translate(self, text, target, source="English", max_new_tokens=512):
|
| 149 |
+
target, source = lang_name(target), lang_name(source)
|
| 150 |
+
if target == source:
|
| 151 |
+
return text
|
| 152 |
+
return self.text(f"Translate this {source} text to {target}. Reply with only the translation.\n\n{text}",
|
| 153 |
+
max_new_tokens=max_new_tokens)
|
| 154 |
+
|
| 155 |
+
def chat(self, image, question, lang="en", translate_question=True, max_new_tokens=256, temperature=0.0):
|
| 156 |
+
"""Ask about an image in English (en), Igbo (ig), Yoruba (yo) or Hausa (ha).
|
| 157 |
+
|
| 158 |
+
Cascade on one loaded model: question -> English (plain N-ATLaS) -> AtlasVision answers in English
|
| 159 |
+
-> answer -> target language (plain N-ATLaS). Returns both so the English draft can be checked.
|
| 160 |
+
"""
|
| 161 |
+
name = lang_name(lang)
|
| 162 |
+
q_en = question if (name == "English" or not translate_question) else self.translate(question, "English", name)
|
| 163 |
+
a_en = self.ask(image, q_en, max_new_tokens, temperature)
|
| 164 |
+
answer = a_en if name == "English" else self.translate(a_en, name, "English", max_new_tokens=2 * max_new_tokens)
|
| 165 |
+
return {"answer": answer, "english_answer": a_en, "english_question": q_en}
|
| 166 |
+
|
| 167 |
+
def describe(self, image, lang="en", detailed=True, **kw):
|
| 168 |
+
q = "Describe this image in detail." if detailed else "Describe this image briefly."
|
| 169 |
+
return self.chat(image, q, lang=lang, translate_question=False, **kw)["answer"]
|
| 170 |
+
|
| 171 |
|
| 172 |
def main():
|
| 173 |
ap = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
|
| 174 |
ap.add_argument("--image", required=True, help="path or URL")
|
| 175 |
ap.add_argument("--question", default=None)
|
| 176 |
+
ap.add_argument("--stage", default="2b", choices=["1", "2", "2b"])
|
| 177 |
+
ap.add_argument("--lang", default="en", choices=sorted(LANGUAGES))
|
| 178 |
+
ap.add_argument("--repo", default=REPO, help="HF repo id, or a local folder with stage1/ stage2/ stage2b/")
|
| 179 |
ap.add_argument("--llm", default=LLM)
|
| 180 |
ap.add_argument("--vision", default=VISION)
|
| 181 |
ap.add_argument("--load-in-4bit", action="store_true")
|
|
|
|
| 184 |
ap.add_argument("--interactive", action="store_true")
|
| 185 |
a = ap.parse_args()
|
| 186 |
|
| 187 |
+
question = a.question or ("Describe this image briefly." if a.stage == "1" else "Describe this image in detail.")
|
| 188 |
model = AtlasVision(a.stage, a.repo, a.llm, a.vision, a.load_in_4bit)
|
| 189 |
image = load_image(a.image)
|
| 190 |
+
r = model.chat(image, question, lang=a.lang, max_new_tokens=a.max_new_tokens, temperature=a.temperature)
|
| 191 |
+
print(f"\nQ: {question}\nA: {r['answer']}", flush=True)
|
| 192 |
+
if a.lang != "en":
|
| 193 |
+
print(f"[English draft: {r['english_answer']}]", flush=True)
|
| 194 |
while a.interactive:
|
| 195 |
try:
|
| 196 |
q = input("\nQ (empty to quit): ").strip()
|
|
|
|
| 198 |
break
|
| 199 |
if not q:
|
| 200 |
break
|
| 201 |
+
r = model.chat(image, q, lang=a.lang, max_new_tokens=a.max_new_tokens, temperature=a.temperature)
|
| 202 |
+
print(f"A: {r['answer']}", flush=True)
|
| 203 |
|
| 204 |
|
| 205 |
if __name__ == "__main__":
|
code/chat.py
ADDED
|
@@ -0,0 +1,206 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/usr/bin/env python3
|
| 2 |
+
"""AtlasVision inference: SigLIP2 vision encoder + N-ATLaS (Llama-3 8B) via a trained MLP projector.
|
| 3 |
+
|
| 4 |
+
pip install -U torch transformers peft safetensors pillow accelerate huggingface_hub
|
| 5 |
+
export HF_TOKEN=hf_... # needs accepted access to the gated NCAIR1/N-ATLaS
|
| 6 |
+
|
| 7 |
+
python chat.py --image photo.jpg --question "What is happening in this picture?"
|
| 8 |
+
python chat.py --image photo.jpg --question "Kedu ihe di na foto a?" --lang ig
|
| 9 |
+
python chat.py --image photo.jpg --stage 1 # stage-1 captioner
|
| 10 |
+
python chat.py --image photo.jpg --load-in-4bit # ~8 GB GPU
|
| 11 |
+
python chat.py --image photo.jpg --interactive
|
| 12 |
+
|
| 13 |
+
Stages: "2b" (default, de-biased), "2" (kept for reproducibility), "1" (captioner).
|
| 14 |
+
Needs ~18 GB of GPU memory in bf16; --load-in-4bit fits ~8 GB GPUs such as a Colab T4.
|
| 15 |
+
"""
|
| 16 |
+
import argparse
|
| 17 |
+
import contextlib
|
| 18 |
+
import io
|
| 19 |
+
import os
|
| 20 |
+
import sys
|
| 21 |
+
|
| 22 |
+
import torch
|
| 23 |
+
import torch.nn as nn
|
| 24 |
+
from PIL import Image
|
| 25 |
+
|
| 26 |
+
REPO = "Modularcomputing/AtlasVision"
|
| 27 |
+
LLM = "NCAIR1/N-ATLaS"
|
| 28 |
+
VISION = "google/siglip2-base-patch16-224"
|
| 29 |
+
USER_HEADER = "<|start_header_id|>user<|end_header_id|>\n\n"
|
| 30 |
+
ASSIST_HEADER = "<|eot_id|><|start_header_id|>assistant<|end_header_id|>\n\n"
|
| 31 |
+
EOT = "<|eot_id|>"
|
| 32 |
+
LANGUAGES = {"en": "English", "ig": "Igbo", "yo": "Yoruba", "ha": "Hausa"}
|
| 33 |
+
|
| 34 |
+
|
| 35 |
+
class ProjectionMLP(nn.Module):
|
| 36 |
+
def __init__(self, vision_dim, text_dim):
|
| 37 |
+
super().__init__()
|
| 38 |
+
self.net = nn.Sequential(nn.Linear(vision_dim, text_dim), nn.GELU(), nn.Linear(text_dim, text_dim))
|
| 39 |
+
|
| 40 |
+
def forward(self, x):
|
| 41 |
+
return self.net(x)
|
| 42 |
+
|
| 43 |
+
|
| 44 |
+
def fetch(repo, filename):
|
| 45 |
+
if os.path.isdir(repo):
|
| 46 |
+
return os.path.join(repo, filename)
|
| 47 |
+
from huggingface_hub import hf_hub_download
|
| 48 |
+
return hf_hub_download(repo, filename)
|
| 49 |
+
|
| 50 |
+
|
| 51 |
+
def load_image(src):
|
| 52 |
+
if isinstance(src, Image.Image):
|
| 53 |
+
return src.convert("RGB")
|
| 54 |
+
if isinstance(src, str) and src.startswith(("http://", "https://")):
|
| 55 |
+
import urllib.request
|
| 56 |
+
with urllib.request.urlopen(src) as r:
|
| 57 |
+
return Image.open(io.BytesIO(r.read())).convert("RGB")
|
| 58 |
+
return Image.open(src).convert("RGB")
|
| 59 |
+
|
| 60 |
+
|
| 61 |
+
def lang_name(lang):
|
| 62 |
+
return "English" if lang is None else LANGUAGES.get(str(lang).lower(), str(lang).title())
|
| 63 |
+
|
| 64 |
+
|
| 65 |
+
class AtlasVision:
|
| 66 |
+
def __init__(self, stage="2b", repo=REPO, llm=LLM, vision=VISION, load_in_4bit=False, device=None):
|
| 67 |
+
from safetensors.torch import load_file
|
| 68 |
+
from transformers import AutoImageProcessor, AutoModel, AutoModelForCausalLM, AutoTokenizer
|
| 69 |
+
|
| 70 |
+
self.stage = str(stage)
|
| 71 |
+
self.device = torch.device(device or ("cuda" if torch.cuda.is_available() else "cpu"))
|
| 72 |
+
cuda = self.device.type == "cuda"
|
| 73 |
+
self.dtype = torch.bfloat16 if (not cuda or torch.cuda.is_bf16_supported()) else torch.float16
|
| 74 |
+
|
| 75 |
+
self.tok = AutoTokenizer.from_pretrained(llm)
|
| 76 |
+
self.pad_id = self.tok.pad_token_id if self.tok.pad_token_id is not None else self.tok.eos_token_id
|
| 77 |
+
self.eot_id = self.tok.convert_tokens_to_ids(EOT)
|
| 78 |
+
self.processor = AutoImageProcessor.from_pretrained(vision)
|
| 79 |
+
|
| 80 |
+
full_vision = AutoModel.from_pretrained(vision, dtype=self.dtype)
|
| 81 |
+
self.vision = full_vision.vision_model.to(self.device).eval()
|
| 82 |
+
vision_dim = full_vision.config.vision_config.hidden_size
|
| 83 |
+
del full_vision
|
| 84 |
+
|
| 85 |
+
kw = {"dtype": self.dtype}
|
| 86 |
+
if load_in_4bit:
|
| 87 |
+
from transformers import BitsAndBytesConfig
|
| 88 |
+
kw["quantization_config"] = BitsAndBytesConfig(load_in_4bit=True, bnb_4bit_quant_type="nf4",
|
| 89 |
+
bnb_4bit_compute_dtype=self.dtype)
|
| 90 |
+
kw["device_map"] = {"": self.device.index or 0}
|
| 91 |
+
self.llm = AutoModelForCausalLM.from_pretrained(llm, **kw)
|
| 92 |
+
if not load_in_4bit:
|
| 93 |
+
self.llm.to(self.device)
|
| 94 |
+
text_dim = self.llm.config.hidden_size
|
| 95 |
+
|
| 96 |
+
if self.stage != "1":
|
| 97 |
+
from peft import PeftModel
|
| 98 |
+
sub = f"stage{self.stage}/lora_adapter"
|
| 99 |
+
if os.path.isdir(repo):
|
| 100 |
+
self.llm = PeftModel.from_pretrained(self.llm, os.path.join(repo, sub))
|
| 101 |
+
else:
|
| 102 |
+
self.llm = PeftModel.from_pretrained(self.llm, repo, subfolder=sub)
|
| 103 |
+
self.llm.eval()
|
| 104 |
+
|
| 105 |
+
self.projector = ProjectionMLP(vision_dim, text_dim)
|
| 106 |
+
self.projector.load_state_dict(load_file(fetch(repo, f"stage{self.stage}/projector.safetensors")))
|
| 107 |
+
self.projector.to(self.device, dtype=torch.float32).eval()
|
| 108 |
+
|
| 109 |
+
self.prefix = torch.tensor([self.tok(USER_HEADER, add_special_tokens=True).input_ids], device=self.device)
|
| 110 |
+
self.history = []
|
| 111 |
+
|
| 112 |
+
def _gen(self, max_new_tokens, temperature, **inputs):
|
| 113 |
+
kw = dict(max_new_tokens=max_new_tokens, repetition_penalty=1.1,
|
| 114 |
+
eos_token_id=self.eot_id, pad_token_id=self.pad_id)
|
| 115 |
+
kw.update(dict(do_sample=True, temperature=temperature, top_p=0.9) if temperature > 0 else dict(do_sample=False))
|
| 116 |
+
return self.llm.generate(**inputs, **kw)
|
| 117 |
+
|
| 118 |
+
@torch.no_grad()
|
| 119 |
+
def ask(self, image, question, max_new_tokens=256, temperature=0.0):
|
| 120 |
+
"""One vision-language turn, in English. Follow-ups reuse self.history."""
|
| 121 |
+
pv = self.processor(images=load_image(image), return_tensors="pt").pixel_values.to(self.device, self.dtype)
|
| 122 |
+
img = self.projector(self.vision(pixel_values=pv).last_hidden_state.float()).to(self.dtype)
|
| 123 |
+
|
| 124 |
+
text = ""
|
| 125 |
+
for i, (q, a) in enumerate(self.history):
|
| 126 |
+
text += (q if i == 0 else USER_HEADER + q) + ASSIST_HEADER + a + EOT
|
| 127 |
+
text += (question if not self.history else USER_HEADER + question) + ASSIST_HEADER
|
| 128 |
+
ids = torch.tensor([self.tok(text, add_special_tokens=False).input_ids], device=self.device)
|
| 129 |
+
|
| 130 |
+
emb = self.llm.get_input_embeddings()
|
| 131 |
+
embeds = torch.cat([emb(self.prefix).to(self.dtype), img, emb(ids).to(self.dtype)], dim=1)
|
| 132 |
+
mask = torch.ones(embeds.shape[:2], dtype=torch.long, device=self.device)
|
| 133 |
+
out = self._gen(max_new_tokens, temperature, inputs_embeds=embeds, attention_mask=mask)
|
| 134 |
+
answer = self.tok.decode(out[0], skip_special_tokens=True).strip()
|
| 135 |
+
self.history.append((question, answer))
|
| 136 |
+
return answer
|
| 137 |
+
|
| 138 |
+
@torch.no_grad()
|
| 139 |
+
def text(self, prompt, max_new_tokens=512, temperature=0.0):
|
| 140 |
+
"""Plain N-ATLaS: LoRA switched off, no image. Translation, explanation, any text task."""
|
| 141 |
+
ids = self.tok(USER_HEADER + prompt + ASSIST_HEADER, add_special_tokens=True,
|
| 142 |
+
return_tensors="pt").input_ids.to(self.device)
|
| 143 |
+
off = self.llm.disable_adapter() if hasattr(self.llm, "disable_adapter") else contextlib.nullcontext()
|
| 144 |
+
with off:
|
| 145 |
+
out = self._gen(max_new_tokens, temperature, input_ids=ids, attention_mask=torch.ones_like(ids))
|
| 146 |
+
return self.tok.decode(out[0, ids.shape[1]:], skip_special_tokens=True).strip()
|
| 147 |
+
|
| 148 |
+
def translate(self, text, target, source="English", max_new_tokens=512):
|
| 149 |
+
target, source = lang_name(target), lang_name(source)
|
| 150 |
+
if target == source:
|
| 151 |
+
return text
|
| 152 |
+
return self.text(f"Translate this {source} text to {target}. Reply with only the translation.\n\n{text}",
|
| 153 |
+
max_new_tokens=max_new_tokens)
|
| 154 |
+
|
| 155 |
+
def chat(self, image, question, lang="en", translate_question=True, max_new_tokens=256, temperature=0.0):
|
| 156 |
+
"""Ask about an image in English (en), Igbo (ig), Yoruba (yo) or Hausa (ha).
|
| 157 |
+
|
| 158 |
+
Cascade on one loaded model: question -> English (plain N-ATLaS) -> AtlasVision answers in English
|
| 159 |
+
-> answer -> target language (plain N-ATLaS). Returns both so the English draft can be checked.
|
| 160 |
+
"""
|
| 161 |
+
name = lang_name(lang)
|
| 162 |
+
q_en = question if (name == "English" or not translate_question) else self.translate(question, "English", name)
|
| 163 |
+
a_en = self.ask(image, q_en, max_new_tokens, temperature)
|
| 164 |
+
answer = a_en if name == "English" else self.translate(a_en, name, "English", max_new_tokens=2 * max_new_tokens)
|
| 165 |
+
return {"answer": answer, "english_answer": a_en, "english_question": q_en}
|
| 166 |
+
|
| 167 |
+
def describe(self, image, lang="en", detailed=True, **kw):
|
| 168 |
+
q = "Describe this image in detail." if detailed else "Describe this image briefly."
|
| 169 |
+
return self.chat(image, q, lang=lang, translate_question=False, **kw)["answer"]
|
| 170 |
+
|
| 171 |
+
|
| 172 |
+
def main():
|
| 173 |
+
ap = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
|
| 174 |
+
ap.add_argument("--image", required=True, help="path or URL")
|
| 175 |
+
ap.add_argument("--question", default=None)
|
| 176 |
+
ap.add_argument("--stage", default="2b", choices=["1", "2", "2b"])
|
| 177 |
+
ap.add_argument("--lang", default="en", choices=sorted(LANGUAGES))
|
| 178 |
+
ap.add_argument("--repo", default=REPO, help="HF repo id, or a local folder with stage1/ stage2/ stage2b/")
|
| 179 |
+
ap.add_argument("--llm", default=LLM)
|
| 180 |
+
ap.add_argument("--vision", default=VISION)
|
| 181 |
+
ap.add_argument("--load-in-4bit", action="store_true")
|
| 182 |
+
ap.add_argument("--max-new-tokens", type=int, default=256)
|
| 183 |
+
ap.add_argument("--temperature", type=float, default=0.0)
|
| 184 |
+
ap.add_argument("--interactive", action="store_true")
|
| 185 |
+
a = ap.parse_args()
|
| 186 |
+
|
| 187 |
+
question = a.question or ("Describe this image briefly." if a.stage == "1" else "Describe this image in detail.")
|
| 188 |
+
model = AtlasVision(a.stage, a.repo, a.llm, a.vision, a.load_in_4bit)
|
| 189 |
+
image = load_image(a.image)
|
| 190 |
+
r = model.chat(image, question, lang=a.lang, max_new_tokens=a.max_new_tokens, temperature=a.temperature)
|
| 191 |
+
print(f"\nQ: {question}\nA: {r['answer']}", flush=True)
|
| 192 |
+
if a.lang != "en":
|
| 193 |
+
print(f"[English draft: {r['english_answer']}]", flush=True)
|
| 194 |
+
while a.interactive:
|
| 195 |
+
try:
|
| 196 |
+
q = input("\nQ (empty to quit): ").strip()
|
| 197 |
+
except EOFError:
|
| 198 |
+
break
|
| 199 |
+
if not q:
|
| 200 |
+
break
|
| 201 |
+
r = model.chat(image, q, lang=a.lang, max_new_tokens=a.max_new_tokens, temperature=a.temperature)
|
| 202 |
+
print(f"A: {r['answer']}", flush=True)
|
| 203 |
+
|
| 204 |
+
|
| 205 |
+
if __name__ == "__main__":
|
| 206 |
+
sys.exit(main())
|