diff --git a/nodes.py b/nodes.py index 594b9c3..ac83acf 100644 --- a/nodes.py +++ b/nodes.py @@ -2128,39 +2128,16 @@ class WanVideoSampler: rope_function = "default" #echoshot does not support comfy rope function log.info(f"Number of shots in prompt: {shot_num}, Shot token lengths: {shot_len}") - #region transformer settings - #rope - freqs = None - transformer.rope_embedder.k = None - transformer.rope_embedder.num_frames = None - if "default" in rope_function or bidirectional_sampling: - d = transformer.dim // transformer.num_heads - freqs = torch.cat([ - rope_params(1024, d - 4 * (d // 6), L_test=latent_video_length, k=riflex_freq_index), - rope_params(1024, 2 * (d // 6)), - rope_params(1024, 2 * (d // 6)) - ], - dim=1) - elif "comfy" in rope_function: - transformer.rope_embedder.k = riflex_freq_index - transformer.rope_embedder.num_frames = latent_video_length - - transformer.rope_func = rope_function - for block in transformer.blocks: - block.rope_func = rope_function - if transformer.vace_layers is not None: - for block in transformer.vace_blocks: - block.rope_func = rope_function - - #blockswap init mm.unload_all_models() mm.soft_empty_cache() gc.collect() - + + #region transformer settings if transformer_options is not None: block_swap_args = transformer_options.get("block_swap_args", None) + #blockswap init if block_swap_args is not None: transformer.use_non_blocking = block_swap_args.get("use_non_blocking", False) for name, param in transformer.named_parameters(): @@ -2286,6 +2263,29 @@ class WanVideoSampler: import copy sample_scheduler_flipped = copy.deepcopy(sample_scheduler) + #rope + freqs = None + transformer.rope_embedder.k = None + transformer.rope_embedder.num_frames = None + if "default" in rope_function or bidirectional_sampling: + d = transformer.dim // transformer.num_heads + freqs = torch.cat([ + rope_params(1024, d - 4 * (d // 6), L_test=latent_video_length, k=riflex_freq_index), + rope_params(1024, 2 * (d // 6)), + rope_params(1024, 2 * (d // 6)) + ], + dim=1) + elif "comfy" in rope_function: + transformer.rope_embedder.k = riflex_freq_index + transformer.rope_embedder.num_frames = latent_video_length + + transformer.rope_func = rope_function + for block in transformer.blocks: + block.rope_func = rope_function + if transformer.vace_layers is not None: + for block in transformer.vace_blocks: + block.rope_func = rope_function + #region model pred def predict_with_cfg(z, cfg_scale, positive_embeds, negative_embeds, timestep, idx, image_cond=None, clip_fea=None, control_latents=None, vace_data=None, unianim_data=None, audio_proj=None, control_camera_latents=None,