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
« 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"""
6import sys
7import numpy as np
9from unittest.mock import patch
10from unittest.mock import MagicMock
11from tests.torch_mock import TorchMock
12from PIL import Image
14mock_torch = TorchMock()
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())
28sys.path.append("wrapper")
29sys.path.append("wrapper/hunyuanavatar")
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"]
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
44 mock_audio_input = np.zeros(16000 * 2, dtype=np.float32) # 2s @ 16kHz
46 with patch.object(_encode_data_module, 'librosa') as mock_librosa, \
47 patch.object(_encode_data_module, 'torch') as mock_t:
49 mock_librosa.load.return_value = (mock_audio_input, 16000)
51 fake_tensor = MagicMock()
52 mock_t.cat.return_value = fake_tensor
54 result_features, result_len = get_audio_feature(mock_feature_extractor, "/tmp/test.wav")
56 assert result_features is fake_tensor
57 assert result_len == len(mock_audio_input) // 640
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()
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'):
70 mock_t.device.return_value = MagicMock()
71 mock_t.float16 = mock_torch.float16
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 )
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
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 }
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
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 }
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 )
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()
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}
142 mock_text_encoder.text2tokens.side_effect = _fake_text2tokens
143 mock_text_encoder_2.text2tokens.side_effect = _fake_text2tokens
145 ref_image = Image.new("RGB", (120, 80))
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]))
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:
162 mock_t.float16 = mock_torch.float16
163 mock_t.from_numpy.side_effect = lambda x: MagicMock()
164 mock_t.device.return_value = MagicMock()
166 fake_audio_features = MagicMock()
167 fake_audio_features.__getitem__ = MagicMock(return_value=MagicMock())
168 mock_get_audio.return_value = (fake_audio_features, 10)
170 # Make rearrange return our controlled fake_pixel_values
171 mock_rearrange.return_value = fake_pixel_values
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 )
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 )
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"
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()
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}
208 mock_text_encoder.text2tokens.side_effect = _fake_text2tokens
209 mock_text_encoder_2.text2tokens.side_effect = _fake_text2tokens
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))
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]))
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:
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
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 )
251 assert result is not None
252 assert result["audio_len"] == 10