(config: Dict[str, Any])
| 208 | |
| 209 | |
| 210 | def load_wan_image_encoder(config: Dict[str, Any]): |
| 211 | from lightx2v.models.input_encoders.hf.wan.xlm_roberta.model import CLIPModel |
| 212 | |
| 213 | image_encoder = None |
| 214 | if config["task"] in ["i2v", "flf2v", "animate", "s2v"] and config.get("use_image_encoder", True): |
| 215 | # offload config |
| 216 | clip_offload = config.get("clip_cpu_offload", config.get("cpu_offload", False)) |
| 217 | if clip_offload: |
| 218 | clip_device = torch.device("cpu") |
| 219 | else: |
| 220 | clip_device = torch.device(AI_DEVICE) |
| 221 | # quant_config |
| 222 | clip_quantized = config.get("clip_quantized", False) |
| 223 | if clip_quantized: |
| 224 | clip_quant_scheme = config.get("clip_quant_scheme", None) |
| 225 | assert clip_quant_scheme is not None |
| 226 | tmp_clip_quant_scheme = clip_quant_scheme.split("-")[0] |
| 227 | clip_model_name = f"models_clip_open-clip-xlm-roberta-large-vit-huge-14-{tmp_clip_quant_scheme}.pth" |
| 228 | clip_quantized_ckpt = find_torch_model_path(config, "clip_quantized_ckpt", clip_model_name) |
| 229 | clip_original_ckpt = None |
| 230 | else: |
| 231 | clip_quantized_ckpt = None |
| 232 | clip_quant_scheme = None |
| 233 | clip_model_name = "models_clip_open-clip-xlm-roberta-large-vit-huge-14.pth" |
| 234 | clip_original_ckpt = find_torch_model_path(config, "clip_original_ckpt", clip_model_name) |
| 235 | |
| 236 | image_encoder = CLIPModel( |
| 237 | dtype=torch.float16, |
| 238 | device=clip_device, |
| 239 | checkpoint_path=clip_original_ckpt, |
| 240 | clip_quantized=clip_quantized, |
| 241 | clip_quantized_ckpt=clip_quantized_ckpt, |
| 242 | quant_scheme=clip_quant_scheme, |
| 243 | cpu_offload=clip_offload, |
| 244 | use_31_block=config.get("use_31_block", True), |
| 245 | load_from_rank0=config.get("load_from_rank0", False), |
| 246 | ) |
| 247 | |
| 248 | return image_encoder |
| 249 | |
| 250 | |
| 251 | def get_vae_parallel(config: Dict[str, Any]): |
no test coverage detected