Fix T5 double RAM use
This commit is contained in:
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user