File size: 6,412 Bytes
9894238
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
#!/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