Coverage for tests/test_hunyuanavatar_encode_data.py: 100%

125 statements  

« prev     ^ index     » next       coverage.py v7.15.4, created at 2026-08-09 04:47 +0000

1#!/usr/bin/env python3 

2""" 

3Tests for wrapper/hunyuanavatar/encode_data.py 

4""" 

5 

6import sys 

7import numpy as np 

8 

9from unittest.mock import patch 

10from unittest.mock import MagicMock 

11from tests.torch_mock import TorchMock 

12from PIL import Image 

13 

14mock_torch = TorchMock() 

15 

16mock_modules = { 

17 "nvidia_smi": MagicMock(), 

18 "torch": mock_torch, 

19 "torchvision": MagicMock(), 

20 "torchvision.transforms": MagicMock(), 

21 "torchvision.transforms.functional": MagicMock(), 

22 "transformers": MagicMock(), 

23 "einops": MagicMock(), 

24 "librosa": MagicMock(), 

25} 

26mock_modules.update(mock_torch.get_sub_modules()) 

27 

28sys.path.append("wrapper") 

29sys.path.append("wrapper/hunyuanavatar") 

30 

31with patch.dict(sys.modules, mock_modules): 

32 from hunyuanavatar.encode_data import get_audio_feature 

33 from hunyuanavatar.encode_data import VideoAudioTextLoaderVal 

34 _encode_data_module = sys.modules["hunyuanavatar.encode_data"] 

35 

36 

37def test_get_audio_feature_basic() -> None: 

38 """Test get_audio_feature returns expected tuple structure.""" 

39 mock_feature_extractor = MagicMock() 

40 fake_input_features = MagicMock() 

41 fake_input_features.__len__ = MagicMock(return_value=1) 

42 mock_feature_extractor.return_value.input_features = fake_input_features 

43 

44 mock_audio_input = np.zeros(16000 * 2, dtype=np.float32) # 2s @ 16kHz 

45 

46 with patch.object(_encode_data_module, 'librosa') as mock_librosa, \ 

47 patch.object(_encode_data_module, 'torch') as mock_t: 

48 

49 mock_librosa.load.return_value = (mock_audio_input, 16000) 

50 

51 fake_tensor = MagicMock() 

52 mock_t.cat.return_value = fake_tensor 

53 

54 result_features, result_len = get_audio_feature(mock_feature_extractor, "/tmp/test.wav") 

55 

56 assert result_features is fake_tensor 

57 assert result_len == len(mock_audio_input) // 640 

58 

59 

60def test_video_audio_text_loader_val_init() -> None: 

61 """Test VideoAudioTextLoaderVal initializes correctly.""" 

62 mock_text_encoder = MagicMock() 

63 mock_text_encoder_2 = MagicMock() 

64 mock_feature_extractor = MagicMock() 

65 

66 with patch.object(_encode_data_module, 'torch') as mock_t, \ 

67 patch.object(_encode_data_module, 'transforms'), \ 

68 patch.object(_encode_data_module, 'CLIPImageProcessor'): 

69 

70 mock_t.device.return_value = MagicMock() 

71 mock_t.float16 = mock_torch.float16 

72 

73 loader = VideoAudioTextLoaderVal( 

74 image_size=704, 

75 text_encoder=mock_text_encoder, 

76 text_encoder_2=mock_text_encoder_2, 

77 feature_extractor=mock_feature_extractor, 

78 ) 

79 

80 assert loader.image_size == 704 

81 assert loader.text_encoder is mock_text_encoder 

82 assert loader.text_encoder_2 is mock_text_encoder_2 

83 assert loader.feature_extractor is mock_feature_extractor 

84 

85 

86def test_video_audio_text_loader_val_get_text_tokens() -> None: 

87 """Test VideoAudioTextLoaderVal.get_text_tokens static method.""" 

88 mock_text_encoder = MagicMock() 

89 fake_input_ids = MagicMock() 

90 fake_attn_mask = MagicMock() 

91 fake_input_ids.squeeze.return_value = fake_input_ids 

92 fake_attn_mask.squeeze.return_value = fake_attn_mask 

93 mock_text_encoder.text2tokens.return_value = { 

94 "input_ids": fake_input_ids, 

95 "attention_mask": fake_attn_mask, 

96 } 

97 

98 text_ids, text_mask = VideoAudioTextLoaderVal.get_text_tokens( 

99 mock_text_encoder, "test description" 

100 ) 

101 mock_text_encoder.text2tokens.assert_called_once_with( 

102 "test description", data_type="video" 

103 ) 

104 assert text_ids is fake_input_ids 

105 assert text_mask is fake_attn_mask 

106 

107 

108def test_video_audio_text_loader_val_get_text_tokens_image_dtype() -> None: 

109 """Test get_text_tokens with image dtype_encode.""" 

110 mock_text_encoder = MagicMock() 

111 fake_input_ids = MagicMock() 

112 fake_attn_mask = MagicMock() 

113 fake_input_ids.squeeze.return_value = fake_input_ids 

114 fake_attn_mask.squeeze.return_value = fake_attn_mask 

115 mock_text_encoder.text2tokens.return_value = { 

116 "input_ids": fake_input_ids, 

117 "attention_mask": fake_attn_mask, 

118 } 

119 

120 text_ids, text_mask = VideoAudioTextLoaderVal.get_text_tokens( 

121 mock_text_encoder, "a photo", dtype_encode="image" 

122 ) 

123 mock_text_encoder.text2tokens.assert_called_once_with( 

124 "a photo", data_type="image" 

125 ) 

126 

127 

128def test_video_audio_text_loader_encode_data_returns_expected_structure() -> None: 

129 """Test encode_data returns the expected dict structure with correct keys and values.""" 

130 mock_text_encoder = MagicMock() 

131 mock_text_encoder_2 = MagicMock() 

132 mock_feature_extractor = MagicMock() 

133 

134 # Prepare fake returns for text tokens 

135 def _fake_text2tokens(text: str, data_type: str = "video") -> dict: 

136 fake_ids = MagicMock() 

137 fake_ids.squeeze.return_value = MagicMock() 

138 fake_mask = MagicMock() 

139 fake_mask.squeeze.return_value = MagicMock() 

140 return {"input_ids": fake_ids, "attention_mask": fake_mask} 

141 

142 mock_text_encoder.text2tokens.side_effect = _fake_text2tokens 

143 mock_text_encoder_2.text2tokens.side_effect = _fake_text2tokens 

144 

145 ref_image = Image.new("RGB", (120, 80)) 

146 

147 # Set up a fake pixel_values tensor that supports the operations in encode_data 

148 fake_np_image = np.zeros((64, 64, 3), dtype=np.uint8) 

149 fake_item = MagicMock() 

150 fake_item.permute.return_value.data.cpu.return_value.numpy.return_value.astype.return_value = fake_np_image 

151 fake_pixel_values = MagicMock() 

152 fake_pixel_values.__getitem__ = MagicMock(return_value=fake_item) 

153 fake_pixel_values.__iter__ = MagicMock(return_value=iter([fake_item])) 

154 

155 with patch.object(_encode_data_module, 'torch') as mock_t, \ 

156 patch.object(_encode_data_module, 'transforms'), \ 

157 patch.object(_encode_data_module, 'CLIPImageProcessor'), \ 

158 patch.object(_encode_data_module, 'rearrange') as mock_rearrange, \ 

159 patch.object(_encode_data_module, 'ToPILImage'), \ 

160 patch.object(_encode_data_module, 'get_audio_feature') as mock_get_audio: 

161 

162 mock_t.float16 = mock_torch.float16 

163 mock_t.from_numpy.side_effect = lambda x: MagicMock() 

164 mock_t.device.return_value = MagicMock() 

165 

166 fake_audio_features = MagicMock() 

167 fake_audio_features.__getitem__ = MagicMock(return_value=MagicMock()) 

168 mock_get_audio.return_value = (fake_audio_features, 10) 

169 

170 # Make rearrange return our controlled fake_pixel_values 

171 mock_rearrange.return_value = fake_pixel_values 

172 

173 # Create loader 

174 loader = VideoAudioTextLoaderVal( 

175 image_size=704, 

176 text_encoder=mock_text_encoder, 

177 text_encoder_2=mock_text_encoder_2, 

178 feature_extractor=mock_feature_extractor, 

179 ) 

180 

181 result = loader.encode_data( 

182 ref_image=ref_image, 

183 audio_path="/tmp/test_audio.wav", 

184 prompt="test prompt", 

185 fps=12.5, 

186 ) 

187 

188 assert result is not None 

189 assert "text_prompt" in result 

190 assert "audio_len" in result 

191 assert result["audio_len"] == 10 

192 assert result["audio_path"] == "/tmp/test_audio.wav" 

193 

194 

195def test_video_audio_text_loader_encode_data_large_image() -> None: 

196 """Test encode_data with a large image that triggers the rescaling branch (lines 95-98).""" 

197 mock_text_encoder = MagicMock() 

198 mock_text_encoder_2 = MagicMock() 

199 mock_feature_extractor = MagicMock() 

200 

201 def _fake_text2tokens(text: str, data_type: str = "video") -> dict: 

202 fake_ids = MagicMock() 

203 fake_ids.squeeze.return_value = MagicMock() 

204 fake_mask = MagicMock() 

205 fake_mask.squeeze.return_value = MagicMock() 

206 return {"input_ids": fake_ids, "attention_mask": fake_mask} 

207 

208 mock_text_encoder.text2tokens.side_effect = _fake_text2tokens 

209 mock_text_encoder_2.text2tokens.side_effect = _fake_text2tokens 

210 

211 # Large image: 1000x2000 triggers the area-cap branch in encode_data. 

212 # At image_size=704: new_w*new_h exceeds the cap of 704*1216 (portrait max area), 

213 # so the code re-scales to fit within that budget. 

214 ref_image = Image.new("RGB", (1000, 2000)) 

215 

216 fake_np_image = np.zeros((64, 64, 3), dtype=np.uint8) 

217 fake_item = MagicMock() 

218 fake_item.permute.return_value.data.cpu.return_value.numpy.return_value.astype.return_value = fake_np_image 

219 fake_pixel_values = MagicMock() 

220 fake_pixel_values.__getitem__ = MagicMock(return_value=fake_item) 

221 fake_pixel_values.__iter__ = MagicMock(return_value=iter([fake_item])) 

222 

223 with patch.object(_encode_data_module, 'torch') as mock_t, \ 

224 patch.object(_encode_data_module, 'transforms'), \ 

225 patch.object(_encode_data_module, 'CLIPImageProcessor'), \ 

226 patch.object(_encode_data_module, 'rearrange') as mock_rearrange, \ 

227 patch.object(_encode_data_module, 'ToPILImage'), \ 

228 patch.object(_encode_data_module, 'get_audio_feature') as mock_get_audio: 

229 

230 mock_t.float16 = mock_torch.float16 

231 mock_t.from_numpy.side_effect = lambda x: MagicMock() 

232 mock_t.device.return_value = MagicMock() 

233 fake_audio_features = MagicMock() 

234 fake_audio_features.__getitem__ = MagicMock(return_value=MagicMock()) 

235 mock_get_audio.return_value = (fake_audio_features, 10) 

236 mock_rearrange.return_value = fake_pixel_values 

237 

238 loader = VideoAudioTextLoaderVal( 

239 image_size=704, 

240 text_encoder=mock_text_encoder, 

241 text_encoder_2=mock_text_encoder_2, 

242 feature_extractor=mock_feature_extractor, 

243 ) 

244 result = loader.encode_data( 

245 ref_image=ref_image, 

246 audio_path="/tmp/test_audio.wav", 

247 prompt="large image test", 

248 fps=12.5, 

249 ) 

250 

251 assert result is not None 

252 assert result["audio_len"] == 10