diff --git a/hyvideo/text_encoder/__init__.py b/hyvideo/text_encoder/__init__.py index 494992b..41a9713 100644 --- a/hyvideo/text_encoder/__init__.py +++ b/hyvideo/text_encoder/__init__.py @@ -21,6 +21,7 @@ def load_text_encoder( text_encoder_path=None, logger=None, device=None, + quantization_config=None, ): if text_encoder_path is None: text_encoder_path = TEXT_ENCODER_PATH[text_encoder_type] @@ -34,14 +35,16 @@ def load_text_encoder( text_encoder.final_layer_norm = text_encoder.text_model.final_layer_norm elif text_encoder_type == "llm": text_encoder = AutoModel.from_pretrained( - text_encoder_path, low_cpu_mem_usage=True + text_encoder_path, + low_cpu_mem_usage=True, + quantization_config=quantization_config ) text_encoder.final_layer_norm = text_encoder.norm else: raise ValueError(f"Unsupported text encoder type: {text_encoder_type}") # from_pretrained will ensure that the model is in eval mode. - if text_encoder_precision is not None: + if text_encoder_precision is not None and quantization_config is None: text_encoder = text_encoder.to(dtype=PRECISION_TO_TYPE[text_encoder_precision]) text_encoder.requires_grad_(False) @@ -49,7 +52,7 @@ def load_text_encoder( if logger is not None: logger.info(f"Text encoder to dtype: {text_encoder.dtype}") - if device is not None: + if device is not None and quantization_config is None: text_encoder = text_encoder.to(device) return text_encoder, text_encoder_path @@ -118,6 +121,7 @@ class TextEncoder(nn.Module): reproduce: bool = False, logger=None, device=None, + quantization_config=None, ): super().__init__() self.text_encoder_type = text_encoder_type @@ -183,6 +187,7 @@ class TextEncoder(nn.Module): text_encoder_path=self.model_path, logger=self.logger, device=device, + quantization_config=quantization_config, ) self.dtype = self.model.dtype self.device = self.model.device diff --git a/nodes.py b/nodes.py index ef8374a..dca64c1 100644 --- a/nodes.py +++ b/nodes.py @@ -374,6 +374,7 @@ class DownloadAndLoadHyVideoTextEncoder: "optional": { "apply_final_norm": ("BOOLEAN", {"default": False}), "hidden_state_skip_layer": ("INT", {"default": 2}), + "quantization": (['disabled', 'bnb_nf4'], {"default": 'disabled'}), } } @@ -383,11 +384,22 @@ class DownloadAndLoadHyVideoTextEncoder: CATEGORY = "HunyuanVideoWrapper" DESCRIPTION = "Loads Hunyuan text_encoder model from 'ComfyUI/models/LLM'" - def loadmodel(self, llm_model, clip_model, precision, apply_final_norm=False, hidden_state_skip_layer=2, use_prompt_templates=True): + def loadmodel(self, llm_model, clip_model, precision, apply_final_norm=False, hidden_state_skip_layer=2, use_prompt_templates=True, quantization="disabled"): device = mm.get_torch_device() offload_device = mm.unet_offload_device() dtype = {"bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32}[precision] + if quantization == "bnb_nf4": + from transformers import BitsAndBytesConfig + + quantization_config = BitsAndBytesConfig( + load_in_4bit=True, + bnb_4bit_quant_type="nf4", + bnb_4bit_use_double_quant=True, + bnb_4bit_compute_dtype=torch.bfloat16 + ) + else: + quantization_config = None if clip_model != "disabled": clip_model_path = os.path.join(folder_paths.models_dir, "clip", "clip-vit-large-patch14") if not os.path.exists(clip_model_path): @@ -443,6 +455,7 @@ class DownloadAndLoadHyVideoTextEncoder: apply_final_norm=apply_final_norm, logger=log, device=device, + quantization_config=quantization_config )