Coverage for wrapper/hunyuanavatar/config.py: 100%

106 statements  

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

1import argparse 

2import re 

3import collections.abc 

4 

5from typing import Tuple 

6from typing import Any 

7from typing import Optional 

8 

9from hymm_sp.constants import TEXT_ENCODER_PATH 

10from hymm_sp.constants import TOKENIZER_PATH 

11from hymm_sp.constants import PROMPT_TEMPLATE 

12from hymm_sp.constants import TEXT_PROJECTION 

13from hymm_sp.constants import PRECISIONS 

14 

15 

16def as_tuple(x: Any) -> Tuple: 

17 if isinstance(x, collections.abc.Iterable) and not isinstance(x, str): 

18 return tuple(x) 

19 if x is None or isinstance(x, (int, float, str)): 

20 return (x,) 

21 raise ValueError(f"Unknown type {type(x)}") 

22 

23 

24def parse_args( 

25 namespace: Optional[argparse.Namespace] = None 

26) -> argparse.Namespace: 

27 parser = argparse.ArgumentParser(description="Hunyuan Multimodal training/inference script") 

28 parser = add_extra_args(parser) 

29 # args = parser.parse_args(namespace=namespace) 

30 # (hqiu) accept other arguments from run_httpserver.py 

31 args, unknown_args = parser.parse_known_args(namespace=namespace) 

32 if unknown_args: 

33 print(f"Additional arguments: {unknown_args}") 

34 assert args is not None 

35 args = sanity_check_args(args) 

36 return args 

37 

38 

39def add_extra_args( 

40 parser: argparse.ArgumentParser 

41) -> argparse.ArgumentParser: 

42 parser = add_network_args(parser) 

43 parser = add_extra_models_args(parser) 

44 parser = add_denoise_schedule_args(parser) 

45 parser = add_evaluation_args(parser) 

46 return parser 

47 

48 

49def add_network_args( 

50 parser: argparse.ArgumentParser 

51) -> argparse.ArgumentParser: 

52 group = parser.add_argument_group(title="Network") 

53 group.add_argument("--model", type=str, default="HYVideo-T/2", 

54 help="Model architecture to use. It it also used to determine the experiment directory.") 

55 group.add_argument("--latent-channels", type=str, default=None, 

56 help="Number of latent channels of DiT. If None, it will be determined by `vae`. If provided, " 

57 "it still needs to match the latent channels of the VAE model.") 

58 group.add_argument("--rope-theta", type=int, default=256, help="Theta used in RoPE.") 

59 return parser 

60 

61 

62def add_extra_models_args(parser: argparse.ArgumentParser) -> argparse.ArgumentParser: 

63 group = parser.add_argument_group(title="Extra Models (VAE, Text Encoder, Tokenizer)") 

64 

65 # VAE 

66 group.add_argument("--vae", type=str, default="884-16c-hy0801", help="Name of the VAE model.") 

67 group.add_argument("--vae-precision", type=str, default="fp16", 

68 help="Precision mode for the VAE model.") 

69 group.add_argument("--vae-tiling", action="store_true", default=True, help="Enable tiling for the VAE model.") 

70 group.add_argument("--text-encoder", type=str, default="llava-llama-3-8b", choices=list(TEXT_ENCODER_PATH), 

71 help="Name of the text encoder model.") 

72 group.add_argument("--text-encoder-precision", type=str, default="fp16", choices=PRECISIONS, 

73 help="Precision mode for the text encoder model.") 

74 group.add_argument("--text-states-dim", type=int, default=4096, help="Dimension of the text encoder hidden states.") 

75 group.add_argument("--text-len", type=int, default=256, help="Maximum length of the text input.") 

76 group.add_argument("--tokenizer", type=str, default="llava-llama-3-8b", choices=list(TOKENIZER_PATH), 

77 help="Name of the tokenizer model.") 

78 group.add_argument("--text-encoder-infer-mode", type=str, default="encoder", choices=["encoder", "decoder"], 

79 help="Inference mode for the text encoder model. It should match the text encoder type. T5 and " 

80 "CLIP can only work in 'encoder' mode, while Llava/GLM can work in both modes.") 

81 group.add_argument("--prompt-template-video", type=str, default='li-dit-encode-video', choices=PROMPT_TEMPLATE, 

82 help="Video prompt template for the decoder-only text encoder model.") 

83 group.add_argument("--hidden-state-skip-layer", type=int, default=2, 

84 help="Skip layer for hidden states.") 

85 group.add_argument("--apply-final-norm", action="store_true", 

86 help="Apply final normalization to the used text encoder hidden states.") 

87 

88 # - CLIP 

89 group.add_argument("--text-encoder-2", type=str, default='clipL', choices=list(TEXT_ENCODER_PATH), 

90 help="Name of the second text encoder model.") 

91 group.add_argument("--text-encoder-precision-2", type=str, default="fp16", choices=PRECISIONS, 

92 help="Precision mode for the second text encoder model.") 

93 group.add_argument("--text-states-dim-2", type=int, default=768, 

94 help="Dimension of the second text encoder hidden states.") 

95 group.add_argument("--tokenizer-2", type=str, default='clipL', choices=list(TOKENIZER_PATH), 

96 help="Name of the second tokenizer model.") 

97 group.add_argument("--text-len-2", type=int, default=77, help="Maximum length of the second text input.") 

98 group.set_defaults(use_attention_mask=True) 

99 group.add_argument("--text-projection", type=str, default="single_refiner", choices=TEXT_PROJECTION, 

100 help="A projection layer for bridging the text encoder hidden states and the diffusion model " 

101 "conditions.") 

102 return parser 

103 

104 

105def add_denoise_schedule_args(parser: argparse.ArgumentParser) -> argparse.ArgumentParser: 

106 group = parser.add_argument_group(title="Denoise schedule") 

107 group.add_argument("--flow-shift-eval-video", type=float, default=None, 

108 help="Shift factor for flow matching schedulers when using video data.") 

109 group.add_argument("--flow-reverse", action="store_true", default=True, 

110 help="If reverse, learning/sampling from t=1 -> t=0.") 

111 group.add_argument("--flow-solver", type=str, default="euler", help="Solver for flow matching.") 

112 group.add_argument("--use-linear-quadratic-schedule", action="store_true", 

113 help="Use linear quadratic schedule for flow matching." 

114 "Follow MovieGen (https://ai.meta.com/static-resource/movie-gen-research-paper)") 

115 group.add_argument("--linear-schedule-end", type=int, default=25, 

116 help="End step for linear quadratic schedule for flow matching.") 

117 return parser 

118 

119 

120def add_evaluation_args(parser: argparse.ArgumentParser) -> argparse.ArgumentParser: 

121 group = parser.add_argument_group(title="Validation Loss Evaluation") 

122 parser.add_argument("--precision", type=str, default="bf16", choices=PRECISIONS, 

123 help="Precision mode. Options: fp32, fp16, bf16. Applied to the backbone model and optimizer.") 

124 parser.add_argument("--reproduce", action="store_true", 

125 help="Enable reproducibility by setting random seeds and deterministic algorithms.") 

126 parser.add_argument("--ckpt", type=str, help="Path to the checkpoint to evaluate.") 

127 parser.add_argument("--load-key", type=str, default="module", choices=["module", "ema"], 

128 help="Key to load the model states. 'module' for the main model, 'ema' for the EMA model.") 

129 parser.add_argument("--cpu-offload", action="store_true", help="Use CPU offload for the model load.") 

130 parser.add_argument("--infer-min", action="store_true", help="infer 5s.") 

131 group.add_argument("--use-fp8", action="store_true", help="Enable use fp8 for inference acceleration.") 

132 group.add_argument("--video-size", type=int, nargs='+', default=512, 

133 help="Video size for training. If a single value is provided, it will be used for both width " 

134 "and height. If two values are provided, they will be used for width and height " 

135 "respectively.") 

136 group.add_argument("--sample-n-frames", type=int, default=1, 

137 help="How many frames to sample from a video. if using 3d vae, the number should be 4n+1") 

138 group.add_argument("--infer-steps", type=int, default=100, help="Number of denoising steps for inference.") 

139 group.add_argument("--val-disable-autocast", action="store_true", 

140 help="Disable autocast for denoising loop and vae decoding in pipeline sampling.") 

141 group.add_argument("--num-images", type=int, default=1, help="Number of images to generate for each prompt.") 

142 group.add_argument("--seed", type=int, default=1024, help="Seed for evaluation.") 

143 group.add_argument("--save-path-suffix", type=str, default="", help="Suffix for the directory of saved samples.") 

144 group.add_argument("--pos-prompt", type=str, default='', help="Prompt for sampling during evaluation.") 

145 group.add_argument("--neg-prompt", type=str, default='', help="Negative prompt for sampling during evaluation.") 

146 group.add_argument("--image-size", type=int, default=704) 

147 group.add_argument("--pad-face-size", type=float, default=0.7, help="Pad bbox for face align.") 

148 group.add_argument("--image-path", type=str, default="", help="") 

149 group.add_argument("--save-path", type=str, default=None, help="Path to save the generated samples.") 

150 group.add_argument("--input", type=str, default=None, help="test data.") 

151 group.add_argument("--item-name", type=str, default=None, help="") 

152 group.add_argument("--cfg-scale", type=float, default=7.5, help="Classifier free guidance scale.") 

153 group.add_argument("--ip-cfg-scale", type=float, default=0, help="Classifier free guidance scale.") 

154 group.add_argument("--use-deepcache", type=int, default=1) 

155 return parser 

156 

157 

158def sanity_check_args(args: argparse.Namespace) -> argparse.Namespace: 

159 # VAE channels 

160 vae_pattern = r"\d{2,3}-\d{1,2}c-\w+" 

161 if not re.match(vae_pattern, args.vae): 

162 raise ValueError( 

163 f"Invalid VAE model: {args.vae}. Must be in the format of '{vae_pattern}'." 

164 ) 

165 vae_channels = int(args.vae.split("-")[1][:-1]) 

166 if args.latent_channels is None: 

167 args.latent_channels = vae_channels 

168 if vae_channels != args.latent_channels: 

169 raise ValueError(f"Latent ({args.latent_channels}) must match VAE ({vae_channels}).") 

170 return args