diff --git a/backend/app/main.py b/backend/app/main.py index 70d9b5a..082ec1e 100644 --- a/backend/app/main.py +++ b/backend/app/main.py @@ -21,6 +21,9 @@ async def lifespan(app: FastAPI): caps = probe_upscale_capabilities() if not caps["realesrgan_pytorch"] and not caps["realesrgan_ncnn"]: asyncio.create_task(ensure_ncnn_installed()) + # Pre-download SAM model in background so first click is fast + from app.services.sam_service import ensure_sam_installed + asyncio.create_task(ensure_sam_installed()) yield diff --git a/backend/app/routers/ai_tools.py b/backend/app/routers/ai_tools.py index 243c5ba..b99be59 100644 --- a/backend/app/routers/ai_tools.py +++ b/backend/app/routers/ai_tools.py @@ -348,3 +348,63 @@ async def get_config(): }, } } + + +# ─── SAM (Segment Anything) ────────────────────────────────────────────────── + +class SegmentPointRequest(BaseModel): + image: str # base64 PNG/JPEG + points: list[list[int]] # [[x, y], ...] original image coords + labels: list[int] # 1=include, 0=exclude — same length as points + + +@router.post("/segment/point") +async def segment_point(req: SegmentPointRequest): + """ + Run SAM point-prompt segmentation. + Returns a binary mask PNG (white = selected area). + Auto-downloads the SAM ViT-B model (~375 MB) on first call. + """ + if not req.points: + raise HTTPException(status_code=400, detail="At least one point required.") + if len(req.points) != len(req.labels): + raise HTTPException(status_code=400, detail="points and labels must have the same length.") + + try: + image_bytes = base64.b64decode(req.image) + except Exception as e: + raise HTTPException(status_code=400, detail=f"Could not decode image: {e}") + + from app.services.sam_service import predict_points, get_install_status + try: + mask_bytes = await predict_points( + image_bytes, + [tuple(p) for p in req.points], + req.labels, + ) + return { + "mask": base64.b64encode(mask_bytes).decode(), + "sam_install": get_install_status(), + } + except RuntimeError as e: + raise HTTPException(status_code=503, detail=str(e)) + except Exception as e: + import traceback; traceback.print_exc() + raise HTTPException(status_code=500, detail=str(e)) + + +@router.get("/segment/install-status") +def segment_install_status(): + """Poll SAM model download progress.""" + from app.services.sam_service import get_install_status, sam_model_available + status = get_install_status() + status["model_ready"] = sam_model_available() + return status + + +@router.post("/segment/install") +async def segment_install(): + """Trigger SAM model download explicitly (also auto-triggered on first /segment/point call).""" + from app.services.sam_service import ensure_sam_installed, get_install_status + asyncio.create_task(ensure_sam_installed()) + return get_install_status() diff --git a/backend/app/services/sam_service.py b/backend/app/services/sam_service.py new file mode 100644 index 0000000..b10e7ea --- /dev/null +++ b/backend/app/services/sam_service.py @@ -0,0 +1,184 @@ +""" +SAM (Segment Anything Model) service. + +Auto-downloads the ViT-B checkpoint (~375 MB) on first use. +Caches the loaded model in memory; re-uses predictor across calls. + +Prediction API: + predict_points(image_bytes, points, labels) -> mask_bytes (PNG, white=selected) + points: list of (x, y) in original image pixels + labels: list of 1 (include) or 0 (exclude), same length as points +""" + +import asyncio +import io +import os +import urllib.request +from dataclasses import dataclass +from enum import Enum +from pathlib import Path +from typing import Optional + +import numpy as np +from PIL import Image + +# ── Model download ──────────────────────────────────────────────────────────── + +SAM_DIR = Path("/app/data/models/sam") +SAM_FILENAME = "sam_vit_b_01ec64.pth" +SAM_URL = f"https://dl.fbaipublicfiles.com/segment_anything/{SAM_FILENAME}" +SAM_PATH = SAM_DIR / SAM_FILENAME + + +class SamInstallState(str, Enum): + idle = "idle" + downloading = "downloading" + done = "done" + failed = "failed" + + +@dataclass +class SamInstallStatus: + state: SamInstallState = SamInstallState.idle + progress: int = 0 + message: str = "" + error: str = "" + + +_install_status = SamInstallStatus() +_install_lock = asyncio.Lock() + + +def get_install_status() -> dict: + s = _install_status + return {"state": s.state.value, "progress": s.progress, + "message": s.message, "error": s.error} + + +def sam_model_available() -> bool: + return SAM_PATH.exists() and SAM_PATH.stat().st_size > 100_000_000 + + +async def ensure_sam_installed() -> bool: + """Download SAM ViT-B checkpoint if not present. Returns True on success.""" + global _install_status + + if sam_model_available(): + _install_status = SamInstallStatus(state=SamInstallState.done, progress=100, + message="SAM model ready.") + return True + + async with _install_lock: + if sam_model_available(): + _install_status = SamInstallStatus(state=SamInstallState.done, progress=100, + message="SAM model ready.") + return True + + if _install_status.state == SamInstallState.downloading: + return False + + try: + SAM_DIR.mkdir(parents=True, exist_ok=True) + _install_status = SamInstallStatus( + state=SamInstallState.downloading, progress=0, + message="Downloading SAM ViT-B model (~375 MB)…", + ) + + def _download(): + def _progress(count, block, total): + if total > 0: + _install_status.progress = min(99, int(count * block * 99 / total)) + tmp = SAM_PATH.with_suffix(".tmp") + urllib.request.urlretrieve(SAM_URL, tmp, _progress) + tmp.rename(SAM_PATH) + + loop = asyncio.get_event_loop() + await loop.run_in_executor(None, _download) + + _install_status = SamInstallStatus(state=SamInstallState.done, progress=100, + message="SAM model ready.") + return True + + except Exception as exc: + _install_status = SamInstallStatus( + state=SamInstallState.failed, error=str(exc), + message="SAM download failed.", + ) + print(f"[sam] Download failed: {exc}") + return False + + +# ── Model cache ─────────────────────────────────────────────────────────────── + +_predictor = None +_predictor_lock = asyncio.Lock() + + +def _load_predictor(): + """Load SAM model and return a SamPredictor. Called in thread pool.""" + global _predictor + if _predictor is not None: + return _predictor + + import torch + from segment_anything import sam_model_registry, SamPredictor + + if torch.cuda.is_available(): + device = "cuda" + elif hasattr(torch.backends, "mps") and torch.backends.mps.is_available(): + device = "mps" + else: + device = "cpu" + + print(f"[sam] Loading SAM ViT-B on {device}…") + sam = sam_model_registry["vit_b"](checkpoint=str(SAM_PATH)) + sam.to(device) + _predictor = SamPredictor(sam) + print("[sam] Model loaded.") + return _predictor + + +# ── Prediction ──────────────────────────────────────────────────────────────── + +def _predict_sync(image_bytes: bytes, + points: list[tuple[int, int]], + labels: list[int]) -> bytes: + """ + Run SAM prediction synchronously (call via run_in_executor). + Returns PNG bytes: white = selected, black = background. + """ + predictor = _load_predictor() + + image = Image.open(io.BytesIO(image_bytes)).convert("RGB") + img_array = np.array(image) + + predictor.set_image(img_array) + + pt_array = np.array(points, dtype=np.float32) # [[x, y], ...] + lbl_array = np.array(labels, dtype=np.int32) # [1=fg, 0=bg, ...] + + masks, scores, _ = predictor.predict( + point_coords=pt_array, + point_labels=lbl_array, + multimask_output=True, + ) + + # Pick the highest-confidence mask + best = masks[int(np.argmax(scores))] # bool array H×W + + mask_img = Image.fromarray((best * 255).astype(np.uint8), mode="L") + buf = io.BytesIO() + mask_img.save(buf, format="PNG") + return buf.getvalue() + + +async def predict_points(image_bytes: bytes, + points: list[tuple[int, int]], + labels: list[int]) -> bytes: + """Async wrapper for SAM point prediction.""" + if not sam_model_available(): + ok = await ensure_sam_installed() + if not ok: + raise RuntimeError("SAM model not available.") + loop = asyncio.get_event_loop() + return await loop.run_in_executor(None, _predict_sync, image_bytes, points, labels) diff --git a/frontend/src/js/config.js b/frontend/src/js/config.js index 466112c..6678895 100644 --- a/frontend/src/js/config.js +++ b/frontend/src/js/config.js @@ -105,25 +105,8 @@ config.TOOLS = [ attributes: {}, }, { - name: 'ai_inpaint', - title: 'AI Inpaint', - on_activate: 'on_activate', - attributes: {}, - }, - { - name: 'ai_lama_erase', - title: 'AI Magic Erase (LaMa) - Paint over to erase', - attributes: { - size: { - value: 30, - min: 5, - max: 200, - }, - }, - }, - { - name: 'ai_smart_inpaint', - title: 'AI Smart Inpaint - Paint + describe replacement', + name: 'ai_edit', + title: 'AI Edit — paint mask, then Erase / Replace / Upscale / Expand', on_activate: 'on_activate', attributes: { size: { @@ -133,12 +116,6 @@ config.TOOLS = [ }, }, }, - { - name: 'ai_replace_selection', - title: 'AI Replace Selection - Use any selection tool first', - on_activate: 'on_activate', - attributes: {}, - }, { name: 'magic_wand', title: 'Magic Wand (Color Select)', diff --git a/frontend/src/js/tools/ai_edit.js b/frontend/src/js/tools/ai_edit.js new file mode 100644 index 0000000..14105da --- /dev/null +++ b/frontend/src/js/tools/ai_edit.js @@ -0,0 +1,578 @@ +/** + * AI Edit — unified smart selection + inpainting tool. + * + * Workflow: + * 1. CLICK mode (default): click any object → SAM auto-selects it (red overlay) + * • Alt+click → subtract from selection (deselect over-selected area) + * • Multiple clicks accumulate on the mask + * 2. BRUSH + / BRUSH − tabs: paint to add or erase from the mask by hand + * (refine what SAM missed or got wrong) + * 3. Action bar: Erase | Replace… | Upscale | Expand | Clear + * + * SAM model (~375 MB) auto-downloads on first click; progress shown inline. + * Falls back gracefully to brush-only if SAM is unavailable. + */ + +import app from './../app.js'; +import config from './../config.js'; +import Base_layers_class from './../core/base-layers.js'; +import alertify from './../../../node_modules/alertifyjs/build/alertify.min.js'; + +var instance = null; + +const BRUSH_DEFAULT = 30; +const OVERLAY_COLOR = 'rgba(255, 55, 55, 0.50)'; +const ERASE_COLOR = 'rgba(0, 0, 0, 0.70)'; // brush-erase preview + +class Tools_ai_edit_class { + + constructor() { + if (instance) return instance; + instance = this; + this.Base_layers = new Base_layers_class(); + this.name = 'ai_edit'; + this.title = 'AI Edit'; + // interaction state + this._mode = 'sam'; // 'sam' | 'brush_add' | 'brush_sub' + this._painting = false; + this._samWorking = false; + this._isRunning = false; + this._hasMask = false; + // DOM elements + this._maskCanvas = null; + this._maskCtx = null; + this._overlayEl = null; + this._panel = null; + } + + // ── Tool lifecycle ──────────────────────────────────────────────────────── + + on_activate() { + if (!config.layer || config.layer.type !== 'image') { + alertify.error('Select an image layer first.'); + return; + } + this._initMask(); + this._mountOverlay(); + this._mountPanel(); + } + + on_leave() { + this._removeOverlay(); + this._removePanel(); + this._painting = false; + } + + // ── Input routing ───────────────────────────────────────────────────────── + + mousedown(e) { + if (!config.layer || config.layer.type !== 'image') return; + if (this._mode === 'sam') { + this._handleSamClick(e); + } else { + this._painting = true; + this._brushPaint(e); + } + } + + mousemove(e) { + if (this._mode !== 'sam' && this._painting) this._brushPaint(e); + } + + mouseup() { + if (this._painting) { + this._painting = false; + if (this._hasMask) this._showActions(); + } + } + + // ── Coordinate mapping ──────────────────────────────────────────────────── + + _screenToImage(e) { + const canvasEl = document.getElementById('canvas_minipaint') || document.querySelector('canvas'); + if (!canvasEl) return null; + const rect = canvasEl.getBoundingClientRect(); + const scaleX = config.layer.width_original / (config.WIDTH * config.ZOOM); + const scaleY = config.layer.height_original / (config.HEIGHT * config.ZOOM); + const ix = ((e.clientX - rect.left) - config.layer.x * config.ZOOM) * scaleX; + const iy = ((e.clientY - rect.top) - config.layer.y * config.ZOOM) * scaleY; + return { ix, iy, scaleX, scaleY }; + } + + // ── SAM click selection ─────────────────────────────────────────────────── + + async _handleSamClick(e) { + if (this._samWorking) return; + const coords = this._screenToImage(e); + if (!coords) return; + + const label = e.altKey ? 0 : 1; // alt = exclude, normal = include + const x = Math.round(coords.ix); + const y = Math.round(coords.iy); + + // Clamp to image bounds + const w = config.layer.width_original; + const h = config.layer.height_original; + if (x < 0 || y < 0 || x >= w || y >= h) return; + + this._samWorking = true; + this._setSamCursor('wait'); + + // Collect any existing points for multi-click accumulation + if (!this._samPoints) this._samPoints = []; + if (!this._samLabels) this._samLabels = []; + this._samPoints.push([x, y]); + this._samLabels.push(label); + + try { + const imageB64 = this._getLayerB64(); + const base = window.API_BASE_URL || ''; + const r = await fetch(`${base}/api/segment/point`, { + method: 'POST', + headers: { 'Content-Type': 'application/json' }, + body: JSON.stringify({ + image: imageB64, + points: this._samPoints, + labels: this._samLabels, + }), + }); + + if (r.status === 503) { + // SAM model downloading — poll and retry + const data = await r.json().catch(() => ({})); + await this._waitForSamModel(data.detail || ''); + // Remove the point we just added so user can retry cleanly + this._samPoints.pop(); + this._samLabels.pop(); + this._samWorking = false; + this._setSamCursor('crosshair'); + return; + } + + if (!r.ok) { + const err = await r.json().catch(() => ({})); + throw new Error(err.detail || 'SAM failed'); + } + + const data = await r.json(); + await this._applySamMask(data.mask, label === 0); + this._hasMask = true; + this._showActions(); + + } catch (err) { + alertify.error('SAM failed: ' + (err.message || err)); + // Pop failed point + this._samPoints.pop(); + this._samLabels.pop(); + } + + this._samWorking = false; + this._setSamCursor('crosshair'); + } + + async _applySamMask(maskB64, isSubtract) { + return new Promise((resolve) => { + const img = new Image(); + img.onload = () => { + // Draw SAM mask onto our persistent mask canvas + const tmp = document.createElement('canvas'); + tmp.width = this._maskCanvas.width; + tmp.height = this._maskCanvas.height; + const tctx = tmp.getContext('2d'); + tctx.drawImage(img, 0, 0, tmp.width, tmp.height); + + if (isSubtract) { + // Erase mask where SAM says to subtract + this._maskCtx.globalCompositeOperation = 'destination-out'; + this._maskCtx.drawImage(tmp, 0, 0); + this._maskCtx.globalCompositeOperation = 'source-over'; + } else { + this._maskCtx.drawImage(tmp, 0, 0); + } + this._redrawOverlay(); + resolve(); + }; + img.src = 'data:image/png;base64,' + maskB64; + }); + } + + async _waitForSamModel(detail) { + // SAM model is downloading — show progress bar and poll + return new Promise((resolve) => { + alertify.message( + `