add bitsandbytes nf4 quant as option for the llm
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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
|
||||
)
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user