ChatGLM from safetensors

This commit is contained in:
kijai
2024-07-07 17:58:46 +03:00
parent db684d4519
commit 3378d3f7af
5 changed files with 108 additions and 5 deletions
+42
View File
@@ -0,0 +1,42 @@
{
"_name_or_path": "THUDM/chatglm3-6b-base",
"model_type": "chatglm",
"architectures": [
"ChatGLMModel"
],
"auto_map": {
"AutoConfig": "configuration_chatglm.ChatGLMConfig",
"AutoModel": "modeling_chatglm.ChatGLMForConditionalGeneration",
"AutoModelForCausalLM": "modeling_chatglm.ChatGLMForConditionalGeneration",
"AutoModelForSeq2SeqLM": "modeling_chatglm.ChatGLMForConditionalGeneration",
"AutoModelForSequenceClassification": "modeling_chatglm.ChatGLMForSequenceClassification"
},
"add_bias_linear": false,
"add_qkv_bias": true,
"apply_query_key_layer_scaling": true,
"apply_residual_connection_post_layernorm": false,
"attention_dropout": 0.0,
"attention_softmax_in_fp32": true,
"bias_dropout_fusion": true,
"ffn_hidden_size": 13696,
"fp32_residual_connection": false,
"hidden_dropout": 0.0,
"hidden_size": 4096,
"kv_channels": 128,
"layernorm_epsilon": 1e-05,
"multi_query_attention": true,
"multi_query_group_num": 2,
"num_attention_heads": 32,
"num_layers": 28,
"original_rope": true,
"padded_vocab_size": 65024,
"post_layer_norm": true,
"rmsnorm": true,
"seq_length": 32768,
"use_cache": true,
"torch_dtype": "float16",
"transformers_version": "4.30.2",
"tie_word_embeddings": false,
"eos_token_id": 2,
"pad_token_id": 0
}
Binary file not shown.
+12
View File
@@ -0,0 +1,12 @@
{
"name_or_path": "THUDM/chatglm3-6b-base",
"remove_space": false,
"do_lower_case": false,
"tokenizer_class": "ChatGLMTokenizer",
"auto_map": {
"AutoTokenizer": [
"tokenization_chatglm.ChatGLMTokenizer",
null
]
}
}
Binary file not shown.
+54 -5
View File
@@ -3,17 +3,18 @@ import os
import random import random
import re import re
import gc import gc
import sys import json
import comfy.model_management as mm import comfy.model_management as mm
from comfy.utils import ProgressBar, load_torch_file from comfy.utils import ProgressBar, load_torch_file
import folder_paths import folder_paths
script_directory = os.path.dirname(os.path.abspath(__file__)) script_directory = os.path.dirname(os.path.abspath(__file__))
sys.path.append(script_directory)
folder_paths.add_model_folder_path("LLM", os.path.join(folder_paths.models_dir, "LLM", "checkpoints"))
from .kolors.pipelines.pipeline_stable_diffusion_xl_chatglm_256 import StableDiffusionXLPipeline from .kolors.pipelines.pipeline_stable_diffusion_xl_chatglm_256 import StableDiffusionXLPipeline
from .kolors.models.modeling_chatglm import ChatGLMModel from .kolors.models.modeling_chatglm import ChatGLMModel, ChatGLMConfig
from .kolors.models.tokenization_chatglm import ChatGLMTokenizer from .kolors.models.tokenization_chatglm import ChatGLMTokenizer
from diffusers import UNet2DConditionModel from diffusers import UNet2DConditionModel
from diffusers import (DPMSolverMultistepScheduler, from diffusers import (DPMSolverMultistepScheduler,
@@ -83,6 +84,52 @@ class DownloadAndLoadKolorsModel:
return (kolors_model,) return (kolors_model,)
class LoadChatGLM3:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"chatglm3_checkpoint": (folder_paths.get_filename_list("LLM"),),
"precision": ([ 'fp16', 'quant4', 'quant8'],
{
"default": 'fp16'
}),
},
}
RETURN_TYPES = ("CHATGLM3MODEL",)
RETURN_NAMES = ("chatglm3_model",)
FUNCTION = "loadmodel"
CATEGORY = "KwaiKolorsWrapper"
def loadmodel(self, chatglm3_checkpoint, precision):
pbar = ProgressBar(2)
chatglm3_path = folder_paths.get_full_path("LLM", chatglm3_checkpoint)
print("Load TEXT_ENCODER...")
text_encoder_config = os.path.join(script_directory, 'configs', 'text_encoder_config.json')
with open(text_encoder_config, 'r') as file:
config = json.load(file)
text_encoder_config = ChatGLMConfig(**config)
text_encoder = ChatGLMModel(text_encoder_config)
text_encoder.load_state_dict(load_torch_file(chatglm3_path))
if precision == 'quant8':
text_encoder.quantize(8)
elif precision == 'quant4':
text_encoder.quantize(4)
tokenizer_path = os.path.join(script_directory,'configs',"tokenizer")
tokenizer = ChatGLMTokenizer.from_pretrained(tokenizer_path)
pbar.update(1)
chatglm3_model = {
'text_encoder': text_encoder,
'tokenizer': tokenizer
}
return (chatglm3_model,)
class DownloadAndLoadChatGLM3: class DownloadAndLoadChatGLM3:
@classmethod @classmethod
def INPUT_TYPES(s): def INPUT_TYPES(s):
@@ -396,11 +443,13 @@ NODE_CLASS_MAPPINGS = {
"DownloadAndLoadKolorsModel": DownloadAndLoadKolorsModel, "DownloadAndLoadKolorsModel": DownloadAndLoadKolorsModel,
"DownloadAndLoadChatGLM3": DownloadAndLoadChatGLM3, "DownloadAndLoadChatGLM3": DownloadAndLoadChatGLM3,
"KolorsSampler": KolorsSampler, "KolorsSampler": KolorsSampler,
"KolorsTextEncode": KolorsTextEncode "KolorsTextEncode": KolorsTextEncode,
"LoadChatGLM3": LoadChatGLM3
} }
NODE_DISPLAY_NAME_MAPPINGS = { NODE_DISPLAY_NAME_MAPPINGS = {
"DownloadAndLoadKolorsModel": "(Down)load Kolors Model", "DownloadAndLoadKolorsModel": "(Down)load Kolors Model",
"DownloadAndLoadChatGLM3": "(Down)load ChatGLM3 Model", "DownloadAndLoadChatGLM3": "(Down)load ChatGLM3 Model",
"KolorsSampler": "Kolors Sampler", "KolorsSampler": "Kolors Sampler",
"KolorsTextEncode": "Kolors Text Encode" "KolorsTextEncode": "Kolors Text Encode",
"LoadChatGLM3": "Load ChatGLM3 Model"
} }