fp8 text encoder fixes

This commit is contained in:
kijai
2025-08-01 17:42:32 +03:00
parent 8321e19009
commit 83bb0e7f35
2 changed files with 23 additions and 5 deletions
+16 -5
View File
@@ -324,12 +324,18 @@ class WanVideoTextEncode:
device_to = mm.get_torch_device()
else:
device_to = torch.device("cpu")
if encoder.quantization == "fp8_e4m3fn":
cast_dtype = torch.float8_e4m3fn
else:
cast_dtype = encoder.dtype
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 encoder.dtype
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])
with torch.autocast(device_type=mm.get_autocast_device(device_to), dtype=dtype, enabled=encoder.quantization != 'disabled'):
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
if use_disk_cache and context is not None:
pass
@@ -453,12 +459,17 @@ class WanVideoTextEncodeSingle:
device_to = mm.get_torch_device()
else:
device_to = torch.device("cpu")
if encoder.quantization == "fp8_e4m3fn":
cast_dtype = torch.float8_e4m3fn
else:
cast_dtype = encoder.dtype
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 encoder.dtype
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])
with torch.autocast(device_type=mm.get_autocast_device(device_to), dtype=dtype, enabled=encoder.quantization != 'disabled'):
with torch.autocast(device_type=mm.get_autocast_device(device_to), dtype=encoder.dtype, enabled=encoder.quantization != 'disabled'):
encoded = encoder([prompt], device_to)
if force_offload:
@@ -1848,7 +1859,7 @@ class WanVideoSampler:
shot_num = len(text_embeds["prompt_embeds"])
shot_len = [latent_video_length//shot_num] * (shot_num-1)
shot_len.append(latent_video_length-sum(shot_len))
log.info(f"Number of shots in prompt: {shot_num}, Shot token lengths: {shot_len}")
log.info(f"EchoShot - Number of shots in prompt: {shot_num}, Shot token lengths: {shot_len}")
#region transformer settings
#rope
+7
View File
@@ -1387,6 +1387,13 @@ class LoadWanVideoT5TextEncoder:
model_path = folder_paths.get_full_path("text_encoders", model_name)
sd = load_torch_file(model_path, safe_load=True)
if quantization == "disabled":
for k, v in sd.items():
if isinstance(v, torch.Tensor):
if v.dtype == torch.float8_e4m3fn:
quantization = "fp8_e4m3fn"
break
if "token_embedding.weight" not in sd and "shared.weight" not in sd:
raise ValueError("Invalid T5 text encoder model, this node expects the 'umt5-xxl' model")