Fix T5 double RAM use

This commit is contained in:
kijai
2025-08-03 12:28:25 +03:00
parent 26c1911e41
commit 5cde6f2216
+7 -1
View File
@@ -387,7 +387,12 @@ class WanVideoTextEncode:
params_to_keep = {'norm', 'pos_embedding', 'token_embedding'}
for name, param in encoder.model.named_parameters():
dtype_to_use = dtype if any(keyword in name for keyword in params_to_keep) else cast_dtype
set_module_tensor_to_device(encoder.model, name, device=device_to, dtype=dtype_to_use, value=encoder.state_dict[name])
value = encoder.state_dict[name] if hasattr(encoder, 'state_dict') else encoder.model.state_dict()[name]
set_module_tensor_to_device(encoder.model, name, device=device_to, dtype=dtype_to_use, value=value)
if hasattr(encoder, 'state_dict'):
del encoder.state_dict
mm.soft_empty_cache()
gc.collect()
with torch.autocast(device_type=mm.get_autocast_device(device_to), dtype=encoder.dtype, enabled=encoder.quantization != 'disabled'):
# Encode positive if not loaded from cache
@@ -411,6 +416,7 @@ class WanVideoTextEncode:
if force_offload:
encoder.model.to(offload_device)
mm.soft_empty_cache()
gc.collect()
prompt_embeds_dict = {
"prompt_embeds": context,