CoolFace
Apppublic

arabago96/Ai_spatial_modeling

sourceHugging Faceapache-2.0updated 2mo agoView on Hugging Face
1likes
utils_birefnet.py46 linesDownload Raw Back to root
1from typing import *2from transformers import AutoModelForImageSegmentation3import torch4from torchvision import transforms5from PIL import Image6 7class BiRefNet:8    def __init__(self, model_name: str = "ZhengPeng7/BiRefNet"):9        self.model = AutoModelForImageSegmentation.from_pretrained(10            model_name, trust_remote_code=True11        )12        self.model.eval()13        self.transform_image = transforms.Compose(14            [15                transforms.Resize((1024, 1024)),16                transforms.ToTensor(),17                transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]),18            ]19        )20    21    def to(self, device: str):22        self.model.to(device)23 24    def cuda(self):25        self.model.cuda()26 27    def cpu(self):28        self.model.cpu()29        30    def __call__(self, image: Image.Image) -> Image.Image:31        image_size = image.size32        # Always convert to RGB for the transform (handles RGBA, L, LA, CMYK, P, etc.)33        rgb_image = image.convert('RGB')34 35        input_images = self.transform_image(rgb_image).unsqueeze(0).to("cuda")36        # Prediction37        with torch.no_grad():38            preds = self.model(input_images)[-1].sigmoid().cpu()39        pred = preds[0].squeeze()40        pred_pil = transforms.ToPILImage()(pred)41        mask = pred_pil.resize(image_size)42        # Convert to RGBA so putalpha works regardless of the original mode43        rgba_image = rgb_image.convert('RGBA')44        rgba_image.putalpha(mask)45        return rgba_image46