imgtools-ai / main.py
ilhamdev's picture
feat(ai): configure Gradio SDK with ZeroGPU support and REST API mounting
824ac86
Raw History Blame Contribute Delete
21.5 kB
import os
import io
import json
import asyncio
import logging
from contextlib import asynccontextmanager
from typing import Dict, Any, Optional, Literal, List
from fastapi import FastAPI, UploadFile, File, Form, HTTPException, Request
from fastapi.middleware.cors import CORSMiddleware
from fastapi.responses import StreamingResponse, JSONResponse
from PIL import Image, ImageOps, ImageFilter
# Konfigurasi Logging
logging.basicConfig(level=logging.INFO, format="%(asctime)s [%(levelname)s] %(message)s")
logger = logging.getLogger("imgtools-ai")
# Konfigurasi Environment Variables
ALLOWED_ORIGINS_RAW = os.getenv("ALLOWED_ORIGINS", "http://localhost:3000")
ALLOWED_ORIGINS = [origin.strip() for origin in ALLOWED_ORIGINS_RAW.split(",") if origin.strip()]
MAX_PIXELS = int(os.getenv("MAX_PIXELS", 24_000_000)) # Default 24 Megapixels
MAX_BYTES = int(os.getenv("MAX_BYTES", 25 * 1024 * 1024)) # Default 25 MB
MAX_CONCURRENCY = int(os.getenv("MAX_CONCURRENCY", 2)) # Default 2 concurrent requests
# Semaphore untuk membatasi eksekusi serentak pada CPU/RAM terbatas
concurrency_semaphore = asyncio.Semaphore(MAX_CONCURRENCY)
# Dukungan ZeroGPU jika berjalan pada Hugging Face Spaces dengan SDK Gradio
try:
import spaces
def gpu_decorator(fn):
return spaces.GPU(fn)
except Exception:
def gpu_decorator(fn):
return fn
# Cache penyimpanan model ML (Lazy-loaded)
_models: Dict[str, Any] = {}
def model(name: str):
"""
Lazy-load dan cache model ML dalam dictionary.
Mendukung:
- 'remove-bg' -> briaai/RMBG-1.4
- 'upscale' -> caidas/swin2SR-realworld-sr-x4-64-bsrl (device=0 jika CUDA/ZeroGPU, -1 jika CPU)
- 'enhance' -> caidas/swin2SR-classical-sr-x2-64 (device=0 jika CUDA/ZeroGPU, -1 jika CPU)
- 'face' -> ultralytics YOLO("Bingsu/yolov8n-face")
"""
if name not in _models:
import torch
device = 0 if torch.cuda.is_available() else -1
logger.info(f"Memuat model '{name}' ke memori (device={device})...")
if name == "remove-bg":
from transformers import pipeline
_models[name] = pipeline(
"image-segmentation",
model="briaai/RMBG-1.4",
device=device,
trust_remote_code=True,
)
logger.info("Model 'remove-bg' (briaai/RMBG-1.4) berhasil dimuat.")
elif name == "upscale":
from transformers import pipeline
_models[name] = pipeline(
"image-to-image",
model="caidas/swin2SR-realworld-sr-x4-64-bsrl",
device=device,
)
logger.info("Model 'upscale' (caidas/swin2SR-realworld-sr-x4-64-bsrl) berhasil dimuat.")
elif name == "enhance":
from transformers import pipeline
_models[name] = pipeline(
"image-to-image",
model="caidas/swin2SR-classical-sr-x2-64",
device=device,
)
logger.info("Model 'enhance' (caidas/swin2SR-classical-sr-x2-64) berhasil dimuat.")
elif name == "face":
from ultralytics import YOLO
yolo_model = YOLO("Bingsu/yolov8n-face")
if torch.cuda.is_available():
yolo_model.to("cuda")
_models[name] = yolo_model
logger.info("Model 'face' (Bingsu/yolov8n-face) berhasil dimuat.")
else:
raise ValueError(f"Model '{name}' tidak didukung.")
return _models[name]
async def read_image(file: UploadFile) -> Image.Image:
"""
Helper untuk membaca berkas gambar yang diunggah:
- Menolak berkas jika ukuran > MAX_BYTES (HTTP 413).
- Membuka dengan Pillow & menolak jika format tidak dikenal (HTTP 415).
- Menormalkan orientasi EXIF (exif_transpose).
- Downscale proporsional jika jumlah piksel > MAX_PIXELS.
"""
contents = await file.read()
if len(contents) > MAX_BYTES:
raise HTTPException(
status_code=413,
detail=f"Ukuran berkas ({len(contents)} byte) melebihi batas maksimal {MAX_BYTES} byte ({MAX_BYTES // (1024 * 1024)}MB).",
)
try:
img = Image.open(io.BytesIO(contents))
img.load()
except Exception:
raise HTTPException(
status_code=415,
detail="Format gambar tidak didukung atau berkas rusak.",
)
# Normalisasi orientasi foto sesuai metadata EXIF
try:
img = ImageOps.exif_transpose(img)
except Exception:
pass
# Periksa dan downscale jika melebihi batas maksimal piksel
orig_w, orig_h = img.size
total_pixels = orig_w * orig_h
if total_pixels > MAX_PIXELS:
scale = (MAX_PIXELS / total_pixels) ** 0.5
new_w = max(1, int(orig_w * scale))
new_h = max(1, int(orig_h * scale))
logger.info(f"Gambar di-downscale dari {orig_w}x{orig_h} ke {new_w}x{new_h} (> {MAX_PIXELS} px).")
img = img.resize((new_w, new_h), Image.Resampling.LANCZOS)
return img
def png_response(image: Image.Image) -> StreamingResponse:
"""
Helper untuk mengembalikan objek PIL Image sebagai StreamingResponse bertipe image/png.
"""
buf = io.BytesIO()
image.save(buf, format="PNG", optimize=True)
buf.seek(0)
return StreamingResponse(buf, media_type="image/png")
def jpeg_response(image: Image.Image, quality: int = 92) -> StreamingResponse:
"""
Helper untuk mengembalikan objek PIL Image sebagai StreamingResponse bertipe image/jpeg.
"""
buf = io.BytesIO()
rgb_img = image.convert("RGB") if image.mode != "RGB" else image
rgb_img.save(buf, format="JPEG", quality=quality, optimize=True)
buf.seek(0)
return StreamingResponse(buf, media_type="image/jpeg")
@asynccontextmanager
async def lifespan(app: FastAPI):
"""
Event startup & shutdown lifecycle:
Melakukan warm-up model 'remove-bg' saat server dinyalakan.
Jika gagal karena keterbatasan memori saat boot, log warning saja.
"""
logger.info("Memulai layanan ImgTools AI...")
try:
logger.info("Melakukan warm-up model 'remove-bg'...")
model("remove-bg")
except Exception as e:
logger.warning(f"Warm-up model ditangguhkan (akan dimuat saat request pertama): {e}")
yield
logger.info("Mematikan layanan ImgTools AI...")
_models.clear()
# Inisialisasi Aplikasi FastAPI
app = FastAPI(
title="ImgTools AI",
description="REST API AI untuk manipulasi, segmentasi, peningkatan kualitas, dan deteksi wajah pada Hugging Face Spaces",
version="1.0.0",
lifespan=lifespan,
)
# Konfigurasi CORS
is_wildcard = "*" in ALLOWED_ORIGINS or len(ALLOWED_ORIGINS) == 0 or "http://localhost:3000" in ALLOWED_ORIGINS
app.add_middleware(
CORSMiddleware,
allow_origins=["*"] if is_wildcard else ALLOWED_ORIGINS,
allow_credentials=False if is_wildcard else True,
allow_methods=["*"],
allow_headers=["*"],
)
@app.exception_handler(Exception)
async def global_exception_handler(request: Request, exc: Exception):
"""
Handler global untuk menangani unhandled exception secara konsisten (HTTP 500).
"""
if isinstance(exc, HTTPException):
return JSONResponse(status_code=exc.status_code, content={"detail": exc.detail})
logger.exception(f"Unhandled error pada endpoint {request.url.path}: {exc}")
return JSONResponse(
status_code=500,
content={"detail": f"Terjadi kesalahan internal pada server: {str(exc)[:150]}"},
)
@app.get("/health")
def health_check():
"""
Health check endpoint untuk Docker & orchestrator.
"""
return {"status": "ok"}
@app.get("/")
def root():
"""
Root info endpoint menampilkan status dan route yang tersedia.
"""
return {
"service": "ImgTools AI",
"status": "online",
"endpoints": {
"health": "/health",
"remove_bg": "/api/remove-bg",
"upscale": "/api/upscale",
"enhance": "/api/enhance",
"face_blur": "/api/face-blur",
"raw_to_jpg": "/api/raw-to-jpg",
},
}
@app.post("/api/remove-bg")
async def remove_background(file: UploadFile = File(...)):
"""
Hapus latar belakang foto menggunakan model briaai/RMBG-1.4:
1. Baca dan validasi gambar masukan.
2. Resize sementara ke 1024x1024 untuk inferensi optimal model.
3. Ekstraksi mask segmentasi.
4. Resize mask kembali ke resolusi asli.
5. Pasang mask sebagai alpha channel pada gambar asli RGBA.
6. Kembalikan respons PNG transparan.
"""
async with concurrency_semaphore:
try:
# Baca gambar asli
orig_img = await read_image(file)
orig_w, orig_h = orig_img.size
# Konversi ke RGB untuk inferensi
rgb_img = orig_img.convert("RGB")
# Resize ke 1024x1024 untuk model
img_1024 = rgb_img.resize((1024, 1024), Image.Resampling.BILINEAR)
# Inferensi model remove-bg
pipe = model("remove-bg")
loop = asyncio.get_running_loop()
output = await loop.run_in_executor(None, pipe, img_1024)
# Ekstraksi mask hasil segmentasi
if isinstance(output, list) and len(output) > 0:
first = output[0]
mask_candidate = first["mask"] if isinstance(first, dict) and "mask" in first else first
else:
mask_candidate = output
# Dapatkan channel grayscale (L) dari mask
if isinstance(mask_candidate, Image.Image):
if mask_candidate.mode == "RGBA":
mask = mask_candidate.split()[-1]
elif mask_candidate.mode != "L":
mask = mask_candidate.convert("L")
else:
mask = mask_candidate
else:
raise HTTPException(status_code=500, detail="Format output model tidak valid.")
# Resize mask balik ke ukuran gambar asli
resized_mask = mask.resize((orig_w, orig_h), Image.Resampling.BILINEAR)
# Pasang mask sebagai alpha channel pada gambar asli
rgba_result = orig_img.convert("RGBA")
rgba_result.putalpha(resized_mask)
return png_response(rgba_result)
except HTTPException:
raise
except Exception as e:
logger.exception(f"Error pada /api/remove-bg: {e}")
raise HTTPException(status_code=500, detail=f"Gagal menghapus latar belakang: {str(e)[:150]}")
@app.post("/api/upscale")
async def upscale(
request: Request,
file: UploadFile = File(...),
scale: int = Form(2),
max_side: int = Form(2048),
):
"""
Tingkatkan resolusi gambar (Super-Resolution):
- validasi scale in {2, 4} (HTTP 422 jika salah)
- thumbnail ke max_side
- jalankan pipeline swin2SR
- kembalikan PNG
"""
# Dukung override via query params jika ada
if "scale" in request.query_params:
try:
scale = int(request.query_params["scale"])
except ValueError:
raise HTTPException(status_code=422, detail="scale harus berupa integer.")
if "max_side" in request.query_params:
try:
max_side = int(request.query_params["max_side"])
except ValueError:
raise HTTPException(status_code=422, detail="max_side harus berupa integer.")
if scale not in (2, 4):
raise HTTPException(status_code=422, detail="scale harus bernilai 2 atau 4.")
async with concurrency_semaphore:
try:
img = await read_image(file)
# Thumbnail ke max_side (mempertahankan rasio aspek)
img.thumbnail((max_side, max_side), Image.Resampling.LANCZOS)
rgb_img = img.convert("RGB")
# Jalankan pipeline upscale swin2SR
pipe = model("upscale")
loop = asyncio.get_running_loop()
result = await loop.run_in_executor(None, pipe, rgb_img)
if isinstance(result, list) and len(result) > 0:
out_img = result[0].get("image", result[0]) if isinstance(result[0], dict) else result[0]
elif isinstance(result, dict) and "image" in result:
out_img = result["image"]
else:
out_img = result
if not isinstance(out_img, Image.Image):
raise ValueError("Output model tidak menghasilkan PIL Image yang valid.")
# Model default menghasilkan 4x upscale; jika scale 2, resize proporsional
if scale == 2:
target_w = max(1, img.width * 2)
target_h = max(1, img.height * 2)
out_img = out_img.resize((target_w, target_h), Image.Resampling.LANCZOS)
return png_response(out_img)
except HTTPException:
raise
except Exception as e:
logger.exception(f"Error pada /api/upscale: {e}")
raise HTTPException(status_code=500, detail=f"Gagal melakukan upscale gambar: {str(e)[:150]}")
@app.post("/api/enhance")
async def enhance(
request: Request,
file: UploadFile = File(...),
strength: float = Form(1.0),
):
"""
Tingkatkan ketajaman dan detail gambar:
- thumbnail ke 1600
- jalankan model enhance (swin2SR-classical-sr-x2-64)
- terapkan UnsharpMask(radius=1.5, percent=int(60*strength))
- kembalikan PNG
"""
if "strength" in request.query_params:
try:
strength = float(request.query_params["strength"])
except ValueError:
raise HTTPException(status_code=422, detail="strength harus berupa float.")
async with concurrency_semaphore:
try:
img = await read_image(file)
# Thumbnail ke 1600
img.thumbnail((1600, 1600), Image.Resampling.LANCZOS)
rgb_img = img.convert("RGB")
# Jalankan model enhance
pipe = model("enhance")
loop = asyncio.get_running_loop()
result = await loop.run_in_executor(None, pipe, rgb_img)
if isinstance(result, list) and len(result) > 0:
out_img = result[0].get("image", result[0]) if isinstance(result[0], dict) else result[0]
elif isinstance(result, dict) and "image" in result:
out_img = result["image"]
else:
out_img = result
if not isinstance(out_img, Image.Image):
raise ValueError("Output model tidak menghasilkan PIL Image yang valid.")
# Terapkan filter penajaman UnsharpMask
percent = max(0, int(60 * strength))
enhanced = out_img.filter(ImageFilter.UnsharpMask(radius=1.5, percent=percent))
return png_response(enhanced)
except HTTPException:
raise
except Exception as e:
logger.exception(f"Error pada /api/enhance: {e}")
raise HTTPException(status_code=500, detail=f"Gagal meningkatkan kualitas gambar: {str(e)[:150]}")
@app.post("/api/face-blur")
async def face_blur(
request: Request,
file: UploadFile = File(...),
blur: int = Form(25),
mode: Literal["gaussian", "pixelate"] = Form("gaussian"),
conf: float = Form(0.35),
boxes: Optional[str] = Form(None),
):
"""
Deteksi wajah & buramkan area sensitif:
- boxes: JSON string list [x, y, w, h] untuk area manual tambahan.
- deteksi wajah via YOLO (conf), digabung dengan boxes manual.
- tiap region: Gaussian blur (radius=blur) ATAU pixelate (resize kecil lalu NEAREST balik).
- clamp koordinat ke batas gambar; lewati region < 3px.
- kembalikan PNG.
"""
if "blur" in request.query_params:
try:
blur = int(request.query_params["blur"])
except ValueError:
raise HTTPException(status_code=422, detail="blur harus berupa integer.")
if "mode" in request.query_params:
m = request.query_params["mode"]
if m in ("gaussian", "pixelate"):
mode = m
else:
raise HTTPException(status_code=422, detail="mode harus 'gaussian' atau 'pixelate'.")
if "conf" in request.query_params:
try:
conf = float(request.query_params["conf"])
except ValueError:
raise HTTPException(status_code=422, detail="conf harus berupa float.")
if "boxes" in request.query_params:
boxes = request.query_params["boxes"]
if mode not in ("gaussian", "pixelate"):
raise HTTPException(status_code=422, detail="mode harus 'gaussian' atau 'pixelate'.")
# Parsing kotak manual jika disediakan
all_boxes: List[List[int]] = []
if boxes:
try:
parsed = json.loads(boxes)
if isinstance(parsed, list):
if len(parsed) > 0 and isinstance(parsed[0], (int, float)) and len(parsed) == 4:
all_boxes.append([int(v) for v in parsed])
else:
for item in parsed:
if isinstance(item, (list, tuple)) and len(item) == 4:
all_boxes.append([int(item[0]), int(item[1]), int(item[2]), int(item[3])])
except Exception as e:
logger.warning(f"Gagal parse boxes JSON: {e}")
raise HTTPException(status_code=422, detail="Format JSON parameter boxes tidak valid.")
async with concurrency_semaphore:
try:
orig_img = await read_image(file)
img = orig_img.copy()
w_img, h_img = img.size
# Deteksi wajah menggunakan model YOLO
try:
yolo = model("face")
rgb_img = img.convert("RGB")
loop = asyncio.get_running_loop()
yolo_results = await loop.run_in_executor(
None, lambda: yolo.predict(source=rgb_img, conf=conf, verbose=False)
)
if yolo_results and len(yolo_results) > 0:
res_boxes = yolo_results[0].boxes
if res_boxes is not None and hasattr(res_boxes, "xyxy"):
xyxy_arr = res_boxes.xyxy.cpu().numpy()
for box in xyxy_arr:
x1, y1, x2, y2 = box
bx = int(x1)
by = int(y1)
bw = int(x2 - x1)
bh = int(y2 - y1)
all_boxes.append([bx, by, bw, bh])
except Exception as e:
logger.warning(f"Deteksi wajah YOLO mengalami kendala: {e}")
# Proses setiap area target
for b in all_boxes:
bx, by, bw, bh = b
# Clamp koordinat ke batas dimensi gambar
x1 = max(0, min(bx, w_img))
y1 = max(0, min(by, h_img))
x2 = max(0, min(bx + bw, w_img))
y2 = max(0, min(by + bh, h_img))
reg_w = x2 - x1
reg_h = y2 - y1
# Lewati region < 3px
if reg_w < 3 or reg_h < 3:
continue
region = img.crop((x1, y1, x2, y2))
if mode == "gaussian":
blurred = region.filter(ImageFilter.GaussianBlur(radius=max(1, blur)))
else: # pixelate
pixel_size = max(2, blur)
small_w = max(1, reg_w // pixel_size)
small_h = max(1, reg_h // pixel_size)
blurred = region.resize((small_w, small_h), Image.Resampling.NEAREST).resize(
(reg_w, reg_h), Image.Resampling.NEAREST
)
img.paste(blurred, (x1, y1))
return png_response(img)
except HTTPException:
raise
except Exception as e:
logger.exception(f"Error pada /api/face-blur: {e}")
raise HTTPException(status_code=500, detail=f"Gagal memproses face-blur: {str(e)[:150]}")
@app.post("/api/raw-to-jpg")
async def raw_to_jpg(file: UploadFile = File(...)):
"""
Konversi gambar kamera RAW ke JPEG:
- rawpy.imread(BytesIO) -> postprocess() -> simpan JPEG quality 92
- Mengembalikan error 422 jika berkas RAW tidak valid atau gagal diproses
"""
async with concurrency_semaphore:
try:
contents = await file.read()
if len(contents) > MAX_BYTES:
raise HTTPException(
status_code=413,
detail=f"Ukuran berkas ({len(contents)} byte) melebihi batas maksimal {MAX_BYTES} byte.",
)
try:
import rawpy
with rawpy.imread(io.BytesIO(contents)) as raw:
rgb = raw.postprocess()
img = Image.fromarray(rgb)
except HTTPException:
raise
except Exception as e:
logger.error(f"Gagal memproses RAW image: {e}")
raise HTTPException(
status_code=422,
detail="Gagal memproses berkas RAW. Pastikan format berkas RAW didukung dan tidak korup.",
)
return jpeg_response(img, quality=92)
except HTTPException:
raise
except Exception as e:
logger.exception(f"Error pada /api/raw-to-jpg: {e}")
raise HTTPException(status_code=500, detail=f"Gagal mengonversi RAW ke JPG: {str(e)[:150]}")