fp8 option for T5

This commit is contained in:
kijai
2025-02-26 11:35:52 +02:00
parent cae79ec20b
commit 7e19d75480
3 changed files with 29 additions and 11 deletions
+16 -7
View File
@@ -498,12 +498,13 @@ class LoadWanVideoT5TextEncoder:
return {
"required": {
"model_name": (folder_paths.get_filename_list("text_encoders"), {"tooltip": "These models are loaded from 'ComfyUI/models/vae'"}),
"precision": (["fp16", "fp32", "bf16"],
"precision": (["fp16", "fp32", "bf16"],
{"default": "bf16"}
),
},
"optional": {
"load_device": (["main_device", "offload_device"], {"default": "offload_device"}),
"quantization": (['disabled', 'fp8_e4m3fn'], {"default": 'disabled', "tooltip": "optional quantization method"}),
}
}
@@ -513,7 +514,7 @@ class LoadWanVideoT5TextEncoder:
CATEGORY = "WanVideoWrapper"
DESCRIPTION = "Loads Hunyuan text_encoder model from 'ComfyUI/models/LLM'"
def loadmodel(self, model_name, precision, load_device="offload_device"):
def loadmodel(self, model_name, precision, load_device="offload_device", quantization="disabled"):
device = mm.get_torch_device()
offload_device = mm.unet_offload_device()
@@ -533,9 +534,14 @@ class LoadWanVideoT5TextEncoder:
device=text_encoder_load_device,
state_dict=sd,
tokenizer_path=tokenizer_path,
quantization=quantization
)
text_encoder = {
"model": T5_text_encoder,
"dtype": dtype,
}
return (T5_text_encoder,)
return (text_encoder,)
class LoadWanVideoClipTextEncoder:
@classmethod
@@ -598,16 +604,19 @@ class WanVideoTextEncode:
device = mm.get_torch_device()
offload_device = mm.unet_offload_device()
encoder = t5["model"]
dtype = t5["dtype"]
t5.model.to(device)
encoder.model.to(device)
context = t5([positive_prompt], device)
context_null = t5([negative_prompt], device)
with torch.autocast(device_type=mm.get_autocast_device(device), dtype=dtype, enabled=True):
context = encoder([positive_prompt], device)
context_null = encoder([negative_prompt], device)
context = [t.to(device) for t in context]
context_null = [t.to(device) for t in context_null]
if force_offload:
t5.model.to(offload_device)
encoder.model.to(offload_device)
prompt_embeds_dict = {
+3 -1
View File
@@ -20,6 +20,8 @@ __all__ = [
from accelerate import init_empty_weights
from accelerate.utils import set_module_tensor_to_device
import comfy.model_management as mm
def pos_interpolate(pos, seq_len):
if pos.size(1) == seq_len:
return pos
@@ -529,6 +531,6 @@ class CLIPModel:
def visual(self, image):
# forward
with torch.cuda.amp.autocast(dtype=self.dtype):
with torch.autocast(device_type=mm.get_autocast_device(self.device), dtype=self.dtype):
out = self.model.visual(image, use_31_block=True)
return out
+10 -3
View File
@@ -480,7 +480,7 @@ class T5EncoderModel:
device=torch.cuda.current_device(),
state_dict=None,
tokenizer_path=None,
shard_fn=None,
quantization="disabled",
):
self.text_len = text_len
self.dtype = dtype
@@ -494,9 +494,16 @@ class T5EncoderModel:
return_tokenizer=False,
dtype=dtype,
device=device).eval().requires_grad_(False)
if quantization == "fp8_e4m3fn":
cast_dtype = torch.float8_e4m3fn
else:
cast_dtype = dtype
params_to_keep = {'norm', 'pos_embedding', 'token_embedding'}
for name, param in model.named_parameters():
set_module_tensor_to_device(model, name, device=device, dtype=dtype, value=state_dict[name])
dtype_to_use = dtype if any(keyword in name for keyword in params_to_keep) else cast_dtype
set_module_tensor_to_device(model, name, device=device, dtype=dtype_to_use, value=state_dict[name])
#model.load_state_dict(state_dict)
self.model = model