56 KiB
56 KiB
In [ ]:
# Imports
from PIL import Image
import torch
from torchvision import transforms
from IPython.display import display
import sys
sys.path.insert(0, "../")
from models.birefnet import BiRefNet
# Load Model
# Option 2 and Option 3 is better for local running -- we can modify codes locally.
# # # Option 1: loading BiRefNet with weights:
# from transformers import AutoModelForImageSegmentation
# birefnet = AutoModelForImageSegmentation.from_pretrained('zhengpeng7/BiRefNet', trust_remote_code=True)
# Option-2: loading weights with BiReNet codes:
birefnet = BiRefNet.from_pretrained(
[
'zhengpeng7/BiRefNet',
'zhengpeng7/BiRefNet-portrait',
'zhengpeng7/BiRefNet-legacy', 'zhengpeng7/BiRefNet-DIS5K-TR_TEs', 'zhengpeng7/BiRefNet-DIS5K', 'zhengpeng7/BiRefNet-HRSOD', 'zhengpeng7/BiRefNet-COD',
'zhengpeng7/BiRefNet_lite', # Modify the `bb` in `config.py` to `swin_v1_tiny`.
][0]
)
# # Option-3: Loading model and weights from local disk:
# from utils import check_state_dict
# birefnet = BiRefNet(bb_pretrained=False)
# state_dict = torch.load('../BiRefNet-general-epoch_244.pth', map_location='cpu')
# state_dict = check_state_dict(state_dict)
# birefnet.load_state_dict(state_dict)
device = 'cuda' if torch.cuda.is_available() else 'cpu'
torch.set_float32_matmul_precision(['high', 'highest'][0])
birefnet.to(device)
birefnet.eval()
print('BiRefNet is ready to use.')
# Input Data
transform_image = transforms.Compose([
transforms.Resize((1024, 1024)),
transforms.ToTensor(),
transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])
])In [ ]:
import os
from glob import glob
from image_proc import refine_foreground
src_dir = '../images_todo'
image_paths = glob(os.path.join(src_dir, '*'))
dst_dir = '../predictions'
os.makedirs(dst_dir, exist_ok=True)
for image_path in image_paths:
print('Processing {} ...'.format(image_path))
image = Image.open(image_path)
input_images = transform_image(image).unsqueeze(0).to(device)
# Prediction
with torch.no_grad():
preds = birefnet(input_images)[-1].sigmoid().cpu()
pred = preds[0].squeeze()
# Show Results
pred_pil = transforms.ToPILImage()(pred)
pred_pil.resize(image.size).save(image_path.replace(src_dir, dst_dir))
# Visualize the last sample:
# Scale proportionally with max length to 1024 for faster showing
scale_ratio = 1024 / max(image.size)
scaled_size = (int(image.size[0] * scale_ratio), int(image.size[1] * scale_ratio))
image_masked = refine_foreground(image, pred_pil)
image_masked.putalpha(pred_pil.resize(image.size))
display(image.resize(scaled_size))
display(pred_pil.resize(scaled_size))
display(image_masked.resize(scaled_size))In [ ]: