some cleanup

This commit is contained in:
kijai
2025-03-28 01:58:47 +02:00
parent 2d56d69ccb
commit 104c632359
2 changed files with 30 additions and 82 deletions
+25 -29
View File
@@ -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"]:
+5 -53
View File
@@ -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] = {