(self, pipe: WanVideoPipeline, inputs_shared, inputs_posi, inputs_nega)
| 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 |
no test coverage detected