#!/usr/bin/env python3 """Local web API for sampling the QuickDraw conditional DDPM.""" from __future__ import annotations import base64 import io import os import random import sys import threading from pathlib import Path import torch from fastapi import FastAPI, HTTPException from fastapi.responses import FileResponse from fastapi.staticfiles import StaticFiles from pydantic import BaseModel, Field from torchvision.utils import make_grid ROOT = Path(__file__).resolve().parents[1] if str(ROOT) not in sys.path: sys.path.insert(0, str(ROOT)) from train_quickdraw_ddpm import ( # noqa: E402 QUICKDRAW_100_CLASSES, SmallConditionalUNet, make_schedule, pick_device, sample, ) from quickdraw_app import cnn_classifier # noqa: E402 DEFAULT_CHECKPOINT = ROOT / "models" / "diffusion" / "checkpoint_step_500000.pt" CHECKPOINT_PATH = Path(os.environ.get("QUICKDRAW_CHECKPOINT", DEFAULT_CHECKPOINT)).expanduser() STATIC_DIR = Path(__file__).resolve().parent / "static" class GenerateRequest(BaseModel): class_name: str = Field(..., min_length=1) count: int = Field(4, ge=1, le=8) guidance_scale: float = Field(3.0, ge=0.0, le=8.0) seed: int | None = Field(None, ge=0, le=2**31 - 1) class RecognizeRequest(BaseModel): image: str = Field(..., min_length=1) top_k: int = Field(5, ge=1, le=10) predictor: str = Field("quickdraw100") class ModelState: def __init__(self) -> None: self.lock = threading.Lock() self.loaded = False self.device = pick_device() self.model: SmallConditionalUNet | None = None self.schedule = None self.classes: list[str] = QUICKDRAW_100_CLASSES self.image_size = 64 self.timesteps = 200 self.base_channels = 64 self.step: int | None = None def load(self) -> None: if self.loaded: return if not CHECKPOINT_PATH.exists(): raise FileNotFoundError(f"Checkpoint not found: {CHECKPOINT_PATH}") checkpoint = torch.load(CHECKPOINT_PATH, map_location=self.device, weights_only=False) self.classes = list(checkpoint.get("classes", QUICKDRAW_100_CLASSES)) self.image_size = int(checkpoint.get("image_size", 64)) self.timesteps = int(checkpoint.get("timesteps", 200)) self.base_channels = int(checkpoint.get("base_channels", 64)) self.step = checkpoint.get("step") model = SmallConditionalUNet(len(self.classes), base_channels=self.base_channels).to(self.device) state_dict = checkpoint.get("model_unwrapped") or checkpoint.get("model") if state_dict is None: raise RuntimeError("Checkpoint does not contain a model state dict.") model.load_state_dict(state_dict) model.eval() self.model = model self.schedule = make_schedule(self.timesteps, self.device) self.loaded = True def generate(self, request: GenerateRequest) -> str: with self.lock: self.load() assert self.model is not None assert self.schedule is not None try: class_index = self.classes.index(request.class_name) except ValueError as exc: raise HTTPException(status_code=400, detail="Unknown class name.") from exc seed = request.seed if request.seed is not None else random.randrange(0, 2**31) torch.manual_seed(seed) if self.device.type == "cuda": torch.cuda.manual_seed_all(seed) labels = torch.full((request.count,), class_index, dtype=torch.long, device=self.device) images = sample( self.model, labels, self.image_size, self.schedule, self.timesteps, self.device, guidance_scale=request.guidance_scale, ) grid = make_grid((images + 1) / 2, nrow=min(request.count, 4), padding=8, pad_value=0) image = grid.mul(255).byte().permute(1, 2, 0).cpu().numpy() from PIL import Image buffer = io.BytesIO() Image.fromarray(image).save(buffer, format="PNG") encoded = base64.b64encode(buffer.getvalue()).decode("ascii") return f"data:image/png;base64,{encoded}" state = ModelState() app = FastAPI(title="QuickDraw Diffusion") app.mount("/static", StaticFiles(directory=STATIC_DIR), name="static") @app.get("/") def index() -> FileResponse: return FileResponse(STATIC_DIR / "index.html") @app.get("/api/status") def status() -> dict: loaded = state.loaded return { "loaded": loaded, "checkpoint": str(CHECKPOINT_PATH), "checkpoint_exists": CHECKPOINT_PATH.exists(), "device": str(state.device), "step": state.step, "image_size": state.image_size, "timesteps": state.timesteps, "cnn_model": str(cnn_classifier.MODEL_PATH), "cnn_model_exists": cnn_classifier.MODEL_PATH.exists(), "quickdraw100_cnn_model": str(cnn_classifier.QUICKDRAW100_MODEL_PATH), "quickdraw100_cnn_model_exists": cnn_classifier.QUICKDRAW100_MODEL_PATH.exists(), } @app.get("/api/classes") def classes() -> dict: return {"classes": state.classes} @app.get("/api/cnn/classes") def cnn_classes(predictor: str = "quickdraw100") -> dict: try: return {"classes": cnn_classifier.classes(predictor=predictor), "predictor": predictor} except FileNotFoundError as exc: raise HTTPException(status_code=500, detail=str(exc)) from exc @app.post("/api/generate") def generate(request: GenerateRequest) -> dict: try: image = state.generate(request) except FileNotFoundError as exc: raise HTTPException(status_code=500, detail=str(exc)) from exc return { "image": image, "class_name": request.class_name, "count": request.count, "guidance_scale": request.guidance_scale, "step": state.step, "device": str(state.device), } @app.post("/api/recognize") def recognize(request: RecognizeRequest) -> dict: image = cnn_classifier.decode_data_url(request.image) try: return cnn_classifier.predict_image(image, request.top_k, predictor=request.predictor) except FileNotFoundError as exc: raise HTTPException(status_code=500, detail=str(exc)) from exc