cleanup
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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",
|
||||
|
||||
Reference in New Issue
Block a user