Initial commit
Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
This commit is contained in:
commit
71db2d1ab9
126 changed files with 3198 additions and 0 deletions
96
server/server.py
Normal file
96
server/server.py
Normal file
|
|
@ -0,0 +1,96 @@
|
|||
"""
|
||||
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,
|
||||
}
|
||||
Loading…
Add table
Add a link
Reference in a new issue