hetpatel-tenup/image-processor
0
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 