| """ |
| Phase 3 — Track C: dataset loading, indexing, VQA metric and agent tools for |
| OCR-focused question answering over TextVQA + a student Portuguese dataset. |
| """ |
| import json |
| import re |
| from pathlib import Path |
| from typing import Optional |
|
|
| import pandas as pd |
| from PIL import Image |
| from opensearchpy import helpers |
|
|
|
|
| |
| |
| |
|
|
| def load_textvqa_sample(n=1000, split="validation", seed=42): |
| """ |
| Load the TextVQA split and return a randomly sampled subset (one row per question). |
| Returns a HuggingFace Dataset with embedded PIL images. |
| """ |
| from datasets import load_dataset |
| ds = load_dataset("lmms-lab/textvqa", split=split) |
| if n and n < len(ds): |
| ds = ds.shuffle(seed=seed).select(range(n)) |
| return ds |
|
|
|
|
| def textvqa_records(ds, dataset_name="textvqa", language="en"): |
| """ |
| Yield normalised records from a TextVQA dataset: |
| {image_id, question_id, question, answers, language, dataset, ocr_tokens_ref}. |
| 'ocr_tokens_ref' is the dataset's own Rosetta-OCR tokens, used only as a |
| reference for CER/WER — not as our extracted OCR. |
| """ |
| for item in ds: |
| yield { |
| "image_id": str(item["image_id"]), |
| "question_id": str(item.get("question_id", "")), |
| "question": item["question"], |
| "answers": list(item.get("answers", [])), |
| "language": language, |
| "dataset": dataset_name, |
| "image_url": item.get("flickr_300k_url") or item.get("flickr_original_url") or "", |
| "ocr_tokens_ref": " ".join(item.get("ocr_tokens", []) or []), |
| } |
|
|
|
|
| def build_textvqa_image_lookup(ds) -> dict: |
| """Build {image_id: PIL.Image} from a TextVQA dataset.""" |
| lookup = {} |
| for item in ds: |
| lookup[str(item["image_id"])] = item["image"].convert("RGB") |
| return lookup |
|
|
|
|
| def load_pt_dataset(root="data/pt_textvqa"): |
| """ |
| Load the student-created Portuguese dataset. |
| Expects root/annotations.json (list of {image_id, source, question, answer, language}) |
| and root/images/<image_id>.<ext>. |
| Returns (records_list, {image_id: PIL.Image}). |
| """ |
| root = Path(root) |
| ann_path = root / "annotations.json" |
| if not ann_path.exists(): |
| raise FileNotFoundError(f"No annotations at {ann_path} — populate the PT dataset first.") |
| annotations = json.loads(ann_path.read_text(encoding="utf-8")) |
|
|
| records, lookup = [], {} |
| for a in annotations: |
| iid = str(a["image_id"]) |
| matches = list((root / "images").glob(f"{iid}.*")) |
| if not matches: |
| print(f" warning: no image file for {iid}") |
| continue |
| lookup[iid] = Image.open(matches[0]).convert("RGB") |
| records.append({ |
| "image_id": iid, |
| "question_id": iid, |
| "question": a["question"], |
| "answers": [a["answer"]] if isinstance(a.get("answer"), str) else list(a.get("answer", [])), |
| "language": a.get("language", "pt"), |
| "dataset": "pt", |
| "ocr_tokens_ref": a.get("ocr_text", ""), |
| }) |
| return records, lookup |
|
|
|
|
| |
| |
| |
|
|
| def create_textvqa_index(client, index_name, image_vec_dim=768, |
| k1=1.2, b=0.75, force_recreate=False, num_shards=1): |
| """Create the OCR-QA index: BM25 over ocr_text + SigLIP image vectors + metadata. |
| |
| num_shards defaults to 1: this index holds only ~1k docs, and the shared |
| classroom cluster has a global shard cap. |
| """ |
| knn_method = { |
| "name": "hnsw", |
| "space_type": "innerproduct", |
| "engine": "faiss", |
| "parameters": {"ef_construction": 256, "m": 48}, |
| } |
| index_body = { |
| "settings": { |
| "index": { |
| "number_of_replicas": 0, |
| "number_of_shards": num_shards, |
| "refresh_interval": "-1", |
| "knn": "true", |
| "similarity": {"custom_bm25": {"type": "BM25", "k1": k1, "b": b}}, |
| } |
| }, |
| "mappings": { |
| "dynamic": "strict", |
| "properties": { |
| "image_id": {"type": "keyword"}, |
| "question_id": {"type": "keyword"}, |
| "dataset": {"type": "keyword"}, |
| "language": {"type": "keyword"}, |
| "question": {"type": "text", "analyzer": "standard"}, |
| "answers": {"type": "keyword"}, |
| "image_url": {"type": "keyword"}, |
| "ocr_text": {"type": "text", "analyzer": "standard", "similarity": "custom_bm25"}, |
| "ocr_tokens": {"type": "keyword"}, |
| "image_vec": {"type": "knn_vector", "dimension": image_vec_dim, "method": knn_method}, |
| }, |
| }, |
| } |
| if client.indices.exists(index=index_name): |
| if not force_recreate: |
| print(f"Index '{index_name}' already exists. Skipping (force_recreate=True to replace).") |
| return |
| client.indices.delete(index=index_name) |
| print(f"Deleted existing index '{index_name}'.") |
| client.indices.create(index=index_name, body=index_body) |
| print(f"Index '{index_name}' created.") |
|
|
|
|
| def generate_textvqa_docs(records, index_name, ocr_df=None, image_vec_df=None): |
| """ |
| Yield index-ready docs. records: list of normalised dicts (see *_records helpers). |
| ocr_df: DataFrame [image_id, ocr_text, ocr_tokens]. image_vec_df: [image_id, embedding]. |
| """ |
| ocr_lookup = {} |
| if ocr_df is not None: |
| ocr_lookup = { |
| r.image_id: (r.ocr_text, list(r.ocr_tokens)) |
| for r in ocr_df.itertuples() |
| } |
| vec_lookup = {} |
| if image_vec_df is not None: |
| vec_lookup = dict(zip(image_vec_df["image_id"], image_vec_df["embedding"])) |
|
|
| for rec in records: |
| iid = rec["image_id"] |
| doc = { |
| "_index": index_name, |
| "image_id": iid, |
| "question_id": rec["question_id"], |
| "dataset": rec["dataset"], |
| "language": rec["language"], |
| "question": rec["question"], |
| "answers": rec["answers"], |
| "image_url": rec.get("image_url", ""), |
| } |
| if iid in ocr_lookup: |
| text, tokens = ocr_lookup[iid] |
| doc["ocr_text"] = text |
| doc["ocr_tokens"] = tokens |
| vec = vec_lookup.get(iid) |
| if vec is not None: |
| doc["image_vec"] = vec if isinstance(vec, list) else vec.tolist() |
| yield doc |
|
|
|
|
| def bulk_index(client, docs, chunk_size=500): |
| ok_count, failed = 0, 0 |
| for ok, _ in helpers.streaming_bulk(client, docs, chunk_size=chunk_size, raise_on_error=False): |
| ok_count += ok |
| failed += (not ok) |
| print(f"Indexed {ok_count} docs, {failed} failures.") |
|
|
|
|
| def get_ocr_text(client, index_name, image_id) -> str: |
| """Fetch the cached OCR text for an image_id from the index.""" |
| resp = client.search( |
| index=index_name, |
| body={"size": 1, "query": {"term": {"image_id": image_id}}}, |
| ) |
| hits = resp["hits"]["hits"] |
| return hits[0]["_source"].get("ocr_text", "") if hits else "" |
|
|
|
|
| def retrieve_by_ocr_text(client, query, index_name, top_k=5, dataset=None): |
| """BM25 retrieval over the ocr_text field (for Recall@K by textual content).""" |
| must = [{"match": {"ocr_text": query}}] |
| body = {"size": top_k, "query": {"bool": {"must": must}}} |
| if dataset: |
| body["query"]["bool"]["filter"] = [{"term": {"dataset": dataset}}] |
| resp = client.search(index=index_name, body=body) |
| return [h["_source"] for h in resp["hits"]["hits"]] |
|
|
|
|
| |
| |
| |
|
|
| OCR_SYSTEM_PROMPT = ( |
| "You answer questions using ONLY text that an OCR engine extracted from an image. " |
| "You cannot see the image itself — work solely from the provided text. Give the " |
| "shortest exact answer (a word or short phrase), with no explanation. If the text " |
| "does not clearly contain the answer, give your best guess from it. Never reply that " |
| "you cannot see an image or that no image was provided." |
| ) |
| VISION_SYSTEM_PROMPT = ( |
| "You are a visual question answering assistant. Read any text visible in the image " |
| "and answer the question. Reply with the shortest exact answer (a word or short " |
| "phrase), with no explanation." |
| ) |
|
|
|
|
| def answer_with_ocr(llm_client, ocr_text, question, model=None) -> str: |
| """Answer using ONLY the pre-extracted OCR text as context (no image sent).""" |
| messages = [ |
| {"role": "system", "content": OCR_SYSTEM_PROMPT}, |
| {"role": "user", "content": f"OCR text from the image:\n\"{ocr_text}\"\n\nQuestion: {question}"}, |
| ] |
| resp = llm_client.chat.completions.create(model=model, messages=messages, max_tokens=64) |
| return resp.choices[0].message.content.strip() |
|
|
|
|
| def answer_with_vision(llm_client, image: Image.Image, question, model=None) -> str: |
| """Answer by sending the image directly to the LVLM (it reads the text itself).""" |
| from vlps.data import image_to_base64 |
| image_url = image_to_base64(image) |
| messages = [ |
| {"role": "system", "content": VISION_SYSTEM_PROMPT}, |
| {"role": "user", "content": [ |
| {"type": "image_url", "image_url": {"url": image_url}}, |
| {"type": "text", "text": question}, |
| ]}, |
| ] |
| resp = llm_client.chat.completions.create(model=model, messages=messages, max_tokens=64) |
| return resp.choices[0].message.content.strip() |
|
|
|
|
| |
| |
| |
|
|
| _PUNCT = re.compile(r"[^\w\s]") |
| _ARTICLES = {"a", "an", "the"} |
|
|
|
|
| def normalise_answer(ans: str) -> str: |
| ans = str(ans).lower().strip() |
| ans = _PUNCT.sub("", ans) |
| ans = " ".join(w for w in ans.split() if w not in _ARTICLES) |
| return ans |
|
|
|
|
| def vqa_accuracy(prediction: str, ground_truths: list) -> float: |
| """ |
| Standard VQA accuracy: min(#humans-that-said-it / 3, 1), averaged over the |
| 10-choose-9 subsets — here approximated by the common closed form. |
| """ |
| pred = normalise_answer(prediction) |
| gts = [normalise_answer(g) for g in ground_truths] |
| if not gts: |
| return 0.0 |
| matches = sum(1 for g in gts if g == pred) |
| return min(matches / 3.0, 1.0) |
|
|