From d60b718eb1e64f6d0849b1cdf94688dd3c576500 Mon Sep 17 00:00:00 2001 From: ZHO-ZHO-ZHO <140084057+ZHO-ZHO-ZHO@users.noreply.github.com> Date: Wed, 7 Feb 2024 20:34:25 +0800 Subject: [PATCH] =?UTF-8?q?V1.5=20=E5=A2=9E=E5=8A=A0=E6=89=B9=E9=87=8F?= =?UTF-8?q?=E5=A4=84=E7=90=86=20+=20=E8=BE=93=E5=87=BAMASK?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- BRIA_RMBG.py | 60 ++++++++++++++++++++++++++++++++-------------------- 1 file changed, 37 insertions(+), 23 deletions(-) diff --git a/BRIA_RMBG.py b/BRIA_RMBG.py index d7cccba..c1e88cb 100644 --- a/BRIA_RMBG.py +++ b/BRIA_RMBG.py @@ -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,