Coverage for tests/test_hunyuanavatar_config.py: 100%
91 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/config.py and encode_data.py
4"""
6import sys
7import argparse
8import pytest
10from unittest.mock import patch, MagicMock
11from tests.torch_mock import TorchMock
13mock_torch = TorchMock()
15# Build mock modules for hymm_sp.constants used by config.py
16_mock_constants = MagicMock()
17_mock_constants.TEXT_ENCODER_PATH = {"llava-llama-3-8b": "/path/llava", "clipL": "/path/clip"}
18_mock_constants.TOKENIZER_PATH = {"llava-llama-3-8b": "/path/tokenizer", "clipL": "/path/clip_tok"}
19_mock_constants.PROMPT_TEMPLATE = ["li-dit-encode-video"]
20_mock_constants.TEXT_PROJECTION = ["single_refiner"]
21_mock_constants.PRECISIONS = ["fp32", "fp16", "bf16"]
23mock_modules = {
24 'nvidia_smi': MagicMock(),
25 'torch': mock_torch,
26 'torchvision': MagicMock(),
27 'torchvision.transforms': MagicMock(),
28 'torchvision.transforms.functional': MagicMock(),
29 'transformers': MagicMock(),
30 'einops': MagicMock(),
31 'hymm_sp': MagicMock(),
32 'hymm_sp.constants': _mock_constants,
33 'hymm_sp.config': MagicMock(),
34}
35mock_modules.update(mock_torch.get_sub_modules())
37sys.path.append("wrapper")
38sys.path.append("wrapper/hunyuanavatar")
40with patch.dict(sys.modules, mock_modules):
41 from hunyuanavatar.config import as_tuple
42 from hunyuanavatar.config import parse_args
43 from hunyuanavatar.config import sanity_check_args
44 from hunyuanavatar.config import add_extra_args
45 from hunyuanavatar.config import add_network_args
46 from hunyuanavatar.config import add_extra_models_args
47 from hunyuanavatar.config import add_denoise_schedule_args
48 from hunyuanavatar.config import add_evaluation_args
51# ── as_tuple ──────────────────────────────────────────────────────────────────
53def test_as_tuple_with_list() -> None:
54 assert as_tuple([1, 2, 3]) == (1, 2, 3)
57def test_as_tuple_with_tuple() -> None:
58 assert as_tuple((4, 5)) == (4, 5)
61def test_as_tuple_with_int() -> None:
62 assert as_tuple(7) == (7,)
65def test_as_tuple_with_float() -> None:
66 assert as_tuple(3.14) == (3.14,)
69def test_as_tuple_with_string() -> None:
70 assert as_tuple("hello") == ("hello",)
73def test_as_tuple_with_none() -> None:
74 assert as_tuple(None) == (None,)
77def test_as_tuple_with_unknown_type() -> None:
78 class _Obj:
79 pass
80 with pytest.raises(ValueError, match="Unknown type"):
81 as_tuple(_Obj())
84# ── parse_args ────────────────────────────────────────────────────────────────
86def test_parse_args_returns_namespace() -> None:
87 args = parse_args()
88 assert isinstance(args, argparse.Namespace)
91def test_parse_args_defaults() -> None:
92 args = parse_args()
93 assert args.vae == "884-16c-hy0801"
94 assert args.latent_channels == 16
95 assert args.rope_theta == 256
96 assert args.flow_solver == "euler"
99def test_parse_args_with_namespace() -> None:
100 ns = argparse.Namespace()
101 args = parse_args(namespace=ns)
102 assert isinstance(args, argparse.Namespace)
105# ── sanity_check_args ─────────────────────────────────────────────────────────
107def test_sanity_check_args_valid() -> None:
108 args = argparse.Namespace(vae="884-16c-hy0801", latent_channels=None)
109 result = sanity_check_args(args)
110 assert result.latent_channels == 16
113def test_sanity_check_args_latent_channels_mismatch() -> None:
114 args = argparse.Namespace(vae="884-16c-hy0801", latent_channels=8)
115 with pytest.raises(ValueError, match="Latent.*must match VAE"):
116 sanity_check_args(args)
119def test_sanity_check_args_invalid_vae_format() -> None:
120 args = argparse.Namespace(vae="invalid-vae", latent_channels=None)
121 with pytest.raises(ValueError, match="Invalid VAE model"):
122 sanity_check_args(args)
125def test_sanity_check_args_latent_matches_vae() -> None:
126 args = argparse.Namespace(vae="884-16c-hy0801", latent_channels=16)
127 result = sanity_check_args(args)
128 assert result.latent_channels == 16
131# ── add_*_args helpers ────────────────────────────────────────────────────────
133def test_add_network_args() -> None:
134 parser = argparse.ArgumentParser()
135 result = add_network_args(parser)
136 assert result is parser
139def test_add_extra_models_args() -> None:
140 parser = argparse.ArgumentParser()
141 result = add_extra_models_args(parser)
142 assert result is parser
145def test_add_denoise_schedule_args() -> None:
146 parser = argparse.ArgumentParser()
147 result = add_denoise_schedule_args(parser)
148 assert result is parser
151def test_add_evaluation_args() -> None:
152 parser = argparse.ArgumentParser()
153 result = add_evaluation_args(parser)
154 assert result is parser
157def test_add_extra_args() -> None:
158 parser = argparse.ArgumentParser()
159 result = add_extra_args(parser)
160 assert result is parser