remove-bg-api / app.py
zgrafic's picture
Upload 4 files
2bacef8 verified
Raw
History Blame Contribute Delete
13.2 kB
import io
import base64
import re
import numpy as np
import cv2
import easyocr
import torch
import requests
from PIL import Image
from fastapi import FastAPI, File, UploadFile, HTTPException, Form
from fastapi.middleware.cors import CORSMiddleware
from fastapi.responses import Response
from rembg import remove, new_session
from simple_lama_inpainting import SimpleLama
from transformers import AutoImageProcessor, AutoModelForImageClassification
app = FastAPI(title="ZGrafic API")
# Permite que tu página en GitHub Pages (u otro origen) llame a este servidor.
app.add_middleware(
CORSMiddleware,
allow_origins=["*"],
allow_methods=["*"],
allow_headers=["*"],
)
@app.on_event("startup")
async def start_telegram_bot():
import asyncio
from bot import run_bot
asyncio.create_task(run_bot())
# BiRefNet es una arquitectura más nueva (2024) que en la práctica
# supera a u2net/isnet en escenas complejas con varios sujetos. Es más
# pesado y lento en CPU, pero vale la pena para fotos difíciles.
session = new_session("birefnet-general")
# Para el endpoint /extract-text: detecta texto (OCR) y reconstruye el
# fondo donde estaba (inpainting), para volverlo editable.
ocr_reader = easyocr.Reader(["es", "en"], gpu=False)
lama = SimpleLama()
TEXT_PADDING = 4
# Identifica la fuente más parecida entre 48 fuentes estándar (licencia
# MIT, sin restricción de uso comercial).
font_processor = AutoImageProcessor.from_pretrained("gaborcselle/font-identifier")
font_model = AutoModelForImageClassification.from_pretrained("gaborcselle/font-identifier")
font_model.eval()
# El alpha matting es pesado: en imágenes grandes puede quedarse sin
# memoria o tardar demasiado en el CPU gratuito. Por encima de este
# tamaño, lo desactivamos automáticamente aunque el usuario lo pida.
MATTING_MAX_PIXELS = 1_500_000 # ~ 1500x1000
def resize_if_needed(pil_img, max_side=4000):
w, h = pil_img.size
if max(w, h) > max_side:
ratio = max_side / max(w, h)
pil_img = pil_img.resize((int(w * ratio), int(h * ratio)), Image.LANCZOS)
return pil_img
def recolor_from_original(png_bytes, original_rgb_array):
"""rembg deja en negro el color de cualquier píxel que considera
transparente (total o parcialmente), lo que genera parches o motas
negras en zonas de textura difícil. Reconstruye el color SIEMPRE desde
la foto original, y usa el resultado de la IA solo para el canal alfa
(transparencia)."""
img = Image.open(io.BytesIO(png_bytes)).convert("RGBA")
arr = np.array(img)
arr[:, :, 0:3] = original_rgb_array
out = Image.fromarray(arr, "RGBA")
buf = io.BytesIO()
out.save(buf, format="PNG")
return buf.getvalue()
def remove_stray_specks(png_bytes, erode_px=31, safe_margin_px=45, min_core_area_ratio=0.0015):
"""Elimina restos de "primer plano" que cuelgan de un mechón/hilo muy
delgado -como un pelo suelto que termina en una mancha-, sin arriesgar
personas reales (incluso si están separadas del resto, como en una
foto grupal).
La clave: erosiona la máscara primero. Cualquier mechón delgado
desaparece por completo con la erosión, mientras que una persona real
(una masa sólida) sobrevive como un núcleo más chico pero presente.
Se arma la "zona segura" a partir de esos núcleos sólidos, y se borra
todo lo que quede afuera -sin importar si técnicamente estaba
"conectado" por un hilo finito en la máscara original."""
img = Image.open(io.BytesIO(png_bytes)).convert("RGBA")
arr = np.array(img)
alpha = arr[:, :, 3]
fg = (alpha > 60).astype(np.uint8)
erode_kernel = np.ones((erode_px, erode_px), np.uint8)
eroded = cv2.erode(fg, erode_kernel, iterations=1)
num_labels, labels, stats, _ = cv2.connectedComponentsWithStats(eroded, connectivity=8)
if num_labels <= 1:
return png_bytes # no quedó ningún núcleo sólido: no tocar nada, por seguridad
total_area = fg.shape[0] * fg.shape[1]
safe_core = np.zeros_like(fg)
for label in range(1, num_labels):
area = stats[label, cv2.CC_STAT_AREA]
if area >= total_area * min_core_area_ratio:
safe_core[labels == label] = 1
if not safe_core.any():
return png_bytes
dilate_kernel = np.ones((safe_margin_px, safe_margin_px), np.uint8)
safe_zone = cv2.dilate(safe_core, dilate_kernel, iterations=1) > 0
stray = (alpha > 10) & (~safe_zone)
if not stray.any():
return png_bytes
alpha[stray] = 0
arr[:, :, 3] = alpha
out = Image.fromarray(arr, "RGBA")
buf = io.BytesIO()
out.save(buf, format="PNG")
return buf.getvalue()
def fill_enclosed_holes(png_bytes, max_hole_area_ratio=0.02):
"""Si queda algún agujero (transparente) completamente rodeado de
primer plano -como una cara que el modelo borró por error-, lo
detecta geométricamente y lo hace opaco. El color ya es correcto en
todos lados gracias a recolor_from_original, así que acá solo hay
que tocar la transparencia."""
img = Image.open(io.BytesIO(png_bytes)).convert("RGBA")
arr = np.array(img)
alpha = arr[:, :, 3]
fg = (alpha > 10).astype(np.uint8) * 255
# Sella canales delgados que conectan un "agujero" (como una cara
# borrada por error) con el fondo real, para que se detecten como
# una isla encerrada en vez de quedar pegados al fondo.
kernel = np.ones((7, 7), np.uint8)
fg_closed = cv2.morphologyEx(fg, cv2.MORPH_CLOSE, kernel)
contours, hierarchy = cv2.findContours(fg_closed, cv2.RETR_CCOMP, cv2.CHAIN_APPROX_SIMPLE)
if hierarchy is None:
return png_bytes
hierarchy = hierarchy[0]
total_area = fg.shape[0] * fg.shape[1]
fill_mask = np.zeros_like(fg)
for i, h in enumerate(hierarchy):
_, _, _, parent = h
if parent != -1: # es un agujero (contorno hijo de otro contorno)
area = cv2.contourArea(contours[i])
if area < total_area * max_hole_area_ratio:
cv2.drawContours(fill_mask, contours, i, 255, thickness=cv2.FILLED)
if not fill_mask.any():
return png_bytes
alpha[fill_mask == 255] = 255
arr[:, :, 3] = alpha
out = Image.fromarray(arr, "RGBA")
buf = io.BytesIO()
out.save(buf, format="PNG")
return buf.getvalue()
@app.get("/")
def health():
return {"status": "ok", "version": "erosion-core-v3"}
@app.post("/remove-bg")
async def remove_bg(file: UploadFile = File(...), matting: bool = Form(True)):
if not file.content_type or not file.content_type.startswith("image/"):
raise HTTPException(status_code=400, detail="El archivo debe ser una imagen.")
input_bytes = await file.read()
try:
pil_img = Image.open(io.BytesIO(input_bytes)).convert("RGB")
pil_img = resize_if_needed(pil_img)
buf_in = io.BytesIO()
pil_img.save(buf_in, format="PNG")
resized_bytes = buf_in.getvalue()
use_matting = matting and (pil_img.width * pil_img.height <= MATTING_MAX_PIXELS)
if use_matting:
output_bytes = remove(
resized_bytes,
session=session,
alpha_matting=True,
alpha_matting_foreground_threshold=240,
alpha_matting_background_threshold=10,
alpha_matting_erode_size=10,
)
else:
# Sin alpha matting: usa la máscara del modelo tal cual, sin el
# refinamiento extra. Más robusto en fotos con mucha gente o
# bordes complejos, aunque el cabello suelto se ve algo peor.
output_bytes = remove(resized_bytes, session=session)
output_bytes = recolor_from_original(output_bytes, np.array(pil_img))
output_bytes = remove_stray_specks(output_bytes)
output_bytes = fill_enclosed_holes(output_bytes)
except Exception as e:
raise HTTPException(status_code=500, detail=f"No se pudo procesar la imagen: {e}")
return Response(content=output_bytes, media_type="image/png")
def estimate_text_color(image_bgr, x1, y1, x2, y2):
"""Aproxima el color del texto separando, dentro del recuadro, los
píxeles del texto de los del fondo (el texto suele ser la clase
minoritaria dentro de su propio recuadro)."""
crop = image_bgr[y1:y2, x1:x2]
if crop.size == 0:
return "#000000"
gray = cv2.cvtColor(crop, cv2.COLOR_BGR2GRAY)
_, thresh = cv2.threshold(gray, 0, 255, cv2.THRESH_BINARY + cv2.THRESH_OTSU)
if np.sum(thresh == 255) > np.sum(thresh == 0):
thresh = cv2.bitwise_not(thresh)
text_pixels = crop[thresh == 255]
if len(text_pixels) == 0:
text_pixels = crop.reshape(-1, 3)
b, g, r = np.mean(text_pixels, axis=0)
return "#{:02x}{:02x}{:02x}".format(int(r), int(g), int(b))
def identify_font(pil_crop):
"""Devuelve el nombre de la fuente (de un set de 48 fuentes estándar)
más parecida a la del recorte de texto dado."""
try:
crop = pil_crop.convert("RGB")
# Los recortes chicos le dan poco detalle al clasificador; agrandarlos
# mejora bastante la precisión del reconocimiento.
if min(crop.size) < 200:
ratio = 200 / min(crop.size)
crop = crop.resize((int(crop.width * ratio), int(crop.height * ratio)), Image.LANCZOS)
inputs = font_processor(images=crop, return_tensors="pt")
with torch.no_grad():
logits = font_model(**inputs).logits
predicted_id = logits.argmax(-1).item()
return font_model.config.id2label[predicted_id]
except Exception:
return None
@app.get("/font-file")
def font_file(name: str):
"""Consigue el archivo .ttf real de una fuente de Google Fonts, para
poder incrustarla de verdad en el PDF (no solo mostrarla en pantalla).
Le pedimos a Google con un User-Agent viejo a propósito: así el
servidor responde con .ttf en vez del .woff2 que le da a los
navegadores modernos (jsPDF necesita .ttf)."""
try:
css_url = f"https://fonts.googleapis.com/css2?family={name.replace(' ', '+')}:wght@700"
headers = {"User-Agent": "Mozilla/5.0 (Windows NT 6.1; rv:2.0.1) Gecko/20100101 Firefox/4.0.1"}
css_resp = requests.get(css_url, headers=headers, timeout=15)
css_resp.raise_for_status()
match = re.search(r"url\((https://fonts\.gstatic\.com/[^)]+\.ttf)\)", css_resp.text)
if not match:
raise HTTPException(status_code=404, detail="No se encontró el archivo .ttf para esa fuente.")
ttf_resp = requests.get(match.group(1), timeout=15)
ttf_resp.raise_for_status()
return Response(content=ttf_resp.content, media_type="font/ttf")
except HTTPException:
raise
except Exception as e:
raise HTTPException(status_code=500, detail=f"No se pudo conseguir la fuente: {e}")
@app.post("/extract-text")
async def extract_text(file: UploadFile = File(...)):
if not file.content_type or not file.content_type.startswith("image/"):
raise HTTPException(status_code=400, detail="El archivo debe ser una imagen.")
input_bytes = await file.read()
try:
pil_img = Image.open(io.BytesIO(input_bytes)).convert("RGB")
pil_img = resize_if_needed(pil_img, max_side=1600)
image_bgr = cv2.cvtColor(np.array(pil_img), cv2.COLOR_RGB2BGR)
h, w = image_bgr.shape[:2]
results = ocr_reader.readtext(image_bgr)
texts = []
mask = np.zeros((h, w), dtype=np.uint8)
for bbox, text, conf in results:
if conf < 0.35 or not text.strip():
continue
xs = [p[0] for p in bbox]
ys = [p[1] for p in bbox]
x1 = max(0, int(min(xs)) - TEXT_PADDING)
y1 = max(0, int(min(ys)) - TEXT_PADDING)
x2 = min(w, int(max(xs)) + TEXT_PADDING)
y2 = min(h, int(max(ys)) + TEXT_PADDING)
if x2 <= x1 or y2 <= y1:
continue
color = estimate_text_color(image_bgr, x1, y1, x2, y2)
crop_rgb = cv2.cvtColor(image_bgr[y1:y2, x1:x2], cv2.COLOR_BGR2RGB)
font_name = identify_font(Image.fromarray(crop_rgb))
texts.append({
"text": text,
"x": x1,
"y": y1,
"width": x2 - x1,
"height": y2 - y1,
"fontSize": max(8, y2 - y1),
"color": color,
"font": font_name,
})
cv2.rectangle(mask, (x1, y1), (x2, y2), 255, thickness=-1)
if texts:
mask_pil = Image.fromarray(mask)
background_pil = lama(pil_img, mask_pil)
else:
background_pil = pil_img
buf = io.BytesIO()
background_pil.save(buf, format="PNG")
background_b64 = base64.b64encode(buf.getvalue()).decode("ascii")
except Exception as e:
raise HTTPException(status_code=500, detail=f"No se pudo procesar la imagen: {e}")
return {
"width": w,
"height": h,
"background": "data:image/png;base64," + background_b64,
"texts": texts,
}