Team Ai
Apppublic

jonilaserson/sam-encoder-api

sourceHugging Faceupdated 3mo agoView on Hugging Face
0likes
app.py310 linesDownload Raw Back to root
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