juanpajedrez
Initial commit: First push chestnut mini binary
f7316f8
Raw
History Blame Contribute Delete
4.72 kB
"""
Author: Juan Pablo Triana Martinez
Gradio app β€” LinkNet Binary Text Segmentation (DocLayNet).
Entry-point for HuggingFace Spaces.
"""
import gradio as gr
import numpy as np
import torch
from PIL import Image
from torchvision import transforms
from model import create_binary_model
# ── Constants ──────────────────────────────────────────────────────────────────
MEAN = [0.9329, 0.9343, 0.9341]
STD = [0.1651, 0.1592, 0.1623]
IMG_SIZE = 512
THRESHOLD = 0.5
MODEL_PATH = "linknet_binary_doclaynet_20_percent_seed_7_ce_0.25_dice_1.0_.pth"
# ── Transforms ─────────────────────────────────────────────────────────────────
inference_transform = transforms.Compose([
transforms.Resize((IMG_SIZE, IMG_SIZE)),
transforms.ToTensor(),
transforms.Normalize(mean=MEAN, std=STD),
])
api_transform = transforms.Compose([
transforms.Resize((IMG_SIZE, IMG_SIZE)),
transforms.ToTensor(),
])
# ── Model ──────────────────────────────────────────────────────────────────────
device = "cpu"
model = create_binary_model().to(device)
model.load_state_dict(torch.load(MODEL_PATH, map_location=device))
model.eval()
# ── Helpers ────────────────────────────────────────────────────────────────────
def _logits_to_binary_mask(logits: torch.Tensor, threshold: float = 0.5) -> torch.Tensor:
probs = torch.sigmoid(logits)
return (probs > threshold).float()
def _overlay_mask(image: Image.Image, mask: torch.Tensor) -> np.ndarray:
"""Green-channel overlay of binary mask on the original image."""
img_np = api_transform(image).permute(1, 2, 0).numpy() # (H, W, 3) in [0,1]
mask_np = mask.squeeze().cpu().numpy() # (H, W) in {0,1}
overlay = img_np.copy()
overlay[:, :, 1] = np.maximum(overlay[:, :, 1], mask_np)
return (overlay * 255).clip(0, 255).astype(np.uint8)
# ── Inference ──────────────────────────────────────────────────────────────────
def inference(image: Image.Image, gt_mask: Image.Image = None):
img_tensor = inference_transform(image).unsqueeze(0) # (1, 3, H, W)
with torch.inference_mode():
logits = model(img_tensor) # (1, 1, H, W)
pred_mask = _logits_to_binary_mask(logits, THRESHOLD) # (1, 1, H, W)
pred_mask_np = pred_mask.squeeze().cpu().numpy() # (H, W) for Gradio display
overlay = _overlay_mask(image, pred_mask)
if gt_mask is None:
return pred_mask_np, overlay, "No ground truth provided β†’ metrics unavailable"
gt_tensor = api_transform(gt_mask.convert("L")).unsqueeze(0) # (1, 1, H, W)
gt_binary = (gt_tensor > 0.5).float()
inter = (pred_mask * gt_binary).sum()
union = ((pred_mask + gt_binary) > 0).float().sum()
iou = (inter / (union + 1e-7)).item()
dice = (2 * inter / (pred_mask.sum() + gt_binary.sum() + 1e-7)).item()
metrics_text = f"IoU (Jaccard): {iou:.4f} Dice (F1): {dice:.4f}"
return pred_mask_np, overlay, metrics_text
# ── Gradio Interface ───────────────────────────────────────────────────────────
demo = gr.Interface(
fn=inference,
inputs=[
gr.Image(type="pil", label="Input Image"),
gr.Image(type="pil", label="Ground Truth Mask (optional)"),
],
outputs=[
gr.Image(label="Predicted Binary Mask"),
gr.Image(label="Overlay"),
gr.Textbox(label="Metrics"),
],
title="LinkNet Binary Text Segmentation β€” DocLayNet ChestNut 🌰",
description=(
"Upload a document image to obtain its binary text segmentation mask. "
"Optionally upload a ground-truth binary mask to compute IoU and Dice metrics."
),
examples=[
["examples/binary_1_img.png", "examples/binary_1_mask.png"],
["examples/binary_2_img.png", "examples/binary_2_mask.png"],
["examples/binary_3_img.png", "examples/binary_3_mask.png"],
],
)
if __name__ == "__main__":
demo.launch()