MCPcopy Create free account
hub / github.com/MeiGen-AI/MultiTalk / prepare_source

Method prepare_source

wan/vace.py:212–278  ·  view source on GitHub ↗
(self, src_video, src_mask, src_ref_images, num_frames,
                       image_size, device)

Source from the content-addressed store, hash-verified

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),

Callers

nothing calls this directly

Calls 4

set_areaMethod · 0.80
set_seq_lenMethod · 0.80
load_video_pairMethod · 0.80
load_videoMethod · 0.80

Tested by

no test coverage detected