Spaces:
Running on Zero
Running on Zero
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)
|