fp8 option for T5
This commit is contained in:
@@ -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 = {
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user