From 104c63235988e84fd3c655648d77ea342e284f4d Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Fri, 28 Mar 2025 01:58:47 +0200 Subject: [PATCH] some cleanup --- nodes.py | 54 +++++++++++++++++------------------- wanvideo/modules/model.py | 58 ++++----------------------------------- 2 files changed, 30 insertions(+), 82 deletions(-) diff --git a/nodes.py b/nodes.py index 5b09d80..4037e7d 100644 --- a/nodes.py +++ b/nodes.py @@ -1708,7 +1708,7 @@ class WanVideoSampler: "flowedit_args": ("FLOWEDITARGS", ), "batched_cfg": ("BOOLEAN", {"default": False, "tooltip": "Batc cond and uncond for faster sampling, possibly faster on some hardware, uses more memory"}), "slg_args": ("SLGARGS", ), - "rope_function": (["default", "comfy"], {"default": "default", "tooltip": "!EXPERIMENTAL! Comfy's RoPE implementation doesn't use complex numbers and can thus be compiled, that should be a lot faster when using torch.compile"}), + "rope_function": (["default", "comfy"], {"default": "comfy", "tooltip": "Comfy's RoPE implementation doesn't use complex numbers and can thus be compiled, that should be a lot faster when using torch.compile"}), "loop_args": ("LOOPARGS", ), "experimental_args": ("EXPERIMENTALARGS", ), } @@ -2006,22 +2006,22 @@ class WanVideoSampler: self.teacache_states_context = [] - if "sparge" in transformer.attention_mode: - from spas_sage_attn.autotune import ( - SparseAttentionMeansim, - extract_sparse_attention_state_dict, - load_sparse_attention_state_dict, - ) + # if "sparge" in transformer.attention_mode: + # from spas_sage_attn.autotune import ( + # SparseAttentionMeansim, + # extract_sparse_attention_state_dict, + # load_sparse_attention_state_dict, + # ) - for idx, block in enumerate(transformer.blocks): - block.self_attn.verbose = True - block.self_attn.inner_attention = SparseAttentionMeansim(l1=0.06, pv_l1=0.065) - if transformer.attention_mode == "spargeattn": - try: - saved_state_dict = torch.load("sparge_wan.pt") - except: - raise ValueError("No saved parameters found for sparse attention, tuning is required first") - load_sparse_attention_state_dict(transformer, saved_state_dict, verbose = True) + # for idx, block in enumerate(transformer.blocks): + # block.self_attn.verbose = True + # block.self_attn.inner_attention = SparseAttentionMeansim(l1=0.06, pv_l1=0.065) + # if transformer.attention_mode == "spargeattn": + # try: + # saved_state_dict = torch.load("sparge_wan.pt") + # except: + # raise ValueError("No saved parameters found for sparse attention, tuning is required first") + # load_sparse_attention_state_dict(transformer, saved_state_dict, verbose = True) if flowedit_args is not None: source_embeds = flowedit_args["source_embeds"] @@ -2137,17 +2137,13 @@ class WanVideoSampler: ) noise_pred_uncond=noise_pred_uncond[0].to(intermediate_device) - noise_pred_text = noise_pred_cond - #https://github.com/WeichenFan/CFG-Zero-star/ if use_cfg_zero_star: - positive_flat = noise_pred_text.view(batch_size, -1) - negative_flat = noise_pred_uncond.view(batch_size, -1) - - alpha = optimized_scale(positive_flat,negative_flat) - alpha = alpha.view(batch_size, 1, 1, 1) - - noise_pred = noise_pred_uncond * alpha + cfg_scale * (noise_pred_text - noise_pred_uncond * alpha) + alpha = optimized_scale( + noise_pred_cond.view(batch_size, -1), + noise_pred_uncond.view(batch_size, -1) + ).view(batch_size, 1, 1, 1) + noise_pred = noise_pred_uncond * alpha + cfg_scale * (noise_pred_cond - noise_pred_uncond * alpha) else: noise_pred = noise_pred_uncond + cfg_scale * (noise_pred_cond - noise_pred_uncond) return noise_pred, [teacache_state_cond, teacache_state_uncond] @@ -2481,10 +2477,10 @@ class WanVideoSampler: log.info(f"TeaCache skipped: {len(state['skipped_steps'])} {name} steps: {state['skipped_steps']}") transformer.teacache_state.clear_all() - if transformer.attention_mode == "spargeattn_tune": - saved_state_dict = extract_sparse_attention_state_dict(transformer) - torch.save(saved_state_dict, "sparge_wan.pt") - save_torch_file(saved_state_dict, "sparge_wan.safetensors") + # if transformer.attention_mode == "spargeattn_tune": + # saved_state_dict = extract_sparse_attention_state_dict(transformer) + # torch.save(saved_state_dict, "sparge_wan.pt") + # save_torch_file(saved_state_dict, "sparge_wan.safetensors") if force_offload: if model["manual_offloading"]: diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index f81fff8..923cf1c 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -91,49 +91,9 @@ def rope_params(max_seq_len, dim, theta=10000, L_test=25, k=0): from comfy.model_management import get_torch_device, get_autocast_device @torch.autocast(device_type=get_autocast_device(get_torch_device()), enabled=False) +@torch.compiler.disable() def rope_apply(x, grid_sizes, freqs): n, c = x.size(2), x.size(3) // 2 - - freqs = freqs.split([c - 2 * (c // 3), c // 3, c // 3], dim=1) - - output = [] - for i, (f, h, w) in enumerate(grid_sizes.tolist()): - seq_len = f * h * w - - @torch.compiler.disable() - def view_as_complex_no_compile(x): - x_i = torch.view_as_complex(x[i, :seq_len].to(torch.float64).reshape(seq_len, n, -1, 2)) - return x_i - - x_i = view_as_complex_no_compile(x) - - f_size = (f, 1, 1, -1) - h_size = (1, h, 1, -1) - w_size = (1, 1, w, -1) - - freq_cat = torch.cat([ - freqs[0][:f].view(*f_size).expand(f, h, w, -1), - freqs[1][:h].view(*h_size).expand(f, h, w, -1), - freqs[2][:w].view(*w_size).expand(f, h, w, -1) - ], dim=-1).reshape(seq_len, 1, -1).to(x.device) - - @torch.compiler.disable() - def view_as_real_no_compile(x_i): - x_i.mul_(freq_cat) - x_i = torch.view_as_real(x_i).flatten(2) - return x_i - - x_i = view_as_real_no_compile(x_i) - del freq_cat - - if seq_len < x.size(1): - x_i = torch.cat([x_i, x[i, seq_len:]], dim=0) - output.append(x_i) - - return torch.stack(output).to(torch.float32) - -def rope_apply_original(x, grid_sizes, freqs): - n, c = x.size(2), x.size(3) // 2 # split freqs freqs = freqs.split([c - 2 * (c // 3), c // 3, c // 3], dim=1) @@ -766,9 +726,6 @@ class WanModel(ModelMixin, ConfigMixin): if model_type == 'i2v': self.img_emb = MLPProj(1280, dim) - # initialize weights - #self.init_weights() - def block_swap(self, blocks_to_swap, offload_txt_emb=False, offload_img_emb=False): print(f"Swapping {blocks_to_swap + 1} transformer blocks") self.blocks_to_swap = blocks_to_swap @@ -790,8 +747,7 @@ class WanModel(ModelMixin, ConfigMixin): mm.soft_empty_cache() gc.collect() - - #print(f"Block {b}: {block_memory:.2f}MB on {block.parameters().__next__().device}") + log.info("----------------------") log.info(f"Block swap memory summary:") log.info(f"Transformer blocks on {self.offload_device}: {total_offload_memory:.2f}MB") @@ -837,8 +793,6 @@ class WanModel(ModelMixin, ConfigMixin): List[Tensor]: List of denoised video tensors with original input shapes [C_out, F, H / 8, W / 8] """ - #if self.model_type == 'i2v': - # assert clip_fea is not None and y is not None # params device = self.patch_embedding.weight.device if freqs is not None and freqs.device != device: @@ -922,7 +876,7 @@ class WanModel(ModelMixin, ConfigMixin): accumulated_rel_l1_distance = torch.tensor(0.0, dtype=torch.float32, device=device) if self.enable_teacache and self.teacache_start_step <= current_step <= self.teacache_end_step: if pred_id is None: - pred_id = self.teacache_state.new_prediction() + pred_id = self.teacache_state.new_prediction(cache_device=self.teacache_cache_device) #log.info(current_step) #log.info(f"TeaCache: Initializing TeaCache variables for model pred: {pred_id}") should_calc = True @@ -993,11 +947,8 @@ class WanModel(ModelMixin, ConfigMixin): accumulated_rel_l1_distance=accumulated_rel_l1_distance, previous_modulated_input=previous_modulated_input ) - #self.teacache_state.report() - # head x = self.head(x, e) - # unpatchify x = self.unpatchify(x, grid_sizes) # type: ignore[arg-type] x = [u.float() for u in x] return (x, pred_id) if pred_id is not None else (x, None) @@ -1034,8 +985,9 @@ class TeaCacheState: self.states = {} self._next_pred_id = 0 - def new_prediction(self): + def new_prediction(self, cache_device='cpu'): """Create new prediction state and return its ID""" + self.cache_device = cache_device pred_id = self._next_pred_id self._next_pred_id += 1 self.states[pred_id] = {