V1.5 增加批量处理 + 输出MASK
This commit is contained in:
+37
-23
@@ -60,34 +60,48 @@ class BRIA_RMBG_Zho:
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
RETURN_NAMES = ("image",)
|
||||
RETURN_TYPES = ("IMAGE", "MASK", )
|
||||
RETURN_NAMES = ("image", "mask", )
|
||||
FUNCTION = "remove_background"
|
||||
CATEGORY = "🧹BRIA RMBG"
|
||||
|
||||
def remove_background(self, rmbgmodel, image):
|
||||
orig_image = tensor2pil(image)
|
||||
w,h = orig_image.size
|
||||
image = resize_image(orig_image)
|
||||
im_np = np.array(image)
|
||||
im_tensor = torch.tensor(im_np, dtype=torch.float32).permute(2,0,1)
|
||||
im_tensor = torch.unsqueeze(im_tensor,0)
|
||||
im_tensor = torch.divide(im_tensor,255.0)
|
||||
im_tensor = normalize(im_tensor,[0.5,0.5,0.5],[1.0,1.0,1.0])
|
||||
if torch.cuda.is_available():
|
||||
im_tensor=im_tensor.cuda()
|
||||
processed_images = []
|
||||
processed_masks = []
|
||||
|
||||
result=rmbgmodel(im_tensor)
|
||||
result = torch.squeeze(F.interpolate(result[0][0], size=(h,w), mode='bilinear') ,0)
|
||||
ma = torch.max(result)
|
||||
mi = torch.min(result)
|
||||
result = (result-mi)/(ma-mi)
|
||||
im_array = (result*255).cpu().data.numpy().astype(np.uint8)
|
||||
pil_im = Image.fromarray(np.squeeze(im_array))
|
||||
new_im = Image.new("RGBA", pil_im.size, (0,0,0,0))
|
||||
new_im.paste(orig_image, mask=pil_im)
|
||||
new_im = pil2tensor(new_im)
|
||||
return (new_im,)
|
||||
for image in image:
|
||||
orig_image = tensor2pil(image)
|
||||
w,h = orig_image.size
|
||||
image = resize_image(orig_image)
|
||||
im_np = np.array(image)
|
||||
im_tensor = torch.tensor(im_np, dtype=torch.float32).permute(2,0,1)
|
||||
im_tensor = torch.unsqueeze(im_tensor,0)
|
||||
im_tensor = torch.divide(im_tensor,255.0)
|
||||
im_tensor = normalize(im_tensor,[0.5,0.5,0.5],[1.0,1.0,1.0])
|
||||
if torch.cuda.is_available():
|
||||
im_tensor=im_tensor.cuda()
|
||||
|
||||
result=rmbgmodel(im_tensor)
|
||||
result = torch.squeeze(F.interpolate(result[0][0], size=(h,w), mode='bilinear') ,0)
|
||||
ma = torch.max(result)
|
||||
mi = torch.min(result)
|
||||
result = (result-mi)/(ma-mi)
|
||||
im_array = (result*255).cpu().data.numpy().astype(np.uint8)
|
||||
pil_im = Image.fromarray(np.squeeze(im_array))
|
||||
new_im = Image.new("RGBA", pil_im.size, (0,0,0,0))
|
||||
new_im.paste(orig_image, mask=pil_im)
|
||||
|
||||
new_im_tensor = pil2tensor(new_im) # 将PIL图像转换为Tensor
|
||||
pil_im_tensor = pil2tensor(pil_im) # 同上
|
||||
|
||||
processed_images.append(new_im_tensor)
|
||||
processed_masks.append(pil_im_tensor)
|
||||
|
||||
new_ims = torch.cat(processed_images, dim=0)
|
||||
new_masks = torch.cat(processed_masks, dim=0)
|
||||
|
||||
return new_ims, new_masks
|
||||
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"BRIA_RMBG_ModelLoader_Zho": BRIA_RMBG_ModelLoader_Zho,
|
||||
|
||||
Reference in New Issue
Block a user