1import torch
2from torchvision.transforms import functional as F
3
4image = Image.open("./retriever_rgb.png").convert("RGB")
5image = F.to_tensor(image).unsqueeze(0).to("cuda").half()
6
7trimap = Image.open("./retriever_trimap.png").convert("L")
8trimap = F.to_tensor(trimap).unsqueeze(0).to("cuda").half()
9
10input = {"image": image, "trimap": trimap}
11
12model = torch.jit.load("./vitmatte_b_dis.pt").to("cuda")
13alpha = model(input)
14
15output = F.to_pil_image(predictions)
16output.save("./predicted.png")