87 lines
2.7 KiB
Python
87 lines
2.7 KiB
Python
"""
|
|
De-risking spike, take 3: instead of embedding + nearest-neighbor lookup,
|
|
just ask a self-hosted vision-language model directly what Pokemon is in
|
|
the photo. No reference images, no vector index -- the model's own
|
|
pretrained world knowledge does the recognition.
|
|
|
|
Usage: .venv_vlm\\Scripts\\python.exe match_spike_vlm.py
|
|
"""
|
|
|
|
from pathlib import Path
|
|
|
|
import torch
|
|
from PIL import Image
|
|
from transformers import AutoProcessor, Qwen2_5_VLForConditionalGeneration
|
|
|
|
HERE = Path(__file__).parent
|
|
QUERY_DIR = HERE / "images" / "query"
|
|
|
|
MODEL_NAME = "Qwen/Qwen2.5-VL-3B-Instruct"
|
|
|
|
PROMPT = (
|
|
"You are the recognition system inside a Pokedex app. Identify the "
|
|
"Pokemon species shown in this image, even if it's a toy, plush, "
|
|
"trading card, sprite, fan art, or in-game screenshot of it. "
|
|
"Reply with ONLY the species name, nothing else. If no Pokemon is "
|
|
"clearly depicted, reply with exactly: unknown"
|
|
)
|
|
|
|
|
|
def load_model():
|
|
processor = AutoProcessor.from_pretrained(MODEL_NAME)
|
|
model = Qwen2_5_VLForConditionalGeneration.from_pretrained(
|
|
MODEL_NAME, torch_dtype=torch.bfloat16, device_map="cuda"
|
|
)
|
|
model.eval()
|
|
return model, processor
|
|
|
|
|
|
def identify(model, processor, path: Path) -> str:
|
|
image = Image.open(path).convert("RGB")
|
|
messages = [
|
|
{
|
|
"role": "user",
|
|
"content": [
|
|
{"type": "image", "image": image},
|
|
{"type": "text", "text": PROMPT},
|
|
],
|
|
}
|
|
]
|
|
text = processor.apply_chat_template(
|
|
messages, tokenize=False, add_generation_prompt=True
|
|
)
|
|
inputs = processor(text=[text], images=[image], return_tensors="pt").to("cuda")
|
|
with torch.no_grad():
|
|
generated = model.generate(**inputs, max_new_tokens=16)
|
|
trimmed = generated[:, inputs["input_ids"].shape[1] :]
|
|
output = processor.batch_decode(
|
|
trimmed, skip_special_tokens=True, clean_up_tokenization_spaces=True
|
|
)[0]
|
|
return output.strip()
|
|
|
|
|
|
def main():
|
|
print(f"Loading {MODEL_NAME} on GPU...")
|
|
model, processor = load_model()
|
|
|
|
query_paths = (
|
|
sorted(QUERY_DIR.glob("*.jpg"))
|
|
+ sorted(QUERY_DIR.glob("*.webp"))
|
|
+ sorted(QUERY_DIR.glob("*.png"))
|
|
)
|
|
|
|
print(f"Identifying {len(query_paths)} query images...\n")
|
|
correct = 0
|
|
for qpath in query_paths:
|
|
expected = qpath.stem.split("_")[0]
|
|
answer = identify(model, processor, qpath)
|
|
is_correct = answer.strip().lower() == expected.lower()
|
|
correct += is_correct
|
|
marker = "OK " if is_correct else "MISS"
|
|
print(f"[{marker}] query={qpath.stem:<18} expected={expected:<10} model_said={answer!r}")
|
|
|
|
print(f"\n{correct}/{len(query_paths)} correct")
|
|
|
|
|
|
if __name__ == "__main__":
|
|
main()
|