From 013fae65fb9d8bc7c6be268cf1b9db3fcec784ec Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Tue, 4 Mar 2025 01:56:13 +0200 Subject: [PATCH] cleanup --- nodes.py | 3 +-- wanvideo/modules/attention.py | 4 ++-- wanvideo/modules/t5.py | 2 +- 3 files changed, 4 insertions(+), 5 deletions(-) diff --git a/nodes.py b/nodes.py index ebbbfa7..cb3688f 100644 --- a/nodes.py +++ b/nodes.py @@ -1268,7 +1268,6 @@ class WanVideoSampler: img_emb = image_embeds.get("image_embeds", None) partial_img_emb = None if img_emb is not None: - print("img_emb shape", img_emb.shape) partial_img_emb = img_emb[:, c, :, :] partial_img_emb[:, 0, :, :] = img_emb[:, 0, :, :].to(intermediate_device) @@ -1377,8 +1376,8 @@ class WanVideoSampler: mm.soft_empty_cache() gc.collect() - print_memory(device) try: + print_memory(device) torch.cuda.reset_peak_memory_stats(device) except: pass diff --git a/wanvideo/modules/attention.py b/wanvideo/modules/attention.py index 85c83e1..ca966c8 100644 --- a/wanvideo/modules/attention.py +++ b/wanvideo/modules/attention.py @@ -62,8 +62,8 @@ def flash_attention( dtype: torch.dtype. Apply when dtype of q/k/v is not float16/bfloat16. """ half_dtypes = (torch.float16, torch.bfloat16) - assert dtype in half_dtypes - assert q.device.type == 'cuda' and q.size(-1) <= 256 + #assert dtype in half_dtypes + #assert q.device.type == 'cuda' and q.size(-1) <= 256 # params b, lq, lk, out_dtype = q.size(0), q.size(1), k.size(1), q.dtype diff --git a/wanvideo/modules/t5.py b/wanvideo/modules/t5.py index 3ba05af..9fc7f67 100644 --- a/wanvideo/modules/t5.py +++ b/wanvideo/modules/t5.py @@ -477,7 +477,7 @@ class T5EncoderModel: self, text_len, dtype=torch.bfloat16, - device=torch.cuda.current_device(), + device=torch.device('cuda'), state_dict=None, tokenizer_path=None, quantization="disabled",