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

1#!/usr/bin/env python3 

2""" 

3Tests for wrapper/hunyuanavatar/config.py and encode_data.py 

4""" 

5 

6import sys 

7import argparse 

8import pytest 

9 

10from unittest.mock import patch, MagicMock 

11from tests.torch_mock import TorchMock 

12 

13mock_torch = TorchMock() 

14 

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"] 

22 

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

36 

37sys.path.append("wrapper") 

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

39 

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 

49 

50 

51# ── as_tuple ────────────────────────────────────────────────────────────────── 

52 

53def test_as_tuple_with_list() -> None: 

54 assert as_tuple([1, 2, 3]) == (1, 2, 3) 

55 

56 

57def test_as_tuple_with_tuple() -> None: 

58 assert as_tuple((4, 5)) == (4, 5) 

59 

60 

61def test_as_tuple_with_int() -> None: 

62 assert as_tuple(7) == (7,) 

63 

64 

65def test_as_tuple_with_float() -> None: 

66 assert as_tuple(3.14) == (3.14,) 

67 

68 

69def test_as_tuple_with_string() -> None: 

70 assert as_tuple("hello") == ("hello",) 

71 

72 

73def test_as_tuple_with_none() -> None: 

74 assert as_tuple(None) == (None,) 

75 

76 

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

82 

83 

84# ── parse_args ──────────────────────────────────────────────────────────────── 

85 

86def test_parse_args_returns_namespace() -> None: 

87 args = parse_args() 

88 assert isinstance(args, argparse.Namespace) 

89 

90 

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" 

97 

98 

99def test_parse_args_with_namespace() -> None: 

100 ns = argparse.Namespace() 

101 args = parse_args(namespace=ns) 

102 assert isinstance(args, argparse.Namespace) 

103 

104 

105# ── sanity_check_args ───────────────────────────────────────────────────────── 

106 

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 

111 

112 

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) 

117 

118 

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) 

123 

124 

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 

129 

130 

131# ── add_*_args helpers ──────────────────────────────────────────────────────── 

132 

133def test_add_network_args() -> None: 

134 parser = argparse.ArgumentParser() 

135 result = add_network_args(parser) 

136 assert result is parser 

137 

138 

139def test_add_extra_models_args() -> None: 

140 parser = argparse.ArgumentParser() 

141 result = add_extra_models_args(parser) 

142 assert result is parser 

143 

144 

145def test_add_denoise_schedule_args() -> None: 

146 parser = argparse.ArgumentParser() 

147 result = add_denoise_schedule_args(parser) 

148 assert result is parser 

149 

150 

151def test_add_evaluation_args() -> None: 

152 parser = argparse.ArgumentParser() 

153 result = add_evaluation_args(parser) 

154 assert result is parser 

155 

156 

157def test_add_extra_args() -> None: 

158 parser = argparse.ArgumentParser() 

159 result = add_extra_args(parser) 

160 assert result is parser