Update infer_flashvsr_v11_full.py

This commit is contained in:
smthemex
2025-11-18 15:48:04 +08:00
committed by GitHub
parent 924eb1d7ad
commit db8918839d
@@ -255,8 +255,18 @@ def init_pipeline_v11(prompt_path,LQ_proj_in_path="./FlashVSR/LQ_proj_in.ckpt",c
pipe.denoising_model().LQ_proj_in.load_state_dict(torch.load(LQ_proj_in_path, map_location="cpu",weights_only=False), strict=True)
pipe.denoising_model().LQ_proj_in.to('cuda')
pipe.vae.model.encoder = None
pipe.vae.model.conv1 = None
if pipe.vae is None: # safetensors cause error or unknown safetensors or pth
from diffusers import AutoencoderKLWan
config=AutoencoderKLWan.load_config(os.path.join(cur_dir,"FlashVSR/config.json"))
vae=AutoencoderKLWan.from_config(config).to(device,dtype=torch.bfloat16)
vae_dict=torch.load(vae_path, map_location="cpu",weights_only=False) if not vae_path.endswith(".safetensors") else load_file(vae_path,device="cpu")
vae.load_state_dict(vae_dict,strict=False)
del vae_dict
pipe.vae=vae
pipe.vae.encoder = None
else:
pipe.vae.model.encoder = None
pipe.vae.model.conv1 = None
#pipe.to('cuda');
pipe.enable_vram_management(num_persistent_param_in_dit=None)
pipe.init_cross_kv(prompt_path); pipe.load_models_to_device(["dit","vae"])