| |
| """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 ( |
| QUICKDRAW_100_CLASSES, |
| SmallConditionalUNet, |
| make_schedule, |
| pick_device, |
| sample, |
| ) |
|
|
| from quickdraw_app import cnn_classifier |
|
|
|
|
| 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 |
|
|