""" 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()