| 31 | } |
| 32 | |
| 33 | def load_data(args): |
| 34 | |
| 35 | transform = transforms.Compose([ |
| 36 | transforms.Resize(size=(224, 224), antialias=True), |
| 37 | # transforms.ToTensor(), |
| 38 | transforms.Normalize(mean, std) |
| 39 | ]) |
| 40 | to_tensor = transforms.ToTensor() |
| 41 | |
| 42 | data_path = 'example_data/' |
| 43 | bg_img = to_tensor(Image.open(data_path + 'bg.png').convert('RGB')).unsqueeze(0) # 1, C, H, W |
| 44 | img_list = [] |
| 45 | start_idx = 0 |
| 46 | for i in range(start_idx, start_idx + args.num_frames * args.stride, args.stride): |
| 47 | img_list.append(to_tensor(Image.open(data_path + str(i) + '.png').convert('RGB'))) |
| 48 | img_list = torch.stack(img_list, dim=0) # T, C, H, W |
| 49 | |
| 50 | img_list = img_list - bg_img + offset |
| 51 | img_list = torch.clamp(img_list, 0.0, 1.0) |
| 52 | img_list_transformed = transform(img_list) # T, C, H, W |
| 53 | |
| 54 | return img_list, img_list_transformed |
| 55 | |
| 56 | def load_model_from_multi_clip(ckpt, model): |
| 57 | |