""" 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()