Views
No views yet
git clone https://github.com/VicFonch/Multi-Input-Resshift-Diffusion-VFI.git1conda create -n multi-input-resshift python=3.12
2conda activate multi-input-resshift
3pip install -r requirements.txt1import os
2from PIL import Image
3import numpy as np
4import matplotlib.pyplot as plt
5
6from torchvision.transforms import Compose, ToTensor, Resize, Normalize
7from utils.utils import denorm
8from model.hub import MultiInputResShiftHub
9
10model = MultiInputResShiftHub.from_pretrained("vfontech/Multiple-Input-Resshift-VFI").cuda()
11model.eval()
12
13img0_path = r"_data\example_images\frame1.png"
14img2_path = r"_data\example_images\frame3.png"
15
16mean = std = [0.5]*3
17transforms = Compose([
18 Resize((256, 448)),
19 ToTensor(),
20 Normalize(mean=mean, std=std),
21])
22
23img0 = transforms(Image.open(img0_path).convert("RGB")).unsqueeze(0).cuda()
24img2 = transforms(Image.open(img2_path).convert("RGB")).unsqueeze(0).cuda()
25tau = 0.5
26
27img1 = model.reverse_process([img0, img2], tau)
28
29plt.figure(figsize=(10, 5))
30plt.subplot(1, 3, 1)
31plt.imshow(denorm(img0, mean=mean, std=std).squeeze().permute(1, 2, 0).cpu().numpy())
32plt.subplot(1, 3, 2)
33plt.imshow(denorm(img1, mean=mean, std=std).squeeze().permute(1, 2, 0).cpu().numpy())
34plt.subplot(1, 3, 3)
35plt.imshow(denorm(img2, mean=mean, std=std).squeeze().permute(1, 2, 0).cpu().numpy())
36plt.show()