jonilaserson/sam-encoder-api
0
1"""2SAM Encoder + WallHead API3--------------------------4Microservice that runs the SAM ViT-B vision encoder plus the png2vtt5wall/door/window heads.6 7POST /probs8 multipart: image=<file> (PNG/JPEG/WebP, should already be ≤1024px)9 form: tta=0|1 (default 0 — single pass; 1 = 7-way TTA)10 Returns: application/octet-stream — numpy .npz with keys:11 wall, door, win: float16 (H, W) probability maps at the input12 image's resolution (after any ≤1024 downscale).13 Runs the full 3-way wall ensemble + door head server-side, so callers14 need no model weights and receive ~6MB instead of 16MB×N embeddings.15 16POST /embed (legacy — kept for older clients)17 multipart: image=<file>18 Returns: .npz with:19 neck: float16 (256, 64, 64) — last_hidden_state (SAM neck output)20 early: float16 (768, 64, 64) — hidden_states[6] block output (BCHW)21 22GET /health23 Returns: {"status": "ok", "device": "cpu"|"cuda"}24"""25 26import io27import os28import numpy as np29from pathlib import Path30 31from fastapi import FastAPI, File, Form, UploadFile, HTTPException32from fastapi.responses import Response, JSONResponse33from PIL import Image34import torch35from transformers import SamProcessor, SamModel36 37# ---------------------------------------------------------------------------38# Model loading (done once at startup)39# ---------------------------------------------------------------------------40 41DEVICE = "cuda" if torch.cuda.is_available() else "cpu"42MODEL_ID = "facebook/sam-vit-base"43MAX_SIDE = 102444EARLY_BLOCK = 6 # 0-indexed; block 6 = hidden_states[7] in transformers (includes embedding)45 46print(f"Loading {MODEL_ID} on {DEVICE}...")47_processor = SamProcessor.from_pretrained(MODEL_ID)48_sam = SamModel.from_pretrained(MODEL_ID).to(DEVICE).eval()49_encoder = _sam.vision_encoder50print("Model ready.")51 52# ---------------------------------------------------------------------------53# WallHead — decoder heads on top of the frozen SAM features54# (architecture mirrors inference.py:_build_wall_head in the main repo)55# ---------------------------------------------------------------------------56 57import torch.nn as nn58 59 60class WallHead(nn.Module):61 def __init__(self, in_ch=256, early_ch=768, mid_ch=128):62 super().__init__()63 self.early_proj = nn.Sequential(64 nn.Conv2d(early_ch, 64, 1, bias=False),65 nn.BatchNorm2d(64), nn.ReLU(inplace=True))66 self.project = nn.Sequential(67 nn.Conv2d(in_ch + 64, mid_ch, 1, bias=False),68 nn.BatchNorm2d(mid_ch), nn.ReLU(inplace=True))69 70 def _up(cin, cout):71 return nn.Sequential(72 nn.ConvTranspose2d(cin, cout, 2, stride=2, bias=False),73 nn.BatchNorm2d(cout), nn.ReLU(inplace=True),74 nn.Conv2d(cout, cout, 3, padding=1, bias=False),75 nn.BatchNorm2d(cout), nn.ReLU(inplace=True))76 self.up1 = _up(mid_ch, 64)77 self.up2 = _up(64, 32)78 self.up3 = _up(32, 16)79 self.up4 = _up(16, 16)80 self.head = nn.Conv2d(16, 3, 1)81 82 def forward(self, x, early_feat):83 ep = self.early_proj(early_feat)84 x = self.project(torch.cat([x, ep], 1))85 x = self.up1(x); x = self.up2(x); x = self.up3(x); x = self.up4(x)86 return self.head(x)87 88 89WALL_ENSEMBLE = [90 ("models/curated_head.pt", 0.41),91 ("models/68map_ep300.pt", 0.41),92 ("models/da_head.pt", 0.18),93]94DOOR_HEAD = "models/door_v2_ep20.pt"95 96_heads: dict = {}97for _path, _w in WALL_ENSEMBLE + [(DOOR_HEAD, None)]:98 _h = WallHead()99 _h.load_state_dict(torch.load(_path, map_location="cpu", weights_only=True))100 _heads[_path] = _h.to(DEVICE).eval()101print(f"{len(_heads)} WallHeads ready.")102 103 104@torch.no_grad()105def _run_heads(pixel_values: torch.Tensor) -> np.ndarray:106 """One SAM pass + all heads. Returns (3, 1024, 1024) float32 probs."""107 out = _encoder(pixel_values, output_hidden_states=True)108 neck = out.last_hidden_state # (1,256,64,64)109 early = out.hidden_states[EARLY_BLOCK + 1].permute(0, 3, 1, 2) # (1,768,64,64)110 111 wall = None112 wsum = 0.0113 for path, w in WALL_ENSEMBLE:114 probs = torch.sigmoid(_heads[path](neck, early))[0] # (3,1024,1024)115 wall = probs[0] * w if wall is None else wall + probs[0] * w116 wsum += w117 wall = wall / wsum118 door_probs = torch.sigmoid(_heads[DOOR_HEAD](neck, early))[0]119 return torch.stack([wall, door_probs[1], door_probs[2]]).cpu().numpy()120 121 122def _probs_single(img: Image.Image) -> np.ndarray:123 """Single-pass probs, cropped to the image's valid region. (3, h, w)."""124 w, h = img.size125 scale = MAX_SIDE / max(w, h)126 probs = _run_heads(_preprocess(img))127 return probs[:, :int(h * scale), :int(w * scale)]128 129 130def _probs_tta7(img: Image.Image) -> np.ndarray:131 """132 7-way TTA: base + hflip + vflip + hvflip + rot90/180/270 (rotations on a133 square-padded canvas). Mirrors inference.py:_all_probs_remote_tta7 exactly.134 """135 w, h = img.size136 sq = max(w, h)137 img_sq = Image.new("RGB", (sq, sq), 0)138 img_sq.paste(img, (0, 0))139 140 flip_variants = [141 (img, None),142 (img.transpose(Image.FLIP_LEFT_RIGHT), "hflip"),143 (img.transpose(Image.FLIP_TOP_BOTTOM), "vflip"),144 (img.transpose(Image.FLIP_LEFT_RIGHT).transpose(Image.FLIP_TOP_BOTTOM), "hvflip"),145 ]146 rot_variants = [147 (img_sq.rotate(90, expand=False), "rot90"),148 (img_sq.rotate(180, expand=False), "rot180"),149 (img_sq.rotate(270, expand=False), "rot270"),150 ]151 152 probs_sum = None153 for aug_img, aug in flip_variants:154 p = _probs_single(aug_img)155 if aug == "hflip":156 p = p[:, :, ::-1].copy()157 elif aug == "vflip":158 p = p[:, ::-1, :].copy()159 elif aug == "hvflip":160 p = p[:, ::-1, ::-1].copy()161 probs_sum = p if probs_sum is None else probs_sum + p162 163 # Rotations run on the sq×sq canvas; SamProcessor scales max-side to 1024164 # for both the flips (max side = sq) and the square, so after un-rotating,165 # the content occupies the same top-left region as the flip outputs.166 Hf, Wf = probs_sum.shape[1], probs_sum.shape[2]167 for aug_img, aug in rot_variants:168 p_sq = _probs_single(aug_img) # (3, S, S) in rotated space169 if aug == "rot90":170 p_sq = np.rot90(p_sq, k=-1, axes=(1, 2)).copy()171 elif aug == "rot180":172 p_sq = np.rot90(p_sq, k=2, axes=(1, 2)).copy()173 elif aug == "rot270":174 p_sq = np.rot90(p_sq, k=1, axes=(1, 2)).copy()175 probs_sum = probs_sum + p_sq[:, :Hf, :Wf]176 177 return probs_sum / (len(flip_variants) + len(rot_variants))178 179 180# ---------------------------------------------------------------------------181# App182# ---------------------------------------------------------------------------183 184app = FastAPI(title="SAM Encoder API", version="2.0")185 186 187def _preprocess(img: Image.Image) -> torch.Tensor:188 """Resize to SAM's expected 1024×1024 input and normalise."""189 # SamProcessor handles resizing and normalisation190 inputs = _processor(images=img, return_tensors="pt")191 return inputs["pixel_values"].to(DEVICE) # (1, 3, 1024, 1024)192 193 194def _embed(pixel_values: torch.Tensor) -> tuple[np.ndarray, np.ndarray]:195 """196 Run SAM vision encoder, return (neck, early) as float16 numpy arrays.197 198 neck: (256, 64, 64) — last_hidden_state (after SAM neck Conv layers)199 early: (768, 64, 64) — block-6 ViT output, permuted to CHW200 """201 with torch.no_grad():202 out = _encoder(pixel_values, output_hidden_states=True)203 204 # last_hidden_state: (1, 256, 64, 64) — already BCHW after SAM neck205 neck = out.last_hidden_state[0].cpu().to(torch.float16).numpy() # (256, 64, 64)206 207 # hidden_states: tuple of length n_blocks+1 (includes embedding layer)208 # hidden_states[EARLY_BLOCK+1] = output after block EARLY_BLOCK209 # shape: (1, 64, 64, 768) — BHWC in HuggingFace SAM implementation210 early_bhwc = out.hidden_states[EARLY_BLOCK + 1][0] # (64, 64, 768)211 early = early_bhwc.permute(2, 0, 1).cpu().to(torch.float16).numpy() # (768, 64, 64)212 213 return neck, early214 215 216@app.get("/health")217def health():218 return JSONResponse({"status": "ok", "device": DEVICE, "heads": len(_heads)})219 220 221def _load_upload(image: UploadFile, raw: bytes) -> Image.Image:222 allowed = {"image/png", "image/jpeg", "image/jpg", "image/webp"}223 ct = (image.content_type or "").lower()224 ext = Path(image.filename or "").suffix.lower()225 if ct not in allowed and ext not in {".png", ".jpg", ".jpeg", ".webp"}:226 raise HTTPException(400, "Unsupported image type. Use PNG, JPEG, or WebP.")227 try:228 img = Image.open(io.BytesIO(raw)).convert("RGB")229 except Exception as e:230 raise HTTPException(400, f"Could not read image: {e}")231 w, h = img.size232 if max(w, h) > MAX_SIDE:233 scale = MAX_SIDE / max(w, h)234 img = img.resize((round(w * scale), round(h * scale)), Image.LANCZOS)235 return img236 237 238@app.post("/probs")239async def probs(image: UploadFile = File(...), tta: int = Form(0)):240 """Full wall/door/window probability maps (ensemble + door head)."""241 raw = await image.read()242 img = _load_upload(image, raw)243 try:244 p = _probs_tta7(img) if tta else _probs_single(img)245 except Exception as e:246 raise HTTPException(500, f"Inference error: {e}")247 248 buf = io.BytesIO()249 np.savez_compressed(buf,250 wall=p[0].astype(np.float16),251 door=p[1].astype(np.float16),252 win=p[2].astype(np.float16))253 buf.seek(0)254 return Response(255 content=buf.read(),256 media_type="application/octet-stream",257 headers={"Content-Disposition": "attachment; filename=probs.npz"},258 )259 260 261@app.post("/embed")262async def embed(image: UploadFile = File(...)):263 # Validate content type264 allowed = {"image/png", "image/jpeg", "image/jpg", "image/webp"}265 ct = (image.content_type or "").lower()266 ext = Path(image.filename or "").suffix.lower()267 if ct not in allowed and ext not in {".png", ".jpg", ".jpeg", ".webp"}:268 raise HTTPException(400, "Unsupported image type. Use PNG, JPEG, or WebP.")269 270 # Load image271 try:272 raw = await image.read()273 img = Image.open(io.BytesIO(raw)).convert("RGB")274 except Exception as e:275 raise HTTPException(400, f"Could not read image: {e}")276 277 # Resize if needed (caller should already send ≤1024px, but be safe)278 w, h = img.size279 if max(w, h) > MAX_SIDE:280 scale = MAX_SIDE / max(w, h)281 img = img.resize((round(w * scale), round(h * scale)), Image.LANCZOS)282 283 # Run encoder284 try:285 pixel_values = _preprocess(img)286 neck, early = _embed(pixel_values)287 except Exception as e:288 raise HTTPException(500, f"Encoder error: {e}")289 290 # Serialise as compressed numpy291 buf = io.BytesIO()292 np.savez_compressed(buf, neck=neck, early=early)293 buf.seek(0)294 295 return Response(296 content=buf.read(),297 media_type="application/octet-stream",298 headers={"Content-Disposition": "attachment; filename=embeddings.npz"},299 )300 301 302# ---------------------------------------------------------------------------303# Entry point304# ---------------------------------------------------------------------------305 306if __name__ == "__main__":307 import uvicorn308 port = int(os.environ.get("PORT", 7860))309 uvicorn.run(app, host="0.0.0.0", port=port)310 