From 7e19d75480305098629d46837dbeb29e05631b50 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Wed, 26 Feb 2025 11:35:52 +0200 Subject: [PATCH] fp8 option for T5 --- nodes.py | 23 ++++++++++++++++------- wanvideo/modules/clip.py | 4 +++- wanvideo/modules/t5.py | 13 ++++++++++--- 3 files changed, 29 insertions(+), 11 deletions(-) diff --git a/nodes.py b/nodes.py index 494dad3..5639dee 100644 --- a/nodes.py +++ b/nodes.py @@ -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 = { diff --git a/wanvideo/modules/clip.py b/wanvideo/modules/clip.py index cbc9c64..8452648 100644 --- a/wanvideo/modules/clip.py +++ b/wanvideo/modules/clip.py @@ -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 diff --git a/wanvideo/modules/t5.py b/wanvideo/modules/t5.py index cc44558..2c77ac3 100644 --- a/wanvideo/modules/t5.py +++ b/wanvideo/modules/t5.py @@ -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