diff --git a/nodes.py b/nodes.py index 112f541..52b40a7 100644 --- a/nodes.py +++ b/nodes.py @@ -666,8 +666,6 @@ class WanVideoModelLoader: if "patch_embedding" in name: dtype_to_use = torch.float32 set_module_tensor_to_device(transformer, name, device=transformer_load_device, dtype=dtype_to_use, value=sd[name]) - for name, param in transformer.named_parameters(): - print(name, param.device, param.dtype) comfy_model.diffusion_model = transformer comfy_model.load_device = transformer_load_device @@ -1758,9 +1756,15 @@ class WanVideoPhantomEmbeds: def INPUT_TYPES(s): return {"required": { "num_frames": ("INT", {"default": 81, "min": 1, "max": 10000, "step": 4, "tooltip": "Number of frames to encode"}), - "phantom_latents": ("LATENT", {"tooltip": "reference latents for the phantom model"}), + "phantom_latent_1": ("LATENT", {"tooltip": "reference latents for the phantom model"}), + "phantom_cfg_scale": ("FLOAT", {"default": 5.0, "min": 0.0, "max": 10.0, "step": 0.01, "tooltip": "CFG scale for the extra phantom cond pass"}), }, + "optional": { + "phantom_latent_2": ("LATENT", {"tooltip": "reference latents for the phantom model"}), + "phantom_latent_3": ("LATENT", {"tooltip": "reference latents for the phantom model"}), + "phantom_latent_4": ("LATENT", {"tooltip": "reference latents for the phantom model"}), + } } RETURN_TYPES = ("WANVIDIMAGE_EMBEDS", ) @@ -1768,9 +1772,15 @@ class WanVideoPhantomEmbeds: FUNCTION = "process" CATEGORY = "WanVideoWrapper" - def process(self, num_frames, phantom_latents, phantom_cfg_scale): + def process(self, num_frames, phantom_cfg_scale, phantom_latent_1, phantom_latent_2=None, phantom_latent_3=None, phantom_latent_4=None): vae_stride = (4, 8, 8) - samples = phantom_latents["samples"].squeeze(0) + samples = phantom_latent_1["samples"].squeeze(0) + if phantom_latent_2 is not None: + samples = torch.cat([samples, phantom_latent_2["samples"].squeeze(0)], dim=1) + if phantom_latent_3 is not None: + samples = torch.cat([samples, phantom_latent_3["samples"].squeeze(0)], dim=1) + if phantom_latent_4 is not None: + samples = torch.cat([samples, phantom_latent_4["samples"].squeeze(0)], dim=1) C, T, H, W = samples.shape target_shape = (16, (num_frames - 1) // vae_stride[0] + 1 + T, @@ -2637,6 +2647,7 @@ class WanVideoSampler: self.teacache_state = [None, None] if phantom_latents is not None: + log.info(f"Phantom latents shape: {phantom_latents.shape}") self.teacache_state = [None, None, None] self.teacache_state_source = [None, None] self.teacache_states_context = []