add bitsandbytes nf4 quant as option for the llm

This commit is contained in:
kijai
2024-12-06 18:58:02 +02:00
parent 0cbcc20a89
commit 6655642ea6
2 changed files with 22 additions and 4 deletions
+8 -3
View File
@@ -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
+14 -1
View File
@@ -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
)