Team Ai
Apppublic

hetpatel-tenup/image-processor

sourceHugging Facemitupdated 2y agoView on Hugging Face
0likes
app.py213 linesDownload Raw Back to root
1import gradio as gr2from PIL import Image, ImageOps3import numpy as np4import os5import cv26import torch7from typing import Tuple8from gradio_imageslider import ImageSlider9from torchvision import transforms10import requests11from io import BytesIO12import zipfile13from model import SimpleGrayScaleModel14from torchvision.transforms import ToPILImage15 16device = "cuda" if torch.cuda.is_available() else "cpu"17 18torch.set_float32_matmul_precision('high')19torch.jit.script = lambda f: f20 21 22# Dummy Model23class DummyModel:24    def __init__(self, num_classes=1, image_size=(1024, 1024)):25        self.num_classes = num_classes26        self.image_size = image_size27 28    def __call__(self, input_tensor):29        batch_size, _, height, width = input_tensor.size()30        dummy_output = torch.rand((batch_size, self.num_classes, height, width))31        return dummy_output32 33 34# Original Model35class OriginalModel:36    def __init__(self):37        self.model = SimpleGrayScaleModel()38 39    def predict(self, input_tensor):40        # Convert tensor back to PIL Image if required41        to_pil = ToPILImage()42        input_image = to_pil(input_tensor.squeeze(0))  # Remove batch dimension43        return self.model.predict(input_image)44 45 46# Function to load appropriate model47def run_model(use_dummy: bool, model_path: str = None):48    if use_dummy:49        print("Using Dummy Model")50        return DummyModel(num_classes=1)51    else:52        print("Using Original Model")53        return OriginalModel()54 55 56# Image Preprocessor57class ImagePreprocessor:58    def __init__(self, resolution: Tuple[int, int] = (1024, 1024)) -> None:59        self.transform_image = transforms.Compose([60            transforms.Resize(resolution),61            transforms.ToTensor(),62            transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]),63        ])64 65    def proc(self, image: Image.Image) -> torch.Tensor:66        image = self.transform_image(image)67        return image68 69 70# Image Refinement (Foreground Extraction)71def refine_foreground(image, mask, r=90):72    if mask.size != image.size:73        mask = mask.resize(image.size)74    image = np.array(image) / 255.075    mask = np.array(mask) / 255.076    estimated_foreground = FB_blur_fusion_foreground_estimator_2(image, mask, r=r)77    image_masked = Image.fromarray((estimated_foreground * 255.0).astype(np.uint8))78    return image_masked79 80 81def FB_blur_fusion_foreground_estimator_2(image, alpha, r=90):82    alpha = alpha[:, :, None]83    F, blur_B = FB_blur_fusion_foreground_estimator(image, image, image, alpha, r)84    return FB_blur_fusion_foreground_estimator(image, F, blur_B, alpha, r=6)[0]85 86 87def FB_blur_fusion_foreground_estimator(image, F, B, alpha, r=90):88    if isinstance(image, Image.Image):89        image = np.array(image) / 255.090    blurred_alpha = cv2.blur(alpha, (r, r))[:, :, None]91    blurred_FA = cv2.blur(F * alpha, (r, r))92    blurred_F = blurred_FA / (blurred_alpha + 1e-5)93    blurred_B1A = cv2.blur(B * (1 - alpha), (r, r))94    blurred_B = blurred_B1A / ((1 - blurred_alpha) + 1e-5)95    F = blurred_F + alpha * (image - alpha * blurred_F - (1 - alpha) * blurred_B)96    F = np.clip(F, 0, 1)97    return F, blurred_B98 99 100# Prediction Function101def predict(images, resolution, use_dummy):102    assert images is not None, 'AssertionError: images cannot be None.'103 104    # Load the appropriate model105    model = run_model(use_dummy, model_path='your_model.pth')106 107    try:108        resolution = tuple(int(val) for val in resolution.strip().split('x'))109        if len(resolution) != 2:110            raise ValueError("Resolution must contain two values (width, height).")111    except:112        resolution = (1024, 1024)  # default resolution113        print('Invalid resolution input. Automatically changed to 1024x1024.')114 115    if isinstance(images, list):116        save_paths = []117        save_dir = 'preds-Model'118        if not os.path.exists(save_dir):119            os.makedirs(save_dir)120        tab_is_batch = True121    else:122        images = [images]123        tab_is_batch = False124 125    for idx_image, image_src in enumerate(images):126        if isinstance(image_src, str):127            if os.path.isfile(image_src):128                image_ori = Image.open(image_src)129            else:130                response = requests.get(image_src)131                image_data = BytesIO(response.content)132                image_ori = Image.open(image_data)133        else:134            image_ori = Image.fromarray(image_src)135 136        image = image_ori.convert('RGB')137        # Preprocess the image138        image_preprocessor = ImagePreprocessor(resolution=tuple(resolution))139        image_proc = image_preprocessor.proc(image)140        image_proc = image_proc.unsqueeze(0)141 142        # Prediction143        if use_dummy:144            preds = model(image_proc.to(device))145        else:146            preds = model.predict(image_proc.to(device))147 148        if isinstance(preds, torch.Tensor):  # Dummy Model149            pred = preds[0].sigmoid().cpu().squeeze()150            pred_pil = transforms.ToPILImage()(pred)151        else:  # Original Model returns PIL image152            pred_pil = preds153 154        # Show Results155        image_masked = refine_foreground(image, pred_pil)156        image_masked.putalpha(pred_pil.resize(image.size))157 158        torch.cuda.empty_cache()159 160        if tab_is_batch:161            save_file_path = os.path.join(save_dir, "{}.png".format(os.path.splitext(os.path.basename(image_src))[0]))162            image_masked.save(save_file_path)163            save_paths.append(save_file_path)164 165    if tab_is_batch:166        zip_file_path = os.path.join(save_dir, "{}.zip".format(save_dir))167        with zipfile.ZipFile(zip_file_path, 'w') as zipf:168            for file in save_paths:169                zipf.write(file, os.path.basename(file))170        return save_paths, zip_file_path171    else:172        return (image_masked, image_ori)173 174 175 176# Gradio Interface177def create_tab_image():178    return gr.Interface(179        fn=predict,180        inputs=[181            gr.Image(label='Upload an image'),182            gr.Dropdown(choices=["720x720", "1024x1024", "1080x1080"], value="1024x1024", label="Resolution"),183            gr.Checkbox(label="Use Dummy Model", value=True)  # Toggle button184        ],185        outputs=ImageSlider(label="Prediction Output", type="pil"),186        api_name="image"187    )188 189 190def create_batch_image():191    return gr.Interface(192        fn=predict,193        inputs=[194            gr.File(label="Upload multiple images", type="filepath", file_count="multiple"),195            gr.Dropdown(choices=["720x720", "1024x1024", "1080x1080"], value="1024x1024", label="Resolution"),196            gr.Checkbox(label="Use Dummy Model", value=True)  # Toggle button197        ],198        outputs=[gr.Gallery(label="Predictions"), gr.File(label="Download masked images.")],199        api_name="batch"200    )201 202 203# Create the demo interface with tabs204demo = gr.TabbedInterface(205    [create_tab_image(), create_batch_image()],206    ['Single Image Processing', 'Batch Image Processing'],207    title="Image Processing with Model"208)209 210# Launch the demo211if __name__ == "__main__":212    demo.launch()213