(config: Dict[str, Any])
| 290 | |
| 291 | |
| 292 | def load_wan_vae_decoder(config: Dict[str, Any]): |
| 293 | from lightx2v.models.video_encoders.hf.wan.vae import WanVAE |
| 294 | from lightx2v.models.video_encoders.hf.wan.vae_2_2 import Wan2_2_VAE |
| 295 | from lightx2v.models.video_encoders.hf.wan.vae_tiny import Wan2_2_VAE_tiny, WanVAE_tiny |
| 296 | |
| 297 | vae_name = config.get("vae_name", "Wan2.1_VAE.pth") |
| 298 | tiny_vae_name = "taew2_1.pth" |
| 299 | |
| 300 | if config.get("model_cls", "") == "wan2.2": |
| 301 | vae_cls = Wan2_2_VAE |
| 302 | tiny_vae_cls = Wan2_2_VAE_tiny |
| 303 | tiny_vae_name = "taew2_2.pth" |
| 304 | else: |
| 305 | vae_cls = WanVAE |
| 306 | tiny_vae_cls = WanVAE_tiny |
| 307 | tiny_vae_name = "taew2_1.pth" |
| 308 | |
| 309 | # offload config |
| 310 | vae_offload = config.get("vae_cpu_offload", config.get("cpu_offload")) |
| 311 | if vae_offload: |
| 312 | vae_device = torch.device("cpu") |
| 313 | else: |
| 314 | vae_device = torch.device(AI_DEVICE) |
| 315 | |
| 316 | vae_config = { |
| 317 | "vae_path": find_torch_model_path(config, "vae_path", vae_name), |
| 318 | "device": vae_device, |
| 319 | "parallel": get_vae_parallel(config), |
| 320 | "use_tiling": config.get("use_tiling_vae", False), |
| 321 | "cpu_offload": vae_offload, |
| 322 | "use_lightvae": config.get("use_lightvae", False), |
| 323 | "dtype": GET_DTYPE(), |
| 324 | "load_from_rank0": config.get("load_from_rank0", False), |
| 325 | } |
| 326 | if config.get("use_tae", False): |
| 327 | tae_path = find_torch_model_path(config, "tae_path", tiny_vae_name) |
| 328 | vae_decoder = tiny_vae_cls(vae_path=tae_path, device=AI_DEVICE, need_scaled=config.get("need_scaled", False)).to(AI_DEVICE) |
| 329 | else: |
| 330 | vae_decoder = vae_cls(**vae_config) |
| 331 | return vae_decoder |
| 332 | |
| 333 | |
| 334 | def load_wan_transformer(config: Dict[str, Any]): |
no test coverage detected