96 lines
3.1 KiB
Python
96 lines
3.1 KiB
Python
"""
|
|
Pokedex image-recognition server.
|
|
|
|
Loads a vision-language model once at startup and keeps it resident on the
|
|
GPU. The Android app POSTs a photo to /identify and gets back a species
|
|
name -- no reference-image index needed, the model's own pretrained
|
|
knowledge does the recognition (see spikes/image_recognition for the
|
|
de-risking spike that validated this approach: 10/12 correct cold, vs 8/12
|
|
for embedding-based nearest-neighbor matching).
|
|
|
|
Run: .venv\\Scripts\\python.exe -m uvicorn server:app --host 0.0.0.0 --port 8420
|
|
"""
|
|
|
|
import io
|
|
import logging
|
|
|
|
import torch
|
|
from fastapi import FastAPI, File, HTTPException, UploadFile
|
|
from PIL import Image
|
|
from transformers import AutoProcessor, Qwen2_5_VLForConditionalGeneration
|
|
|
|
logging.basicConfig(level=logging.INFO)
|
|
logger = logging.getLogger("pokedex-server")
|
|
|
|
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"
|
|
)
|
|
|
|
app = FastAPI(title="Pokedex recognition server")
|
|
_model = None
|
|
_processor = None
|
|
|
|
|
|
@app.on_event("startup")
|
|
def load_model():
|
|
global _model, _processor
|
|
logger.info("Loading %s on GPU...", MODEL_NAME)
|
|
_processor = AutoProcessor.from_pretrained(MODEL_NAME)
|
|
_model = Qwen2_5_VLForConditionalGeneration.from_pretrained(
|
|
MODEL_NAME, torch_dtype=torch.bfloat16, device_map="cuda"
|
|
)
|
|
_model.eval()
|
|
logger.info("Model loaded, ready to serve.")
|
|
|
|
|
|
@app.get("/health")
|
|
def health():
|
|
return {"status": "ok", "model": MODEL_NAME, "ready": _model is not None}
|
|
|
|
|
|
@app.post("/identify")
|
|
async def identify(file: UploadFile = File(...)):
|
|
if _model is None or _processor is None:
|
|
raise HTTPException(503, "Model still loading, try again shortly")
|
|
|
|
raw = await file.read()
|
|
try:
|
|
image = Image.open(io.BytesIO(raw)).convert("RGB")
|
|
except Exception as exc:
|
|
raise HTTPException(400, f"Could not read image: {exc}")
|
|
|
|
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] :]
|
|
raw_answer = _processor.batch_decode(
|
|
trimmed, skip_special_tokens=True, clean_up_tokenization_spaces=False
|
|
)[0].strip()
|
|
|
|
species = raw_answer.strip().rstrip(".")
|
|
recognized = species.lower() != "unknown"
|
|
|
|
logger.info("identify: filename=%s -> %r", file.filename, raw_answer)
|
|
return {
|
|
"recognized": recognized,
|
|
"species": species if recognized else None,
|
|
"raw_response": raw_answer,
|
|
}
|