This commit is contained in:
kijai
2025-03-04 01:56:13 +02:00
parent 8c451f4a2f
commit 013fae65fb
3 changed files with 4 additions and 5 deletions
+1 -2
View File
@@ -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
+2 -2
View File
@@ -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
+1 -1
View File
@@ -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",