timmeyer commited on
Commit
70ff6c9
·
verified ·
1 Parent(s): 09a39d2

Update app.py

Browse files
Files changed (1) hide show
  1. app.py +40 -4
app.py CHANGED
@@ -1,5 +1,41 @@
1
- from transformers import pipeline
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
2
  image_path = "https://farm5.staticflickr.com/4007/4322154488_997e69e4cf_z.jpg"
3
- pipe = pipeline("image-segmentation", model="briaai/RMBG-1.4", trust_remote_code=True)
4
- pillow_mask = pipe(img_path, return_mask = True) # outputs a pillow mask
5
- pillow_image = pipe(image_path) # applies mask on input and returns a pillow image
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from transformers import AutoModelForImageSegmentation
2
+ model = AutoModelForImageSegmentation.from_pretrained("briaai/RMBG-1.4",trust_remote_code=True)
3
+ def preprocess_image(im: np.ndarray, model_input_size: list) -> torch.Tensor:
4
+ if len(im.shape) < 3:
5
+ im = im[:, :, np.newaxis]
6
+ # orig_im_size=im.shape[0:2]
7
+ im_tensor = torch.tensor(im, dtype=torch.float32).permute(2,0,1)
8
+ im_tensor = F.interpolate(torch.unsqueeze(im_tensor,0), size=model_input_size, mode='bilinear')
9
+ image = torch.divide(im_tensor,255.0)
10
+ image = normalize(image,[0.5,0.5,0.5],[1.0,1.0,1.0])
11
+ return image
12
+
13
+ def postprocess_image(result: torch.Tensor, im_size: list)-> np.ndarray:
14
+ result = torch.squeeze(F.interpolate(result, size=im_size, mode='bilinear') ,0)
15
+ ma = torch.max(result)
16
+ mi = torch.min(result)
17
+ result = (result-mi)/(ma-mi)
18
+ im_array = (result*255).permute(1,2,0).cpu().data.numpy().astype(np.uint8)
19
+ im_array = np.squeeze(im_array)
20
+ return im_array
21
+
22
+ device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
23
+ model.to(device)
24
+
25
+ # prepare input
26
  image_path = "https://farm5.staticflickr.com/4007/4322154488_997e69e4cf_z.jpg"
27
+ orig_im = io.imread(im_path)
28
+ orig_im_size = orig_im.shape[0:2]
29
+ image = preprocess_image(orig_im, model_input_size).to(device)
30
+
31
+ # inference
32
+ result=model(image)
33
+
34
+ # post process
35
+ result_image = postprocess_image(result[0][0], orig_im_size)
36
+
37
+ # save result
38
+ pil_im = Image.fromarray(result_image)
39
+ no_bg_image = Image.new("RGBA", pil_im.size, (0,0,0,0))
40
+ orig_image = Image.open(im_path)
41
+ no_bg_image.paste(orig_image, mask=pil_im)