codetechdevx commited on
Commit
17983ff
·
1 Parent(s): ee499f9

updated download function

Browse files
Files changed (1) hide show
  1. app.py +144 -42
app.py CHANGED
@@ -1,17 +1,27 @@
1
  import gradio as gr
2
- from loadimg import load_img
3
  import spaces
4
  from transformers import AutoModelForImageSegmentation
5
  import torch
6
  from torchvision import transforms
 
 
 
7
 
8
- torch.set_float32_matmul_precision(["high", "highest"][0])
 
 
9
 
 
 
10
  birefnet = AutoModelForImageSegmentation.from_pretrained(
11
  "ZhengPeng7/BiRefNet", trust_remote_code=True
12
  )
13
- birefnet.to("cpu")
 
 
14
 
 
15
  transform_image = transforms.Compose(
16
  [
17
  transforms.Resize((1024, 1024)),
@@ -20,52 +30,144 @@ transform_image = transforms.Compose(
20
  ]
21
  )
22
 
23
- def fn(image):
24
- im = load_img(image, output_type="pil")
25
- im = im.convert("RGB")
26
- origin = im.copy()
27
- processed_image = process(im)
28
- return (processed_image, origin)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
29
 
30
- #@spaces.GPU
31
- def process(image):
 
 
 
 
 
 
 
 
 
 
 
32
  image_size = image.size
33
- input_images = transform_image(image).unsqueeze(0).to("cpu")
34
- # Prediction
 
 
35
  with torch.no_grad():
36
- preds = birefnet(input_images)[-1].sigmoid().cpu()
37
- pred = preds[0].squeeze()
38
- pred_pil = transforms.ToPILImage()(pred)
39
- mask = pred_pil.resize(image_size)
40
- image.putalpha(mask)
 
 
 
 
 
41
  return image
42
 
43
- def process_file(f):
44
- name_path = f.rsplit(".", 1)[0] + ".png"
45
- im = load_img(f, output_type="pil")
46
- im = im.convert("RGB")
47
- transparent = process(im)
48
- transparent.save(name_path)
49
- return name_path
50
-
51
- slider1 = gr.ImageSlider(label="Processed Image", type="pil", format="png")
52
- slider2 = gr.ImageSlider(label="Processed Image from URL", type="pil", format="png")
53
- image_upload = gr.Image(label="Upload an image")
54
- image_file_upload = gr.Image(label="Upload an image", type="filepath")
55
- url_input = gr.Textbox(label="Paste an image URL")
56
- output_file = gr.File(label="Output PNG File")
57
-
58
- # Example images
59
- chameleon = load_img("butterfly.jpeg", output_type="pil")
60
- url_example = "https://i.ibb.co/67B6Knk9/students-1807505-1280.jpg"
61
-
62
- tab1 = gr.Interface(fn, inputs=image_upload, outputs=slider1, examples=[chameleon], api_name="image")
63
- tab2 = gr.Interface(fn, inputs=url_input, outputs=slider2, examples=[url_example], api_name="text")
64
- tab3 = gr.Interface(process_file, inputs=image_file_upload, outputs=output_file, examples=["butterfly.jpeg"], api_name="png")
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
65
 
 
66
  demo = gr.TabbedInterface(
67
- [tab1, tab2, tab3], ["Image Upload", "URL Input", "File Output"], title="Background Removal Tool"
 
 
68
  )
69
 
70
  if __name__ == "__main__":
71
- demo.launch(show_error=True)
 
1
  import gradio as gr
2
+ from PIL import Image
3
  import spaces
4
  from transformers import AutoModelForImageSegmentation
5
  import torch
6
  from torchvision import transforms
7
+ import requests
8
+ from io import BytesIO
9
+ import os
10
 
11
+ # --- Model and Processor Setup ---
12
+ # Use a higher precision for matrix multiplication for better performance
13
+ torch.set_float32_matmul_precision("high")
14
 
15
+ # Load the BiRefNet model for image segmentation
16
+ # trust_remote_code=True is required for this model
17
  birefnet = AutoModelForImageSegmentation.from_pretrained(
18
  "ZhengPeng7/BiRefNet", trust_remote_code=True
19
  )
20
+ # Move the model to the available device (GPU if available, otherwise CPU)
21
+ device = "cuda" if torch.cuda.is_available() else "cpu"
22
+ birefnet.to(device)
23
 
24
+ # Define the image transformation pipeline
25
  transform_image = transforms.Compose(
26
  [
27
  transforms.Resize((1024, 1024)),
 
30
  ]
31
  )
32
 
33
+ # --- Helper Function to Load Images ---
34
+ def load_image(image_source, output_type="pil"):
35
+ """
36
+ Loads an image from a file path, URL, or numpy array.
37
+ """
38
+ if image_source is None:
39
+ return None
40
+
41
+ if isinstance(image_source, str):
42
+ if image_source.startswith("http"):
43
+ try:
44
+ response = requests.get(image_source)
45
+ response.raise_for_status()
46
+ image = Image.open(BytesIO(response.content))
47
+ except requests.exceptions.RequestException as e:
48
+ raise gr.Error(f"Could not fetch image from URL: {e}")
49
+ else:
50
+ image = Image.open(image_source)
51
+ elif hasattr(image_source, 'shape'): # Check if it's a numpy-like array
52
+ image = Image.fromarray(image_source)
53
+ else:
54
+ image = image_source # Assume it's already a PIL image
55
 
56
+ if output_type == "pil":
57
+ return image.convert("RGB")
58
+ return image
59
+
60
+ # --- Core Processing Function ---
61
+ # Use @spaces.GPU decorator if you plan to run this on a GPU-enabled Hugging Face Space
62
+ # @spaces.GPU
63
+ def process_image_to_transparent(image: Image.Image) -> Image.Image:
64
+ """
65
+ Takes a PIL image, removes the background, and returns a PIL image with an alpha channel.
66
+ """
67
+ if image is None:
68
+ return None
69
  image_size = image.size
70
+ # Unsqueeze adds a batch dimension, which the model expects
71
+ input_tensor = transform_image(image).unsqueeze(0).to(device)
72
+
73
+ # Prediction without tracking gradients for efficiency
74
  with torch.no_grad():
75
+ # The model returns multiple outputs; the last one is the primary segmentation map
76
+ preds = birefnet(input_tensor)[-1].sigmoid().cpu()
77
+
78
+ # Process the prediction tensor to create a mask
79
+ pred_tensor = preds[0].squeeze()
80
+ mask_pil = transforms.ToPILImage()(pred_tensor)
81
+ mask_resized = mask_pil.resize(image_size)
82
+
83
+ # Apply the mask as an alpha channel to the original image
84
+ image.putalpha(mask_resized)
85
  return image
86
 
87
+ # --- Gradio Interface Functions ---
88
+
89
+ def fn(image_source):
90
+ """
91
+ Handles image uploads and URLs, returning the processed image.
92
+ """
93
+ if image_source is None:
94
+ return None
95
+
96
+ pil_image = load_image(image_source, output_type="pil")
97
+ processed_image = process_image_to_transparent(pil_image)
98
+ return processed_image
99
+
100
+ def process_file(image_filepath):
101
+ """
102
+ Handles a single file upload and returns a downloadable processed file.
103
+ """
104
+ if image_filepath is None:
105
+ return None
106
+
107
+ # Define the output path for the new PNG file
108
+ base_name = os.path.basename(image_filepath.name) # Use .name for Gradio file objects
109
+ name, _ = os.path.splitext(base_name)
110
+ output_path = f"{name}_transparent.png"
111
+
112
+ # Load the image from the provided file path
113
+ pil_image = load_image(image_filepath.name, output_type="pil")
114
+
115
+ # Process the image
116
+ transparent_image = process_image_to_transparent(pil_image)
117
+
118
+ # Save the processed image to the new path
119
+ transparent_image.save(output_path)
120
+
121
+ # Return the path to the newly created file for download
122
+ return output_path
123
+
124
+ # --- Gradio UI Definition ---
125
+
126
+ # Define example images for the interface
127
+ example_image_path = "butterfly.jpeg"
128
+ # You should have a 'butterfly.jpeg' in the same directory or provide a full path
129
+ # For demonstration, let's create a dummy example image if it doesn't exist.
130
+ if not os.path.exists(example_image_path):
131
+ print(f"'{example_image_path}' not found. Creating a dummy image for example.")
132
+ try:
133
+ dummy_img = Image.new('RGB', (200, 200), color = 'red')
134
+ dummy_img.save(example_image_path)
135
+ except Exception as e:
136
+ print(f"Could not create dummy image: {e}")
137
+
138
+ example_url = "https://i.ibb.co/67B6Knk9/students-1807505-1280.jpg"
139
+
140
+ # Define the individual interfaces for each tab
141
+ tab1 = gr.Interface(
142
+ fn,
143
+ inputs=gr.Image(label="Upload an Image", type="pil"),
144
+ outputs=gr.Image(label="Processed Image", format="png"),
145
+ examples=[[example_image_path]],
146
+ api_name="image"
147
+ )
148
+
149
+ tab2 = gr.Interface(
150
+ fn,
151
+ inputs=gr.Textbox(label="Paste an Image URL"),
152
+ outputs=gr.Image(label="Processed Image", format="png"),
153
+ examples=[[example_url]],
154
+ api_name="text"
155
+ )
156
+
157
+ tab3 = gr.Interface(
158
+ process_file,
159
+ inputs=gr.File(label="Upload an Image File"),
160
+ outputs=gr.File(label="Download Processed PNG"),
161
+ examples=[[example_image_path]],
162
+ api_name="png"
163
+ )
164
 
165
+ # Combine the interfaces into a tabbed layout
166
  demo = gr.TabbedInterface(
167
+ [tab1, tab2, tab3],
168
+ ["Image Upload", "URL Input", "File Output"],
169
+ title="Background Removal Tool"
170
  )
171
 
172
  if __name__ == "__main__":
173
+ demo.launch(show_error=True)