Image-Text-to-Text
PEFT
Safetensors
vision-language
multimodal
llava
lora
siglip2
n-atlas
nigerian-languages
UncleanCode commited on
Commit
ed78840
·
verified ·
1 Parent(s): c8317c0

Upload via givemeanode export_data

Browse files
Files changed (3) hide show
  1. README.md +52 -12
  2. chat.py +77 -27
  3. 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/FUTO-NIGERIA/AtlasVision/resolve/main/chat.py
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 | **1.1602** |
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, accuracy rose by **24–28 points**, and the held-out
173
- conversation loss improved as well (1.1271 → 1.1602),
174
- so description quality did not regress. For reference, LLaVA-1.5-7B reports POPE F1 ≈ 0.86 from 665k samples at
175
- 336 px; AtlasVision reaches 0.83 from 150k samples at 224 px.
 
 
 
 
 
 
 
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
- Stage 2b answers Igbo and Yoruba questions **in those languages** without any translation step (stage 2 always replied in English). Hausa still falls back to English, so the `lang=` cascade below remains the reliable route for all three.
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 = {FUTO-NIGERIA},
277
  year = {2026},
278
- url = {https://huggingface.co/FUTO-NIGERIA/AtlasVision}
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?" # stage 2
9
- python chat.py --stage 1 --image photo.jpg # stage 1 captioner
10
- python chat.py --image https://example.com/cat.jpg --question "Kedu ihe dị na foto a?"
11
- python chat.py --image photo.jpg --load-in-4bit # ~8 GB GPU
12
- python chat.py --image photo.jpg --interactive # several questions
13
 
14
- Needs ~18 GB of GPU memory in bf16 (A100, L4, RTX 4090...); --load-in-4bit fits ~8 GB GPUs such as a Colab T4.
 
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 = "FUTO-NIGERIA/AtlasVision"
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 src.startswith(("http://", "https://")):
 
 
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=2, repo=REPO, llm=LLM, vision=VISION, load_in_4bit=False, device=None):
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 == 2:
88
  from peft import PeftModel
 
89
  if os.path.isdir(repo):
90
- self.llm = PeftModel.from_pretrained(self.llm, os.path.join(repo, "stage2/lora_adapter"))
91
  else:
92
- self.llm = PeftModel.from_pretrained(self.llm, repo, subfolder="stage2/lora_adapter")
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 = [] # (question, answer) turns about the current image
 
 
 
 
 
 
101
 
102
  @torch.no_grad()
103
  def ask(self, image, question, max_new_tokens=256, temperature=0.0):
104
- pv = self.processor(images=image, return_tensors="pt").pixel_values.to(self.device, self.dtype)
 
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
- gen = dict(max_new_tokens=max_new_tokens, repetition_penalty=1.1, eos_token_id=self.eot_id, pad_token_id=self.pad_id)
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", type=int, default=2, choices=[1, 2])
132
- ap.add_argument("--repo", default=REPO, help="HF repo id, or a local folder containing stage1/ and stage2/")
 
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
- print(f"\nQ: {question}\nA: {model.ask(image, question, a.max_new_tokens, a.temperature)}", flush=True)
 
 
 
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
- print(f"A: {model.ask(image, q, a.max_new_tokens, a.temperature)}", flush=True)
 
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())