MCPcopy Create free account
hub / github.com/ModelTC/LightX2V / load_wan_vae_decoder

Function load_wan_vae_decoder

lightx2v/disagg/utils.py:292–331  ·  view source on GitHub ↗
(config: Dict[str, Any])

Source from the content-addressed store, hash-verified

290
291
292def 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
334def load_wan_transformer(config: Dict[str, Any]):

Callers 3

load_modelsMethod · 0.90
mainFunction · 0.90
mainFunction · 0.90

Calls 6

find_torch_model_pathFunction · 0.90
GET_DTYPEFunction · 0.90
get_vae_parallelFunction · 0.85
getMethod · 0.45
deviceMethod · 0.45
toMethod · 0.45

Tested by

no test coverage detected