arabago96/Ai_spatial_modeling
1
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 