From 0a506efe245b4a4e9af6cce7241221196e361da4 Mon Sep 17 00:00:00 2001 From: Siddharth Kumar Sah Date: Mon, 13 Apr 2026 00:48:05 +0800 Subject: [PATCH] feat(erase-object): overhaul object eraser with LaMa inpainting improvements Update erase-object pipeline, eraser canvas, and inpainting Python script. Add LaMa model download script and update Dockerfile for model support. Update multi-file tool routes for consistency. --- apps/api/src/routes/pipeline.ts | 4 + apps/api/src/routes/tools/collage.ts | 3 +- apps/api/src/routes/tools/compare.ts | 7 +- apps/api/src/routes/tools/compose.ts | 7 +- apps/api/src/routes/tools/erase-object.ts | 96 ++++++++- apps/api/src/routes/tools/favicon.ts | 5 +- apps/api/src/routes/tools/find-duplicates.ts | 5 +- apps/api/src/routes/tools/image-to-pdf.ts | 5 +- apps/api/src/routes/tools/split.ts | 5 +- apps/api/src/routes/tools/stitch.ts | 3 +- apps/api/src/routes/tools/vectorize.ts | 5 +- apps/api/src/routes/tools/watermark-image.ts | 7 +- .../tools/erase-object-settings.tsx | 53 ++++- .../src/components/tools/eraser-canvas.tsx | 74 +++++-- docker/Dockerfile | 2 +- docker/download_models.py | 29 +++ packages/ai/python/inpaint.py | 190 +++++++++++++----- 17 files changed, 405 insertions(+), 95 deletions(-) diff --git a/apps/api/src/routes/pipeline.ts b/apps/api/src/routes/pipeline.ts index ed5ad0ec..2e7da0e1 100644 --- a/apps/api/src/routes/pipeline.ts +++ b/apps/api/src/routes/pipeline.ts @@ -13,6 +13,7 @@ import { eq } from "drizzle-orm"; import type { FastifyInstance, FastifyReply, FastifyRequest } from "fastify"; import { z } from "zod"; import { db, schema } from "../db/index.js"; +import { autoOrient } from "../lib/auto-orient.js"; import { validateImageBuffer } from "../lib/file-validation.js"; import { sanitizeFilename } from "../lib/filename.js"; import { decodeHeic } from "../lib/heic-converter.js"; @@ -107,6 +108,9 @@ export async function registerPipelineRoutes(app: FastifyInstance): Promise; diff --git a/apps/api/src/routes/tools/compare.ts b/apps/api/src/routes/tools/compare.ts index a8518a6e..69cb3ce5 100644 --- a/apps/api/src/routes/tools/compare.ts +++ b/apps/api/src/routes/tools/compare.ts @@ -3,6 +3,7 @@ import { writeFile } from "node:fs/promises"; import { join } from "node:path"; import type { FastifyInstance } from "fastify"; import sharp from "sharp"; +import { autoOrient } from "../../lib/auto-orient.js"; import { ensureSharpCompat } from "../../lib/heic-converter.js"; import { createWorkspace } from "../../lib/workspace.js"; @@ -42,9 +43,9 @@ export function registerCompare(app: FastifyInstance) { } try { - // Decode HEIC/HEIF if needed - bufferA = await ensureSharpCompat(bufferA); - bufferB = await ensureSharpCompat(bufferB); + // Decode HEIC/HEIF if needed, then normalize EXIF orientation + bufferA = await autoOrient(await ensureSharpCompat(bufferA)); + bufferB = await autoOrient(await ensureSharpCompat(bufferB)); // Normalize both to same size for comparison const metaA = await sharp(bufferA).metadata(); diff --git a/apps/api/src/routes/tools/compose.ts b/apps/api/src/routes/tools/compose.ts index d4ef2e41..84851568 100644 --- a/apps/api/src/routes/tools/compose.ts +++ b/apps/api/src/routes/tools/compose.ts @@ -4,6 +4,7 @@ import { join } from "node:path"; import type { FastifyInstance } from "fastify"; import sharp from "sharp"; import { z } from "zod"; +import { autoOrient } from "../../lib/auto-orient.js"; import { sanitizeFilename } from "../../lib/filename.js"; import { ensureSharpCompat } from "../../lib/heic-converter.js"; import { createWorkspace } from "../../lib/workspace.js"; @@ -81,9 +82,9 @@ export function registerCompose(app: FastifyInstance) { } try { - // Decode HEIC/HEIF if needed - baseBuffer = await ensureSharpCompat(baseBuffer); - overlayBuffer = await ensureSharpCompat(overlayBuffer); + // Decode HEIC/HEIF if needed, then normalize EXIF orientation + baseBuffer = await autoOrient(await ensureSharpCompat(baseBuffer)); + overlayBuffer = await autoOrient(await ensureSharpCompat(overlayBuffer)); // Apply opacity to overlay if needed let processedOverlay = overlayBuffer; diff --git a/apps/api/src/routes/tools/erase-object.ts b/apps/api/src/routes/tools/erase-object.ts index 4674c6b6..2b3cd29c 100644 --- a/apps/api/src/routes/tools/erase-object.ts +++ b/apps/api/src/routes/tools/erase-object.ts @@ -3,13 +3,30 @@ import { writeFile } from "node:fs/promises"; import { basename, join } from "node:path"; import { inpaint } from "@stirling-image/ai"; import type { FastifyInstance, FastifyReply, FastifyRequest } from "fastify"; +import sharp from "sharp"; +import { autoOrient } from "../../lib/auto-orient.js"; import { validateImageBuffer } from "../../lib/file-validation.js"; +import { decodeHeic, encodeHeic } from "../../lib/heic-converter.js"; import { createWorkspace } from "../../lib/workspace.js"; import { updateSingleFileProgress } from "../progress.js"; +const EXT_MAP: Record = { + jpeg: "jpg", + jpg: "jpg", + png: "png", + webp: "webp", + tiff: "tiff", + gif: "gif", + avif: "avif", + heic: "heic", + heif: "heif", +}; + +const BROWSER_PREVIEWABLE = new Set(["png", "jpg", "jpeg", "webp", "gif", "avif", "bmp"]); + /** * Object eraser / inpainting route. - * Accepts an image and a mask image, erases masked areas. + * Accepts an image and a mask image, erases masked areas using LaMa. */ export function registerEraseObject(app: FastifyInstance) { app.post("/api/v1/tools/erase-object", async (request: FastifyRequest, reply: FastifyReply) => { @@ -17,6 +34,8 @@ export function registerEraseObject(app: FastifyInstance) { let maskBuffer: Buffer | null = null; let filename = "image"; let clientJobId: string | null = null; + let format = "png"; + let quality = 95; try { const parts = request.parts(); @@ -35,6 +54,10 @@ export function registerEraseObject(app: FastifyInstance) { } } else if (part.fieldname === "clientJobId") { clientJobId = part.value as string; + } else if (part.fieldname === "format") { + format = (part.value as string) || "png"; + } else if (part.fieldname === "quality") { + quality = Number(part.value) || 95; } } } catch (err) { @@ -64,9 +87,23 @@ export function registerEraseObject(app: FastifyInstance) { try { request.log.info( - { toolId: "erase-object", imageSize: imageBuffer.length, maskSize: maskBuffer.length }, + { + toolId: "erase-object", + imageSize: imageBuffer.length, + maskSize: maskBuffer.length, + format, + }, "Starting object erasure", ); + + // Decode HEIC/HEIF input via system decoder + if (imageValidation.format === "heif") { + imageBuffer = await decodeHeic(imageBuffer); + } + + // Auto-orient to fix EXIF rotation + imageBuffer = await autoOrient(imageBuffer); + const jobId = randomUUID(); const workspacePath = await createWorkspace(jobId); @@ -94,10 +131,58 @@ export function registerEraseObject(app: FastifyInstance) { onProgress, ); + // Convert to the requested output format using Sharp + const needsNodeConversion = ["heic", "heif", "avif"].includes(format); + let outputBuffer: Buffer; + let finalFormat = format; + + if (needsNodeConversion) { + if (format === "heic" || format === "heif") { + outputBuffer = await encodeHeic(resultBuffer, quality); + finalFormat = format; + } else { + outputBuffer = await sharp(resultBuffer).avif({ quality }).toBuffer(); + finalFormat = "avif"; + } + } else if (format === "jpg" || format === "jpeg") { + outputBuffer = await sharp(resultBuffer).jpeg({ quality }).toBuffer(); + finalFormat = "jpg"; + } else if (format === "webp") { + outputBuffer = await sharp(resultBuffer).webp({ quality }).toBuffer(); + finalFormat = "webp"; + } else if (format === "tiff") { + outputBuffer = await sharp(resultBuffer).tiff({ quality }).toBuffer(); + finalFormat = "tiff"; + } else if (format === "gif") { + outputBuffer = await sharp(resultBuffer).gif().toBuffer(); + finalFormat = "gif"; + } else { + outputBuffer = resultBuffer; + finalFormat = "png"; + } + // Save output - const outputFilename = `${filename.replace(/\.[^.]+$/, "")}_erased.png`; + const ext = EXT_MAP[finalFormat] || "png"; + const outputFilename = `${filename.replace(/\.[^.]+$/, "")}_erased.${ext}`; const outputPath = join(workspacePath, "output", outputFilename); - await writeFile(outputPath, resultBuffer); + await writeFile(outputPath, outputBuffer); + + // Generate browser-compatible preview for non-previewable formats + let previewUrl: string | undefined; + if (!BROWSER_PREVIEWABLE.has(finalFormat)) { + try { + const previewInput = + finalFormat === "heic" || finalFormat === "heif" + ? await decodeHeic(outputBuffer) + : outputBuffer; + const previewBuffer = await sharp(previewInput).webp({ quality: 80 }).toBuffer(); + const previewPath = join(workspacePath, "output", "preview.webp"); + await writeFile(previewPath, previewBuffer); + previewUrl = `/api/v1/download/${jobId}/preview.webp`; + } catch { + // Non-fatal - frontend will show fallback + } + } if (clientJobId) { updateSingleFileProgress({ @@ -110,8 +195,9 @@ export function registerEraseObject(app: FastifyInstance) { return reply.send({ jobId, downloadUrl: `/api/v1/download/${jobId}/${encodeURIComponent(outputFilename)}`, + previewUrl, originalSize: imageBuffer.length, - processedSize: resultBuffer.length, + processedSize: outputBuffer.length, }); } catch (err) { request.log.error({ err, toolId: "erase-object" }, "Object erasing failed"); diff --git a/apps/api/src/routes/tools/favicon.ts b/apps/api/src/routes/tools/favicon.ts index ffe872ae..89838ed4 100644 --- a/apps/api/src/routes/tools/favicon.ts +++ b/apps/api/src/routes/tools/favicon.ts @@ -3,6 +3,7 @@ import { basename, extname } from "node:path"; import archiver from "archiver"; import type { FastifyInstance } from "fastify"; import sharp from "sharp"; +import { autoOrient } from "../../lib/auto-orient.js"; import { ensureSharpCompat } from "../../lib/heic-converter.js"; const FAVICON_SIZES = [ @@ -62,8 +63,8 @@ export function registerFavicon(app: FastifyInstance) { archive.pipe(reply.raw); for (const file of uploadedFiles) { - // Decode HEIC/HEIF if needed - const decoded = await ensureSharpCompat(file.buffer); + // Decode HEIC/HEIF if needed, then normalize EXIF orientation + const decoded = await autoOrient(await ensureSharpCompat(file.buffer)); const stem = basename(file.filename, extname(file.filename)); // Single file: flat structure. Multiple files: per-image folders. const prefix = isSingleFile ? "" : `${stem}/`; diff --git a/apps/api/src/routes/tools/find-duplicates.ts b/apps/api/src/routes/tools/find-duplicates.ts index 8a0d446c..dab29e53 100644 --- a/apps/api/src/routes/tools/find-duplicates.ts +++ b/apps/api/src/routes/tools/find-duplicates.ts @@ -1,6 +1,7 @@ import { basename } from "node:path"; import type { FastifyInstance } from "fastify"; import sharp from "sharp"; +import { autoOrient } from "../../lib/auto-orient.js"; import { ensureSharpCompat } from "../../lib/heic-converter.js"; /** @@ -67,9 +68,9 @@ export function registerFindDuplicates(app: FastifyInstance) { } try { - // Decode HEIC/HEIF if needed + // Decode HEIC/HEIF if needed, then normalize EXIF orientation for (const file of files) { - file.buffer = await ensureSharpCompat(file.buffer); + file.buffer = await autoOrient(await ensureSharpCompat(file.buffer)); } // Compute hashes for all images diff --git a/apps/api/src/routes/tools/image-to-pdf.ts b/apps/api/src/routes/tools/image-to-pdf.ts index 08a38f59..b3bca254 100644 --- a/apps/api/src/routes/tools/image-to-pdf.ts +++ b/apps/api/src/routes/tools/image-to-pdf.ts @@ -5,6 +5,7 @@ import type { FastifyInstance } from "fastify"; import PDFDocument from "pdfkit"; import sharp from "sharp"; import { z } from "zod"; +import { autoOrient } from "../../lib/auto-orient.js"; import { ensureSharpCompat } from "../../lib/heic-converter.js"; import { createWorkspace } from "../../lib/workspace.js"; @@ -96,8 +97,8 @@ export function registerImageToPdf(app: FastifyInstance) { for (const file of files) { doc.addPage({ size: [pageW, pageH], margin }); - // Decode HEIC/HEIF if needed, then convert to PNG for PDFKit compatibility - const compatBuffer = await ensureSharpCompat(file.buffer); + // Decode HEIC/HEIF if needed, normalize EXIF orientation, then convert to PNG for PDFKit + const compatBuffer = await autoOrient(await ensureSharpCompat(file.buffer)); const pngBuffer = await sharp(compatBuffer).png().toBuffer(); const meta = await sharp(pngBuffer).metadata(); diff --git a/apps/api/src/routes/tools/split.ts b/apps/api/src/routes/tools/split.ts index aa608d7a..9cb09959 100644 --- a/apps/api/src/routes/tools/split.ts +++ b/apps/api/src/routes/tools/split.ts @@ -4,6 +4,7 @@ import archiver from "archiver"; import type { FastifyInstance } from "fastify"; import sharp from "sharp"; import { z } from "zod"; +import { autoOrient } from "../../lib/auto-orient.js"; import { ensureSharpCompat } from "../../lib/heic-converter.js"; const settingsSchema = z.object({ @@ -58,8 +59,8 @@ export function registerSplit(app: FastifyInstance) { } try { - // Decode HEIC/HEIF if needed - fileBuffer = await ensureSharpCompat(fileBuffer); + // Decode HEIC/HEIF if needed, then normalize EXIF orientation + fileBuffer = await autoOrient(await ensureSharpCompat(fileBuffer)); const metadata = await sharp(fileBuffer).metadata(); const fullW = metadata.width ?? 0; diff --git a/apps/api/src/routes/tools/stitch.ts b/apps/api/src/routes/tools/stitch.ts index 64bd2aee..0cd0f555 100644 --- a/apps/api/src/routes/tools/stitch.ts +++ b/apps/api/src/routes/tools/stitch.ts @@ -4,6 +4,7 @@ import { basename, join } from "node:path"; import type { FastifyInstance } from "fastify"; import sharp from "sharp"; import { z } from "zod"; +import { autoOrient } from "../../lib/auto-orient.js"; import { validateImageBuffer } from "../../lib/file-validation.js"; import { ensureSharpCompat } from "../../lib/heic-converter.js"; import { createWorkspace } from "../../lib/workspace.js"; @@ -72,7 +73,7 @@ export function registerStitch(app: FastifyInstance) { .status(400) .send({ error: `Invalid file "${file.filename}": ${validation.reason}` }); } - file.buffer = await ensureSharpCompat(file.buffer); + file.buffer = await autoOrient(await ensureSharpCompat(file.buffer)); } let settings: z.infer; diff --git a/apps/api/src/routes/tools/vectorize.ts b/apps/api/src/routes/tools/vectorize.ts index 0e81239a..882500b0 100644 --- a/apps/api/src/routes/tools/vectorize.ts +++ b/apps/api/src/routes/tools/vectorize.ts @@ -5,6 +5,7 @@ import type { FastifyInstance } from "fastify"; import potrace from "potrace"; import sharp from "sharp"; import { z } from "zod"; +import { autoOrient } from "../../lib/auto-orient.js"; import { ensureSharpCompat } from "../../lib/heic-converter.js"; import { createWorkspace } from "../../lib/workspace.js"; @@ -79,8 +80,8 @@ export function registerVectorize(app: FastifyInstance) { } try { - // Decode HEIC/HEIF if needed - fileBuffer = await ensureSharpCompat(fileBuffer); + // Decode HEIC/HEIF if needed, then normalize EXIF orientation + fileBuffer = await autoOrient(await ensureSharpCompat(fileBuffer)); // Convert to BMP-compatible format for potrace (PNG) const pngBuffer = await sharp(fileBuffer).grayscale().png().toBuffer(); diff --git a/apps/api/src/routes/tools/watermark-image.ts b/apps/api/src/routes/tools/watermark-image.ts index 1641a635..c6457823 100644 --- a/apps/api/src/routes/tools/watermark-image.ts +++ b/apps/api/src/routes/tools/watermark-image.ts @@ -1,6 +1,7 @@ import type { FastifyInstance } from "fastify"; import sharp from "sharp"; import { z } from "zod"; +import { autoOrient } from "../../lib/auto-orient.js"; import { ensureSharpCompat } from "../../lib/heic-converter.js"; const settingsSchema = z.object({ @@ -68,9 +69,9 @@ export function registerWatermarkImage(app: FastifyInstance) { } try { - // Decode HEIC/HEIF if needed - mainBuffer = await ensureSharpCompat(mainBuffer); - watermarkBuffer = await ensureSharpCompat(watermarkBuffer); + // Decode HEIC/HEIF if needed, then normalize EXIF orientation + mainBuffer = await autoOrient(await ensureSharpCompat(mainBuffer)); + watermarkBuffer = await autoOrient(await ensureSharpCompat(watermarkBuffer)); const mainImage = sharp(mainBuffer); const mainMeta = await mainImage.metadata(); diff --git a/apps/web/src/components/tools/erase-object-settings.tsx b/apps/web/src/components/tools/erase-object-settings.tsx index 0f9381d1..117a8c2f 100644 --- a/apps/web/src/components/tools/erase-object-settings.tsx +++ b/apps/web/src/components/tools/erase-object-settings.tsx @@ -6,6 +6,9 @@ import { generateId } from "@/lib/utils"; import { useFileStore } from "@/stores/file-store"; import type { EraserCanvasRef } from "./eraser-canvas"; +const OUTPUT_FORMATS = ["png", "jpg", "webp", "avif", "tiff", "gif", "heic", "heif"] as const; +const LOSSY_FORMATS = ["jpg", "jpeg", "webp", "avif", "heic", "heif"]; + interface EraseObjectSettingsProps { eraserRef: React.RefObject; hasStrokes: boolean; @@ -29,6 +32,9 @@ export function EraseObjectSettings({ const [elapsed, setElapsed] = useState(0); const elapsedRef = useRef | null>(null); + const [outputFormat, setOutputFormat] = useState("png"); + const [quality, setQuality] = useState(95); + const handleProcess = async () => { if (files.length === 0 || !eraserRef.current) return; @@ -67,6 +73,8 @@ export function EraseObjectSettings({ formData.append("file", files[0]); formData.append("mask", maskFile); formData.append("clientJobId", clientJobId); + formData.append("format", outputFormat); + formData.append("quality", String(quality)); const xhr = new XMLHttpRequest(); xhr.upload.onprogress = (e) => { @@ -87,7 +95,7 @@ export function EraseObjectSettings({ setDownloadUrl(data.downloadUrl); setOriginalSize(data.originalSize); setProcessedSize(data.processedSize); - setProcessedUrl(data.downloadUrl); + setProcessedUrl(data.downloadUrl, data.previewUrl); setSizes(data.originalSize, data.processedSize); } catch { setError("Invalid response"); @@ -166,10 +174,51 @@ export function EraseObjectSettings({ )} + {/* Output Format */} +
+ + +
+ + {/* Quality (lossy formats only) */} + {LOSSY_FORMATS.includes(outputFormat) && ( +
+
+ + {quality} +
+ setQuality(Number(e.target.value))} + className="w-full mt-1" + /> +
+ )} + {/* Hint */} {hasFile && !hasStrokes && (

- Paint over the objects you want to remove on the image. + Paint over the objects you want to remove. Use Ctrl+Z to undo.

)} diff --git a/apps/web/src/components/tools/eraser-canvas.tsx b/apps/web/src/components/tools/eraser-canvas.tsx index f008c106..8eae014c 100644 --- a/apps/web/src/components/tools/eraser-canvas.tsx +++ b/apps/web/src/components/tools/eraser-canvas.tsx @@ -30,6 +30,9 @@ export const EraserCanvas = forwardRef(funct const drawingRef = useRef(false); const currentPointsRef = useRef([]); + // Cursor position for brush preview + const [cursorPos, setCursorPos] = useState(null); + // Measure and fit image to container const measure = useCallback(() => { const img = imgRef.current; @@ -48,6 +51,7 @@ export const EraserCanvas = forwardRef(funct }, []); // Reset strokes when image changes + // biome-ignore lint/correctness/useExhaustiveDependencies: imageSrc triggers intentional reset useEffect(() => { strokesRef.current = []; currentPointsRef.current = []; @@ -55,19 +59,7 @@ export const EraserCanvas = forwardRef(funct setCanvasSize(null); }, [imageSrc, onStrokeChange]); - // Redraw all strokes - const redraw = useCallback(() => { - const ctx = canvasRef.current?.getContext("2d"); - if (!ctx || !canvasSize) return; - - ctx.clearRect(0, 0, canvasSize.w, canvasSize.h); - - for (const stroke of strokesRef.current) { - drawStroke(ctx, stroke); - } - }, [canvasSize]); - - function drawStroke(ctx: CanvasRenderingContext2D, stroke: Stroke) { + const drawStroke = useCallback((ctx: CanvasRenderingContext2D, stroke: Stroke) => { ctx.lineCap = "round"; ctx.lineJoin = "round"; @@ -86,7 +78,33 @@ export const EraserCanvas = forwardRef(funct } ctx.stroke(); } - } + }, []); + + // Redraw all strokes + const redraw = useCallback(() => { + const ctx = canvasRef.current?.getContext("2d"); + if (!ctx || !canvasSize) return; + + ctx.clearRect(0, 0, canvasSize.w, canvasSize.h); + + for (const stroke of strokesRef.current) { + drawStroke(ctx, stroke); + } + }, [canvasSize, drawStroke]); + + // Keyboard shortcut: Ctrl+Z for undo + useEffect(() => { + const handler = (e: KeyboardEvent) => { + if ((e.metaKey || e.ctrlKey) && e.key === "z") { + e.preventDefault(); + strokesRef.current.pop(); + onStrokeChange(strokesRef.current.length > 0); + redraw(); + } + }; + window.addEventListener("keydown", handler); + return () => window.removeEventListener("keydown", handler); + }, [onStrokeChange, redraw]); // Get canvas-relative point from event const getPoint = useCallback((e: React.MouseEvent | React.TouchEvent): Point | null => { @@ -121,9 +139,12 @@ export const EraserCanvas = forwardRef(funct const handleMove = useCallback( (e: React.MouseEvent | React.TouchEvent) => { + // Update cursor position for brush preview + const pt = getPoint(e); + if (pt) setCursorPos(pt); + if (!drawingRef.current) return; if ("touches" in e) e.preventDefault(); - const pt = getPoint(e); if (!pt) return; currentPointsRef.current.push(pt); @@ -159,6 +180,11 @@ export const EraserCanvas = forwardRef(funct } }, [brushSize, onStrokeChange, redraw]); + const handleLeave = useCallback(() => { + setCursorPos(null); + handleUp(); + }, [handleUp]); + // Expose methods useImperativeHandle( ref, @@ -252,15 +278,29 @@ export const EraserCanvas = forwardRef(funct ref={canvasRef} width={canvasSize.w} height={canvasSize.h} - className="absolute inset-0 cursor-crosshair touch-none" + className="absolute inset-0 touch-none" + style={{ cursor: "none" }} onMouseDown={handleDown} onMouseMove={handleMove} onMouseUp={handleUp} - onMouseLeave={handleUp} + onMouseLeave={handleLeave} onTouchStart={handleDown} onTouchMove={handleMove} onTouchEnd={handleUp} /> + {/* Brush cursor preview */} + {cursorPos && ( +
+ )}
)} diff --git a/docker/Dockerfile b/docker/Dockerfile index 9fdbb4b7..d83d946d 100644 --- a/docker/Dockerfile +++ b/docker/Dockerfile @@ -215,7 +215,7 @@ ENV PYTHONWARNINGS=ignore \ # Create non-root user for runtime RUN groupadd -r stirling && useradd -r -g stirling -d /app -s /sbin/nologin stirling -RUN chown -R stirling:stirling /app /data /tmp/workspace /opt/venv +RUN chown -R stirling:stirling /app /data /tmp/workspace /opt/venv /opt/models # Entrypoint fixes volume permissions then drops to stirling via gosu COPY docker/entrypoint.sh /usr/local/bin/entrypoint.sh diff --git a/docker/download_models.py b/docker/download_models.py index c4c7860a..271292b7 100644 --- a/docker/download_models.py +++ b/docker/download_models.py @@ -13,6 +13,11 @@ os.environ["PADDLE_DEVICE"] = "cpu" os.environ["FLAGS_use_cuda"] = "0" os.environ["CUDA_VISIBLE_DEVICES"] = "" +LAMA_MODEL_DIR = "/opt/models/lama" +LAMA_MODEL_URL = "https://huggingface.co/Carve/LaMa-ONNX/resolve/main/lama_fp32.onnx" +LAMA_MODEL_PATH = os.path.join(LAMA_MODEL_DIR, "lama_fp32.onnx") +LAMA_MIN_SIZE = 100_000_000 # ~200 MB + REALESRGAN_MODEL_DIR = "/opt/models/realesrgan" REALESRGAN_MODEL_URL = ( "https://github.com/xinntao/Real-ESRGAN/releases/download/v0.1.0/RealESRGAN_x4plus.pth" @@ -98,6 +103,20 @@ def download_rembg_models(): print(f"All {len(REMBG_MODELS)} rembg models downloaded.\n") +def download_lama_model(): + """Download LaMa ONNX inpainting model from HuggingFace.""" + print("=== Downloading LaMa ONNX model ===") + os.makedirs(LAMA_MODEL_DIR, exist_ok=True) + print(f" Downloading from {LAMA_MODEL_URL}...") + urllib.request.urlretrieve(LAMA_MODEL_URL, LAMA_MODEL_PATH) + + size = os.path.getsize(LAMA_MODEL_PATH) + assert size > LAMA_MIN_SIZE, ( + f"LaMa model too small: {size} bytes (expected > {LAMA_MIN_SIZE})" + ) + print(f" lama_fp32.onnx downloaded ({size / 1_000_000:.1f} MB)\n") + + def download_realesrgan_model(): """Download RealESRGAN_x4plus.pth pretrained weights.""" print("=== Downloading RealESRGAN model ===") @@ -196,6 +215,15 @@ def smoke_test(): import mediapipe as mp print(" MediaPipe import OK") + # LaMa model file must exist + assert os.path.exists(LAMA_MODEL_PATH), ( + f"LaMa model missing: {LAMA_MODEL_PATH}" + ) + assert os.path.getsize(LAMA_MODEL_PATH) > LAMA_MIN_SIZE, ( + "LaMa model file is too small" + ) + print(" LaMa ONNX model file verified") + # RealESRGAN model file must exist assert os.path.exists(REALESRGAN_MODEL_PATH), ( f"RealESRGAN model missing: {REALESRGAN_MODEL_PATH}" @@ -232,6 +260,7 @@ def smoke_test(): def main(): print("Pre-downloading all ML models...\n") + download_lama_model() download_rembg_models() download_realesrgan_model() download_gfpgan_model() diff --git a/packages/ai/python/inpaint.py b/packages/ai/python/inpaint.py index 38af9465..5a50d76e 100644 --- a/packages/ai/python/inpaint.py +++ b/packages/ai/python/inpaint.py @@ -1,5 +1,6 @@ -"""Object erasing / inpainting using OpenCV.""" +"""Object erasing / inpainting using LaMa (Large Mask Inpainting) via ONNX.""" import sys +import os import json @@ -8,68 +9,159 @@ def emit_progress(percent, stage): print(json.dumps({"progress": percent, "stage": stage}), file=sys.stderr, flush=True) +# Resolve the LaMa ONNX model path. +# Docker places it at /opt/models/lama/lama_fp32.onnx. +# For local dev, check a user-writable cache dir. +LAMA_MODEL_DIR = os.environ.get("LAMA_MODEL_DIR", "/opt/models/lama") +LAMA_MODEL_PATH = os.path.join(LAMA_MODEL_DIR, "lama_fp32.onnx") +LAMA_LOCAL_CACHE = os.path.join(os.path.expanduser("~"), ".cache", "stirling-image", "lama") +LAMA_LOCAL_PATH = os.path.join(LAMA_LOCAL_CACHE, "lama_fp32.onnx") +LAMA_HF_URL = "https://huggingface.co/Carve/LaMa-ONNX/resolve/main/lama_fp32.onnx" + +# The ONNX model expects 512x512 fixed input. +MODEL_SIZE = 512 + + +def _get_model_path(): + """Return path to the LaMa ONNX model, downloading if needed.""" + if os.path.exists(LAMA_MODEL_PATH): + return LAMA_MODEL_PATH + if os.path.exists(LAMA_LOCAL_PATH): + return LAMA_LOCAL_PATH + + # Auto-download for local dev + emit_progress(5, "Downloading LaMa model") + os.makedirs(LAMA_LOCAL_CACHE, exist_ok=True) + import urllib.request + urllib.request.urlretrieve(LAMA_HF_URL, LAMA_LOCAL_PATH) + return LAMA_LOCAL_PATH + + +def _preprocess_image(img_array): + """Convert HWC uint8 RGB image to NCHW float32 [0,1] at MODEL_SIZE.""" + import cv2 + import numpy as np + + resized = cv2.resize(img_array, (MODEL_SIZE, MODEL_SIZE), interpolation=cv2.INTER_AREA) + # HWC -> CHW, normalize to [0, 1], add batch dim + chw = np.transpose(resized, (2, 0, 1)).astype(np.float32) / 255.0 + return chw[np.newaxis, ...] # (1, 3, 512, 512) + + +def _preprocess_mask(mask_array): + """Convert HW uint8 grayscale mask to NC(1)HW float32 binary at MODEL_SIZE.""" + import cv2 + import numpy as np + + resized = cv2.resize(mask_array, (MODEL_SIZE, MODEL_SIZE), interpolation=cv2.INTER_NEAREST) + # Threshold to binary 0/1 + binary = (resized > 127).astype(np.float32) + return binary[np.newaxis, np.newaxis, ...] # (1, 1, 512, 512) + + +def _feathered_composite(original, inpainted, mask, feather_radius=5): + """Composite inpainted region into original using a feathered mask. + + This preserves full quality in non-masked areas and smoothly blends + the inpainted region at the boundary. + """ + import cv2 + import numpy as np + + # Dilate mask slightly for smoother transition + kernel = cv2.getStructuringElement(cv2.MORPH_ELLIPSE, (feather_radius, feather_radius)) + dilated = cv2.dilate(mask.astype(np.uint8), kernel, iterations=1) + + # Gaussian blur the dilated mask for feathering + blur_size = feather_radius * 2 + 1 + alpha = cv2.GaussianBlur(dilated.astype(np.float32), (blur_size, blur_size), 0) + alpha = np.clip(alpha, 0.0, 1.0) + + # Expand alpha to 3 channels + alpha_3ch = alpha[:, :, np.newaxis] + + # Composite: original * (1 - alpha) + inpainted * alpha + result = (original.astype(np.float32) * (1.0 - alpha_3ch) + + inpainted.astype(np.float32) * alpha_3ch) + return np.clip(result, 0, 255).astype(np.uint8) + + def main(): input_path = sys.argv[1] mask_path = sys.argv[2] output_path = sys.argv[3] try: - emit_progress(10, "Preparing") + emit_progress(5, "Preparing") from PIL import Image + import numpy as np try: import cv2 - import numpy as np - - emit_progress(20, "Ready") - - img = Image.open(input_path).convert("RGB") - mask = Image.open(mask_path).convert("L") - - # Resize mask to match image if needed - emit_progress(30, "Analyzing mask") - if mask.size != img.size: - mask = mask.resize(img.size, Image.NEAREST) - - img_array = cv2.cvtColor(np.array(img), cv2.COLOR_RGB2BGR) - mask_array = np.array(mask) - - # Threshold mask to binary (ensure clean white/black) - _, mask_binary = cv2.threshold(mask_array, 127, 255, cv2.THRESH_BINARY) - - # Inpaint radius scales with image size for better results - inpaint_radius = max(3, min(img_array.shape[0], img_array.shape[1]) // 200) - - emit_progress(50, "Erasing") - result = cv2.inpaint(img_array, mask_binary, inpaint_radius, cv2.INPAINT_TELEA) - - emit_progress(90, "Saving") - result_rgb = cv2.cvtColor(result, cv2.COLOR_BGR2RGB) - Image.fromarray(result_rgb).save(output_path) - - print(json.dumps({"success": True, "method": "opencv-telea"})) - - except ImportError: - print( - json.dumps( - { - "success": False, - "error": "Object eraser requires OpenCV. Install with: pip install opencv-python-headless", - } - ) - ) + import onnxruntime as ort + except ImportError as e: + print(json.dumps({ + "success": False, + "error": f"Missing dependency: {e}. Requires opencv-python-headless and onnxruntime.", + })) sys.exit(1) - except ImportError: - print( - json.dumps( - { - "success": False, - "error": "Pillow is not installed. Install with: pip install Pillow", - } + emit_progress(10, "Loading model") + model_path = _get_model_path() + + # Configure ONNX Runtime session + providers = ["CPUExecutionProvider"] + if "CUDAExecutionProvider" in ort.get_available_providers(): + providers.insert(0, "CUDAExecutionProvider") + + session = ort.InferenceSession(model_path, providers=providers) + + emit_progress(20, "Loading images") + img = Image.open(input_path).convert("RGB") + mask = Image.open(mask_path).convert("L") + + orig_w, orig_h = img.size + img_array = np.array(img) + mask_array = np.array(mask) + + # Resize mask to match image if needed + if mask_array.shape[:2] != img_array.shape[:2]: + mask_array = cv2.resize( + mask_array, (orig_w, orig_h), interpolation=cv2.INTER_NEAREST ) + + # Threshold mask to binary + _, mask_binary = cv2.threshold(mask_array, 127, 255, cv2.THRESH_BINARY) + + emit_progress(30, "Preprocessing") + img_input = _preprocess_image(img_array) + mask_input = _preprocess_mask(mask_binary) + + emit_progress(40, "Erasing objects") + outputs = session.run( + None, + {"image": img_input, "mask": mask_input}, ) - sys.exit(1) + + emit_progress(75, "Compositing") + # Output shape: (1, 3, 512, 512) with values in [0, 255] + raw_output = outputs[0][0] # (3, 512, 512) + raw_output = np.transpose(raw_output, (1, 2, 0)) # (512, 512, 3) + raw_output = np.clip(raw_output, 0, 255).astype(np.uint8) + + # Resize inpainted result back to original dimensions + inpainted_full = cv2.resize(raw_output, (orig_w, orig_h), interpolation=cv2.INTER_LANCZOS4) + + # Feathered composite: preserve quality outside mask, blend at edges + mask_full = mask_binary.astype(np.float32) / 255.0 + feather_r = max(3, min(orig_w, orig_h) // 200) + result = _feathered_composite(img_array, inpainted_full, mask_full, feather_r) + + emit_progress(90, "Saving") + Image.fromarray(result).save(output_path) + + print(json.dumps({"success": True, "method": "lama-onnx"})) + except Exception as e: print(json.dumps({"success": False, "error": str(e)})) sys.exit(1)