diff --git a/nodes.py b/nodes.py index 0e62cf6..5cbf133 100644 --- a/nodes.py +++ b/nodes.py @@ -4079,6 +4079,16 @@ class WanVideoSampler: gen_video_samples = gen_video_samples[:, :, :-1*miss_lengths[0]] del noise, latent + if force_offload: + if model["manual_offloading"]: + transformer.to(offload_device) + mm.soft_empty_cache() + gc.collect() + try: + print_memory(device) + torch.cuda.reset_peak_memory_stats(device) + except: + pass return {"video": gen_video_samples[0].permute(1, 2, 3, 0).cpu()}, #region normal inference diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index 166d8a6..bef99ab 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -1455,6 +1455,7 @@ class WanModel(ModelMixin, ConfigMixin): # MultiTalk if multitalk_audio is not None: + self.audio_proj.to(self.main_device) audio_cond = multitalk_audio.to(device=x.device, dtype=x.dtype) first_frame_audio_emb_s = audio_cond[:, :1, ...] latter_frame_audio_emb = audio_cond[:, 1:, ...] @@ -1470,6 +1471,7 @@ class WanModel(ModelMixin, ConfigMixin): multitalk_audio_embedding = self.audio_proj(first_frame_audio_emb_s, latter_frame_audio_emb_s) human_num = len(multitalk_audio_embedding) multitalk_audio_embedding = torch.concat(multitalk_audio_embedding.split(1), dim=2).to(x.dtype) + self.audio_proj.to(self.offload_device) # convert ref_target_masks to token_ref_target_masks token_ref_target_masks = None