File size: 14,346 Bytes
5857fdc
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
4debea7
 
 
 
 
5857fdc
 
 
 
 
 
 
 
 
 
 
 
4debea7
 
 
 
 
5857fdc
 
 
 
 
 
 
 
 
 
 
 
 
 
 
af99d79
 
 
 
5857fdc
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
307dc62
 
 
 
 
5857fdc
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1581af1
 
 
5857fdc
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
307dc62
5857fdc
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
307dc62
5857fdc
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
307dc62
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
import os

os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True")

import spaces  # noqa: E402  (must precede torch / CUDA-touching imports)
import torch  # noqa: E402
import gradio as gr  # noqa: E402
from PIL import Image  # noqa: E402
from peft import PeftModel  # noqa: E402
from transformers import AutoModel, AutoProcessor, AutoTokenizer  # noqa: E402

# ---------------------------------------------------------------------------
# Model
# ---------------------------------------------------------------------------
BASE_MODEL = "Qwen/Qwen2.5-VL-7B-Instruct"
ADAPTER_ID = "hmhm1229/ConceptFormer-Qwen"

# Upstream evaluation encodes documents with this generic prompt and pools the
# EOS hidden state (`--pooling eos --append_eos_token --normalize`).
DOC_PROMPT = "What is shown in this image?"
QUERY_MAX_LEN = 256  # scripts/evaluate.sh: --query_max_len 256
DEFAULT_INSTRUCTION = (
    "Given a user query, retrieve a document image that answers the query."
)
MAX_CANDIDATES = 16
# Guard against enormous user uploads (the released benchmark pages are ~850x600,
# well below this cap, so example behaviour matches the paper's setup).
MAX_PIXELS = 1280 * 28 * 28

print(f"Loading processor from {BASE_MODEL}")
try:
    processor = AutoProcessor.from_pretrained(BASE_MODEL, max_pixels=MAX_PIXELS)
except Exception as exc:  # pragma: no cover - processor kwarg drift
    print(f"max_pixels kwarg rejected ({exc!r}); loading default processor")
    processor = AutoProcessor.from_pretrained(BASE_MODEL)
tokenizer = processor.tokenizer
if tokenizer.pad_token_id is None:
    tokenizer.pad_token_id = tokenizer.eos_token_id
tokenizer.padding_side = "right"

print(f"Loading base model {BASE_MODEL}")
model = AutoModel.from_pretrained(
    BASE_MODEL,
    dtype=torch.bfloat16,
    attn_implementation="sdpa",
    trust_remote_code=True,
)
if getattr(model.config, "pad_token_id", None) is None:
    try:
        model.config.pad_token_id = tokenizer.pad_token_id
    except Exception as exc:  # pragma: no cover - transformers v5 config drift
        print(f"Could not set config.pad_token_id: {exc!r}")

# ConceptFormer ships a `<|lcon|>` special token alongside the adapter. It is not
# used at retrieval time, but the reference loader resizes the base embedding
# table when the adapter tokenizer is larger (a no-op for Qwen2.5-VL).
try:
    adapter_tokenizer = AutoTokenizer.from_pretrained(ADAPTER_ID)
    adapter_vocab = len(adapter_tokenizer)
    cur_vocab = int(model.get_input_embeddings().weight.size(0))
    if adapter_vocab > cur_vocab:
        model.resize_token_embeddings(adapter_vocab)
        print(f"Resized embeddings {cur_vocab} -> {adapter_vocab}")
except Exception as exc:  # pragma: no cover
    print(f"Adapter tokenizer inspection failed: {exc!r}")

print(f"Merging ConceptFormer adapter {ADAPTER_ID}")
# `torch_device="cpu"`: PEFT otherwise infers "cuda" from the ZeroGPU-patched
# `torch.cuda.is_available()` and cannot materialise the adapter shards in the
# main process ("No CUDA GPUs are available").
model = PeftModel.from_pretrained(model, ADAPTER_ID, torch_device="cpu")
model = model.merge_and_unload()
model = model.eval().to("cuda")
print("Model ready.")


# ---------------------------------------------------------------------------
# Encoding helpers (mirror conceptformer.retriever.driver.encode)
# ---------------------------------------------------------------------------
def _pool_eos(hidden_states: torch.Tensor, attention_mask: torch.Tensor) -> torch.Tensor:
    lengths = attention_mask.sum(dim=1) - 1
    idx = torch.arange(hidden_states.size(0), device=hidden_states.device)
    reps = hidden_states[idx, lengths]
    return torch.nn.functional.normalize(reps.float(), p=2, dim=-1)


@torch.no_grad()
def _encode_query(text: str) -> torch.Tensor:
    enc = tokenizer(
        [text],
        padding=False,
        truncation=True,
        max_length=QUERY_MAX_LEN - 1,
        add_special_tokens=True,
        return_attention_mask=False,
        return_token_type_ids=False,
    )
    enc["input_ids"] = [ids + [tokenizer.eos_token_id] for ids in enc["input_ids"]]
    batch = tokenizer.pad(
        enc, padding=True, return_attention_mask=True, return_tensors="pt"
    )
    batch = {k: v.to("cuda") for k, v in batch.items()}
    out = model(**batch, return_dict=True, output_hidden_states=True, use_cache=False)
    return _pool_eos(out.hidden_states[-1], batch["attention_mask"])


@torch.no_grad()
def _encode_document(image: Image.Image) -> torch.Tensor:
    messages = [
        {
            "role": "user",
            "content": [
                {"type": "image", "image": image},
                {"type": "text", "text": DOC_PROMPT},
            ],
        }
    ]
    text = processor.apply_chat_template(
        messages, tokenize=False, add_generation_prompt=True
    )
    # `--append_eos_token`: the reference collator appends the EOS id after the
    # prompt. Appending it as text keeps every processor-produced tensor
    # (attention mask, token types, image grid) consistent in length.
    if tokenizer.eos_token:
        text = text + tokenizer.eos_token
    inputs = processor(text=[text], images=[image], return_tensors="pt")
    inputs = {k: v.to("cuda") for k, v in inputs.items()}
    out = model(**inputs, return_dict=True, output_hidden_states=True, use_cache=False)
    return _pool_eos(out.hidden_states[-1], inputs["attention_mask"])


# ---------------------------------------------------------------------------
# Sample corpus (Our World in Data charts, CC BY 4.0, via ConceptFormer-Eval)
# ---------------------------------------------------------------------------
ASSET_DIR = os.path.join(os.path.dirname(os.path.abspath(__file__)), "assets")
SAMPLE_PAGES = sorted(
    os.path.join(ASSET_DIR, f) for f in os.listdir(ASSET_DIR) if f.endswith(".png")
)


def _pretty_name(path: str) -> str:
    name = os.path.splitext(os.path.basename(path))[0]
    if name.startswith("owid_"):
        name = name.split("_", 2)[-1]
    return name.replace("-", " ")


DEFAULT_CANDIDATES = [(p, _pretty_name(p)) for p in SAMPLE_PAGES]


def _normalize_candidates(candidates) -> list:
    """Gallery values arrive as (path, caption) tuples, dicts, or bare paths."""
    paths = []
    for item in candidates or []:
        path = None
        if isinstance(item, (list, tuple)) and item:
            path = item[0]
        elif isinstance(item, dict):
            path = item.get("image") or item.get("path") or item.get("name")
            if isinstance(path, dict):
                path = path.get("path") or path.get("url")
        elif isinstance(item, str):
            path = item
        if isinstance(path, str) and path:
            paths.append(path)
    return paths


# ---------------------------------------------------------------------------
# Inference
# ---------------------------------------------------------------------------
def _estimate_duration(query="", candidates=None, instruction=DEFAULT_INSTRUCTION,
                       top_k=5, *args, **kwargs):
    n = len(_normalize_candidates(candidates if candidates is not None else DEFAULT_CANDIDATES))
    # Measured on ZeroGPU: ~0.3 s per page at MAX_PIXELS plus ~0.4 s for the query,
    # on top of a few seconds of weight streaming. Keep the request tight.
    return int(min(60, 12 + 2 * max(n, 1)))


@spaces.GPU(duration=_estimate_duration)
def retrieve(
    query: str,
    candidates: list = DEFAULT_CANDIDATES,
    instruction: str = DEFAULT_INSTRUCTION,
    top_k: int = 5,
) -> tuple:
    """Rank candidate document pages against a text query with ConceptFormer.

    Args:
        query: the natural-language search query.
        candidates: candidate document page images to rank.
        instruction: retrieval instruction prepended to the query.
        top_k: how many pages to return.

    Returns:
        A ranked gallery of pages, a table of cosine similarity scores, and a status line.
    """
    query = (query or "").strip()
    if not query:
        raise gr.Error("Please enter a query.")

    paths = _normalize_candidates(candidates)
    if not paths:
        raise gr.Error("Please provide at least one candidate document page.")
    if len(paths) > MAX_CANDIDATES:
        raise gr.Error(
            f"This demo ranks at most {MAX_CANDIDATES} pages per query "
            f"(got {len(paths)})."
        )

    instruction = (instruction or "").strip()
    query_text = f"Instruct: {instruction}\nQuery: {query}" if instruction else query

    import time

    start = time.perf_counter()
    q_rep = _encode_query(query_text)

    doc_reps = []
    for path in paths:
        with Image.open(path) as img:
            image = img.convert("RGB")
        doc_reps.append(_encode_document(image))
    doc_reps = torch.cat(doc_reps, dim=0)

    scores = (q_rep @ doc_reps.T)[0].cpu().tolist()
    elapsed = time.perf_counter() - start

    order = sorted(range(len(paths)), key=lambda i: scores[i], reverse=True)
    k = max(1, min(int(top_k), len(order)))

    gallery = [
        (paths[i], f"#{rank + 1} · {scores[i]:.4f} · {_pretty_name(paths[i])}")
        for rank, i in enumerate(order[:k])
    ]
    table = [
        [rank + 1, os.path.basename(paths[i]), round(float(scores[i]), 4)]
        for rank, i in enumerate(order)
    ]
    status = (
        f"Encoded 1 query and {len(paths)} page(s) in {elapsed:.1f}s · "
        f"top score {scores[order[0]]:.4f}"
    )
    return gallery, table, status


def reset_candidates() -> list:
    """Restore the bundled sample corpus in the candidate gallery."""
    return DEFAULT_CANDIDATES


EXAMPLES = [
    [
        "The chart here shows the coverage for Hepatitis B vaccination. Hepatitis B "
        "is a highly contagious viral infection that attacks the liver and is "
        "transmitted through contact with the blood or other body fluids of an "
        "infected person."
    ],
    [
        "One of the strongest determinants of how much meat people eat is how rich "
        "they are. In the scatterplot we see the relationship between per capita "
        "meat supply and average GDP per capita."
    ],
    [
        "Global trends on alcohol abstinence show a mirror image of drinking "
        "prevalence data. This is shown in the charts as the share of adults who "
        "had not drunk in the prior year and those who have never drunk alcohol."
    ],
    [
        "SDG Target 6.2 is to achieve access to adequate and equitable sanitation "
        "and hygiene for all and end open defecation by 2030."
    ],
]

CSS = """
#col-container { max-width: 1150px; margin: 0 auto; }
.dark .gradio-container { color: var(--body-text-color); }
"""

with gr.Blocks() as demo:
    with gr.Column(elem_id="col-container"):
        gr.Markdown(
            """
# 🔎 ConceptFormer — visual document retrieval

Rank document page images against a text query with
[ConceptFormer-Qwen](https://huggingface.co/hmhm1229/ConceptFormer-Qwen)
(LoRA on Qwen2.5-VL-7B-Instruct), from
[*ConceptFormer: Learning Adaptive Latent Concepts for Query-Document Alignment
in Visual Document Retrieval*](https://huggingface.co/papers/2608.15698)
([code](https://github.com/NEUIR/ConceptFormer)).

Queries and pages are embedded separately and scored by cosine similarity —
exactly the encode-then-search path used in the paper's evaluation.
The candidate gallery is preloaded with sample pages; drop in your own to search them.
            """
        )

        with gr.Row():
            query = gr.Textbox(
                label="Query",
                placeholder="Describe the page you are looking for…",
                lines=2,
                scale=4,
            )
            run = gr.Button("Search", variant="primary", scale=1)

        with gr.Row():
            with gr.Column(scale=1):
                candidates = gr.Gallery(
                    value=DEFAULT_CANDIDATES,
                    label="Candidate pages (upload your own)",
                    interactive=True,
                    type="filepath",
                    file_types=["image"],
                    sources=["upload", "clipboard"],
                    columns=3,
                    height=340,
                    show_label=True,
                )
                reset = gr.Button("Reset to sample pages", size="sm")
            with gr.Column(scale=1):
                results = gr.Gallery(
                    label="Ranked results",
                    interactive=False,
                    columns=2,
                    height=340,
                )

        status = gr.Markdown("")
        scores = gr.Dataframe(
            headers=["rank", "page", "score"],
            datatype=["number", "str", "number"],
            label="All candidates by cosine similarity",
            wrap=True,
        )

        with gr.Accordion("Advanced settings", open=False):
            instruction = gr.Textbox(
                label="Instruction prefix",
                value=DEFAULT_INSTRUCTION,
                info="Prepended as `Instruct: …\\nQuery: …`, following the benchmark queries.",
            )
            top_k = gr.Slider(
                label="Pages to show", minimum=1, maximum=MAX_CANDIDATES, step=1, value=5
            )

        gr.Examples(
            examples=EXAMPLES,
            inputs=[query],
            outputs=[results, scores, status],
            fn=retrieve,
            cache_examples=True,
            cache_mode="lazy",
            label="Example queries (Our World in Data charts)",
        )

        gr.Markdown(
            "Sample pages come from the `owid_charts_en` split of "
            "[ConceptFormer-Eval](https://huggingface.co/datasets/hmhm1229/ConceptFormer-Eval); "
            "charts by [Our World in Data](https://ourworldindata.org), CC BY 4.0."
        )

    run.click(
        retrieve,
        inputs=[query, candidates, instruction, top_k],
        outputs=[results, scores, status],
        api_name="retrieve",
    )
    query.submit(
        retrieve,
        inputs=[query, candidates, instruction, top_k],
        outputs=[results, scores, status],
        api_name=False,
    )
    reset.click(reset_candidates, outputs=candidates, api_name=False)

demo.launch(theme=gr.themes.Citrus(), css=CSS, mcp_server=True)