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

1import torch 

2import librosa 

3import numpy as np 

4 

5from typing import Dict 

6from typing import List 

7from typing import Tuple 

8from typing import Any 

9 

10from PIL import Image 

11from einops import rearrange 

12 

13from transformers import CLIPImageProcessor 

14import torchvision.transforms as transforms 

15from torchvision.transforms import ToPILImage 

16 

17 

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 

24 

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) 

33 

34 audio_features = torch.cat(audio_feature_list, dim=-1) 

35 return audio_features, len(audio_input) // 640 

36 

37 

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 

50 

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

59 

60 self.device = torch.device("cuda") 

61 self.weight_dtype = torch.float16 

62 

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 

73 

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 

82 

83 img_size = self.image_size 

84 # ref_image = Image.open(image_path).convert('RGB') 

85 

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 

91 

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 

99 

100 ref_image = ref_image.resize((new_w, new_h), Image.Resampling.LANCZOS) 

101 

102 ref_image = np.array(ref_image) 

103 ref_image = torch.from_numpy(ref_image) 

104 

105 audio_input, audio_len = get_audio_feature(self.feature_extractor, audio_path) 

106 audio_prompts = audio_input[0] 

107 

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

111 

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) 

114 

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) 

122 

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) 

126 

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 }