Piko-9b / tests /test_multimodal_inference.py
Dexy2's picture
Rewrite model card around verified evidence; correct misattributed benchmarks and config path leak
0810902 verified
Raw
History Blame Contribute Delete
3.61 kB
"""Image-input behaviour.
The vision tower in this checkpoint was copied verbatim from Qwen/Qwen3.5-9B and
never trained against Piko's fine-tuned language backbone. Empirically it works
anyway: OCR and document understanding both scored 10/10 on the custom suite, so
the assertions below are genuine regression guards.
What has NOT been tested is anything outside rendered documents — photographs,
handwriting, natural scenes, low-quality scans. Do not add assertions about those
without measuring first.
"""
from __future__ import annotations
from pathlib import Path
import pytest
pytestmark = pytest.mark.slow
def ask_about_image(loaded_model, image: Path, prompt: str, max_new_tokens: int = 96) -> str:
import torch
model, processor = loaded_model
messages = [
{
"role": "user",
"content": [
{"type": "image", "url": str(image)},
{"type": "text", "text": prompt},
],
}
]
inputs = processor.apply_chat_template(
messages,
add_generation_prompt=True,
tokenize=True,
return_dict=True,
return_tensors="pt",
).to(model.device)
with torch.inference_mode():
output = model.generate(**inputs, max_new_tokens=max_new_tokens, do_sample=False)
return processor.decode(
output[0][inputs["input_ids"].shape[1] :], skip_special_tokens=True
).strip()
def test_image_tokens_expand_the_prompt(loaded_model, receipt_image: Path) -> None:
"""Structural test: the processor must inject image tokens.
This passes regardless of whether the model reads the image correctly, and so
separates 'the multimodal plumbing works' from 'the model can see'.
"""
model, processor = loaded_model
text_only = processor.apply_chat_template(
[{"role": "user", "content": [{"type": "text", "text": "Describe it."}]}],
add_generation_prompt=True,
tokenize=True,
return_dict=True,
return_tensors="pt",
)
with_image = processor.apply_chat_template(
[
{
"role": "user",
"content": [
{"type": "image", "url": str(receipt_image)},
{"type": "text", "text": "Describe it."},
],
}
],
add_generation_prompt=True,
tokenize=True,
return_dict=True,
return_tensors="pt",
)
assert with_image["input_ids"].shape[1] > text_only["input_ids"].shape[1]
assert "pixel_values" in with_image
image_token_id = model.config.image_token_id
assert (with_image["input_ids"] == image_token_id).sum() > 0
def test_image_input_does_not_crash(loaded_model, receipt_image: Path) -> None:
text = ask_about_image(loaded_model, receipt_image, "Describe this image in one sentence.")
assert text
assert len(set(text.replace(" ", ""))) > 5, "degenerate output — check device placement"
def test_reads_total_from_receipt(loaded_model, receipt_image: Path) -> None:
"""OCR works on this checkpoint despite the tower never being re-aligned.
Measured at 10/10 on the custom suite, so this is a real regression guard
rather than an aspiration.
"""
assert "27.30" in ask_about_image(
loaded_model, receipt_image, "What is the TOTAL on this receipt? Number only."
)
def test_reads_merchant_from_receipt(loaded_model, receipt_image: Path) -> None:
answer = ask_about_image(loaded_model, receipt_image, "What is the merchant name? Name only.")
assert "northgate" in answer.lower()