lcccluck's picture
Add QuickDraw diffusion model and app code
9894238 verified
Raw
History Blame Contribute Delete
6.41 kB
#!/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