Coverage for wrapper/hunyuanavatar/encode_data.py: 100%
69 statements
« prev ^ index » next coverage.py v7.15.4, created at 2026-08-09 04:47 +0000
« prev ^ index » next coverage.py v7.15.4, created at 2026-08-09 04:47 +0000
1import torch
2import librosa
3import numpy as np
5from typing import Dict
6from typing import List
7from typing import Tuple
8from typing import Any
10from PIL import Image
11from einops import rearrange
13from transformers import CLIPImageProcessor
14import torchvision.transforms as transforms
15from torchvision.transforms import ToPILImage
18def get_audio_feature(
19 feature_extractor: Any,
20 audio_path: str
21) -> Tuple[torch.Tensor, int]:
22 audio_input, sampling_rate = librosa.load(audio_path, sr=16000)
23 assert sampling_rate == 16000
25 audio_feature_list: List[torch.Tensor] = []
26 window = 750 * 640
27 for i in range(0, len(audio_input), window):
28 audio_feature = feature_extractor(audio_input[i:i + window],
29 sampling_rate=sampling_rate,
30 return_tensors="pt",
31 ).input_features
32 audio_feature_list.append(audio_feature)
34 audio_features = torch.cat(audio_feature_list, dim=-1)
35 return audio_features, len(audio_input) // 640
38class VideoAudioTextLoaderVal():
39 def __init__(
40 self,
41 image_size: int,
42 text_encoder: Any,
43 text_encoder_2: Any,
44 feature_extractor: Any,
45 ) -> None:
46 self.image_size = image_size
47 self.text_encoder = text_encoder # llava_text_encoder
48 self.text_encoder_2 = text_encoder_2 # clip_text_encoder
49 self.feature_extractor = feature_extractor
51 self.llava_transform = transforms.Compose(
52 [
53 transforms.Resize((336, 336), interpolation=transforms.InterpolationMode.BILINEAR),
54 transforms.ToTensor(),
55 transforms.Normalize((0.48145466, 0.4578275, 0.4082107), (0.26862954, 0.26130258, 0.27577711)),
56 ]
57 )
58 self.clip_image_processor = CLIPImageProcessor()
60 self.device = torch.device("cuda")
61 self.weight_dtype = torch.float16
63 @staticmethod
64 def get_text_tokens(
65 text_encoder: Any,
66 description: str,
67 dtype_encode: str = "video"
68 ) -> Tuple[torch.Tensor, torch.Tensor]:
69 text_inputs = text_encoder.text2tokens(description, data_type=dtype_encode)
70 text_ids = text_inputs["input_ids"].squeeze(0)
71 text_mask = text_inputs["attention_mask"].squeeze(0)
72 return text_ids, text_mask
74 def encode_data(
75 self,
76 ref_image: Any,
77 audio_path: str,
78 prompt: str,
79 fps: float,
80 ) -> Dict[str, Any]:
81 prompt = "Authentic, Realistic, Natural, High-quality, Lens-Fixed, " + prompt
83 img_size = self.image_size
84 # ref_image = Image.open(image_path).convert('RGB')
86 # Resize reference image
87 w, h = ref_image.size
88 scale = img_size / min(w, h)
89 new_w = round(w * scale / 64) * 64
90 new_h = round(h * scale / 64) * 64
92 if img_size == 704:
93 img_size_long = 1216
94 if new_w * new_h > img_size * img_size_long:
95 import math
96 scale = math.sqrt(img_size * img_size_long / w / h)
97 new_w = round(w * scale / 64) * 64
98 new_h = round(h * scale / 64) * 64
100 ref_image = ref_image.resize((new_w, new_h), Image.Resampling.LANCZOS)
102 ref_image = np.array(ref_image)
103 ref_image = torch.from_numpy(ref_image)
105 audio_input, audio_len = get_audio_feature(self.feature_extractor, audio_path)
106 audio_prompts = audio_input[0]
108 motion_bucket_id_heads: torch.Tensor = torch.from_numpy(np.array([25] * 4))
109 motion_bucket_id_exps: torch.Tensor = torch.from_numpy(np.array([30] * 4))
110 fps_tensor = torch.from_numpy(np.array(fps))
112 to_pil = ToPILImage()
113 pixel_value_ref = rearrange(ref_image.clone().unsqueeze(0), "b h w c -> b c h w") # (b c h w)
115 pixel_value_ref_llava_list = [self.llava_transform(to_pil(image)) for image in pixel_value_ref]
116 pixel_value_ref_llava: torch.Tensor = torch.stack(pixel_value_ref_llava_list, dim=0)
117 pixel_value_ref_clip = self.clip_image_processor(
118 images=Image.fromarray((pixel_value_ref[0].permute(1, 2, 0)).data.cpu().numpy().astype(np.uint8)),
119 return_tensors="pt"
120 ).pixel_values[0]
121 pixel_value_ref_clip = pixel_value_ref_clip.unsqueeze(0)
123 # Encode text prompts
124 text_ids, text_mask = self.get_text_tokens(self.text_encoder, prompt)
125 text_ids_2, text_mask_2 = self.get_text_tokens(self.text_encoder_2, prompt)
127 # Output
128 return {
129 "text_prompt": prompt,
130 "pixel_value_ref": pixel_value_ref.to(dtype=torch.float16), # for vae (1, 3, h, w)
131 "pixel_value_ref_llava": pixel_value_ref_llava.to(dtype=torch.float16), # for llava (1, 3, 336, 336)
132 # for clip_image_encoder (1, 3, 244, 244)
133 "pixel_value_ref_clip": pixel_value_ref_clip.to(dtype=torch.float16),
134 "audio_prompts": audio_prompts.to(dtype=torch.float16),
135 "motion_bucket_id_heads": motion_bucket_id_heads.to(dtype=text_ids.dtype),
136 "motion_bucket_id_exps": motion_bucket_id_exps.to(dtype=text_ids.dtype),
137 "fps": fps_tensor.to(dtype=torch.float16),
138 "text_ids": text_ids.clone(), # for llava_text_encoder
139 "text_mask": text_mask.clone(), # for llava_text_encoder
140 "text_ids_2": text_ids_2.clone(), # for clip_text_encoder
141 "text_mask_2": text_mask_2.clone(), # for clip_text_encoder
142 "audio_len": audio_len,
143 "audio_path": audio_path,
144 }