(self, src_video, src_mask, src_ref_images, num_frames,
image_size, device)
| 210 | return [torch.cat([zz, mm], dim=0) for zz, mm in zip(z, m)] |
| 211 | |
| 212 | def prepare_source(self, src_video, src_mask, src_ref_images, num_frames, |
| 213 | image_size, device): |
| 214 | area = image_size[0] * image_size[1] |
| 215 | self.vid_proc.set_area(area) |
| 216 | if area == 720 * 1280: |
| 217 | self.vid_proc.set_seq_len(75600) |
| 218 | elif area == 480 * 832: |
| 219 | self.vid_proc.set_seq_len(32760) |
| 220 | else: |
| 221 | raise NotImplementedError( |
| 222 | f'image_size {image_size} is not supported') |
| 223 | |
| 224 | image_size = (image_size[1], image_size[0]) |
| 225 | image_sizes = [] |
| 226 | for i, (sub_src_video, |
| 227 | sub_src_mask) in enumerate(zip(src_video, src_mask)): |
| 228 | if sub_src_mask is not None and sub_src_video is not None: |
| 229 | src_video[i], src_mask[ |
| 230 | i], _, _, _ = self.vid_proc.load_video_pair( |
| 231 | sub_src_video, sub_src_mask) |
| 232 | src_video[i] = src_video[i].to(device) |
| 233 | src_mask[i] = src_mask[i].to(device) |
| 234 | src_mask[i] = torch.clamp( |
| 235 | (src_mask[i][:1, :, :, :] + 1) / 2, min=0, max=1) |
| 236 | image_sizes.append(src_video[i].shape[2:]) |
| 237 | elif sub_src_video is None: |
| 238 | src_video[i] = torch.zeros( |
| 239 | (3, num_frames, image_size[0], image_size[1]), |
| 240 | device=device) |
| 241 | src_mask[i] = torch.ones_like(src_video[i], device=device) |
| 242 | image_sizes.append(image_size) |
| 243 | else: |
| 244 | src_video[i], _, _, _ = self.vid_proc.load_video(sub_src_video) |
| 245 | src_video[i] = src_video[i].to(device) |
| 246 | src_mask[i] = torch.ones_like(src_video[i], device=device) |
| 247 | image_sizes.append(src_video[i].shape[2:]) |
| 248 | |
| 249 | for i, ref_images in enumerate(src_ref_images): |
| 250 | if ref_images is not None: |
| 251 | image_size = image_sizes[i] |
| 252 | for j, ref_img in enumerate(ref_images): |
| 253 | if ref_img is not None: |
| 254 | ref_img = Image.open(ref_img).convert("RGB") |
| 255 | ref_img = TF.to_tensor(ref_img).sub_(0.5).div_( |
| 256 | 0.5).unsqueeze(1) |
| 257 | if ref_img.shape[-2:] != image_size: |
| 258 | canvas_height, canvas_width = image_size |
| 259 | ref_height, ref_width = ref_img.shape[-2:] |
| 260 | white_canvas = torch.ones( |
| 261 | (3, 1, canvas_height, canvas_width), |
| 262 | device=device) # [-1, 1] |
| 263 | scale = min(canvas_height / ref_height, |
| 264 | canvas_width / ref_width) |
| 265 | new_height = int(ref_height * scale) |
| 266 | new_width = int(ref_width * scale) |
| 267 | resized_image = F.interpolate( |
| 268 | ref_img.squeeze(1).unsqueeze(0), |
| 269 | size=(new_height, new_width), |
nothing calls this directly
no test coverage detected