MCPcopy Create free account
hub / github.com/modelscope/DiffSynth-Studio / process

Method process

diffsynth/pipelines/wan_video_new.py:948–1004  ·  view source on GitHub ↗
(self, pipe: WanVideoPipeline, inputs_shared, inputs_posi, inputs_nega)

Source from the content-addressed store, hash-verified

946 )
947
948 def process(self, pipe: WanVideoPipeline, inputs_shared, inputs_posi, inputs_nega):
949 if inputs_shared.get("vap_video") is None:
950 return inputs_shared, inputs_posi, inputs_nega
951 else:
952 # 1. encode vap prompt
953 pipe.load_models_to_device(["text_encoder"])
954 vap_prompt, negative_vap_prompt = inputs_posi.get("vap_prompt", ""), inputs_nega.get("negative_vap_prompt", "")
955 vap_prompt_emb = pipe.prompter.encode_prompt(vap_prompt, positive=inputs_posi.get('positive',None), device=pipe.device)
956 negative_vap_prompt_emb = pipe.prompter.encode_prompt(negative_vap_prompt, positive=inputs_nega.get('positive',None), device=pipe.device)
957 inputs_posi.update({"context_vap":vap_prompt_emb})
958 inputs_nega.update({"context_vap":negative_vap_prompt_emb})
959 # 2. prepare vap image clip embedding
960 pipe.load_models_to_device(["vae", "image_encoder"])
961 vap_video, end_image = inputs_shared.get("vap_video"), inputs_shared.get("end_image")
962
963 num_frames, height, width, mot_num = inputs_shared.get("num_frames"),inputs_shared.get("height"), inputs_shared.get("width"), inputs_shared.get("mot_num",1)
964
965 image_vap = pipe.preprocess_image(vap_video[0].resize((width, height))).to(pipe.device)
966
967 vap_clip_context = pipe.image_encoder.encode_image([image_vap])
968 if end_image is not None:
969 vap_end_image = pipe.preprocess_image(vap_video[-1].resize((width, height))).to(pipe.device)
970 if pipe.dit.has_image_pos_emb:
971 vap_clip_context = torch.concat([vap_clip_context, pipe.image_encoder.encode_image([vap_end_image])], dim=1)
972 vap_clip_context = vap_clip_context.to(dtype=pipe.torch_dtype, device=pipe.device)
973 inputs_shared.update({"vap_clip_feature":vap_clip_context})
974
975 # 3. prepare vap latents
976 msk = torch.ones(1, num_frames, height//8, width//8, device=pipe.device)
977 msk[:, 1:] = 0
978 if end_image is not None:
979 msk[:, -1:] = 1
980 last_image_vap = pipe.preprocess_image(vap_video[-1].resize((width, height))).to(pipe.device)
981 vae_input = torch.concat([image_vap.transpose(0,1), torch.zeros(3, num_frames-2, height, width).to(image_vap.device), last_image_vap.transpose(0,1)],dim=1)
982 else:
983 vae_input = torch.concat([image_vap.transpose(0, 1), torch.zeros(3, num_frames-1, height, width).to(image_vap.device)], dim=1)
984
985 msk = torch.concat([torch.repeat_interleave(msk[:, 0:1], repeats=4, dim=1), msk[:, 1:]], dim=1)
986 msk = msk.view(1, msk.shape[1] // 4, 4, height//8, width//8)
987 msk = msk.transpose(1, 2)[0]
988
989 tiled,tile_size,tile_stride = inputs_shared.get("tiled"), inputs_shared.get("tile_size"), inputs_shared.get("tile_stride")
990
991 y = pipe.vae.encode([vae_input.to(dtype=pipe.torch_dtype, device=pipe.device)], device=pipe.device, tiled=tiled, tile_size=tile_size, tile_stride=tile_stride)[0]
992 y = y.to(dtype=pipe.torch_dtype, device=pipe.device)
993 y = torch.concat([msk, y])
994 y = y.unsqueeze(0)
995 y = y.to(dtype=pipe.torch_dtype, device=pipe.device)
996
997 vap_video = pipe.preprocess_video(vap_video)
998 vap_latent = pipe.vae.encode(vap_video, device=pipe.device, tiled=tiled, tile_size=tile_size, tile_stride=tile_stride).to(dtype=pipe.torch_dtype, device=pipe.device)
999
1000 vap_latent = torch.concat([vap_latent,y], dim=1).to(dtype=pipe.torch_dtype, device=pipe.device)
1001 inputs_shared.update({"vap_hidden_state":vap_latent})
1002 pipe.load_models_to_device([])
1003
1004 return inputs_shared, inputs_posi, inputs_nega
1005

Callers 1

Calls 8

preprocess_videoMethod · 0.80
load_models_to_deviceMethod · 0.45
encode_promptMethod · 0.45
updateMethod · 0.45
toMethod · 0.45
preprocess_imageMethod · 0.45
encode_imageMethod · 0.45
encodeMethod · 0.45

Tested by

no test coverage detected