Files
SeargeDP-ComfyUI_Searge_LLM/LLM_Node.py
T

70 lines
2.4 KiB
Python

from transformers import (
AutoTokenizer,
AutoModelForSequenceClassification,
AutoModelForCausalLM,
AutoModelForSeq2SeqLM,
AutoConfig
)
import os
import folder_paths
GLOBAL_MODELS_DIR = os.path.join(folder_paths.models_dir, "LLM_checkpoints")
class LLM_Node:
def __init__(self, device="cuda"):
self.device = device
@classmethod
def INPUT_TYPES(cls):
# Get a list of directories in the checkpoints_path
model_options = next(os.walk(GLOBAL_MODELS_DIR), (None, [], None))[1]
return {
"required": {
"text": ("STRING", {"multiline": True, "default": ""}),
"model": ("STRING", {"default": model_options[0] if model_options else "none", "options": model_options}),
"max_tokens": ("INT", {"default": 2000}),
}
}
RETURN_TYPES = ("STRING",)
RETURN_NAMES = ("string",)
OUTPUT_NODE = False
FUNCTION = "llm_text"
CATEGORY = "LLM/text"
def llm_text(self, text, model, max_tokens: int = 2000):
model_path = os.path.join(GLOBAL_MODELS_DIR, model)
# Load the model and tokenizer based on the model's configuration
config = AutoConfig.from_pretrained(model_path)
tokenizer = AutoTokenizer.from_pretrained(model_path)
# Distinguish between different model types appropriately
if config.model_type == "t5":
model = AutoModelForSeq2SeqLM.from_pretrained(model_path)
elif config.model_type in ["gpt2"]:
model = AutoModelForCausalLM.from_pretrained(model_path)
elif config.model_type == "bert":
model = AutoModelForSequenceClassification.from_pretrained(model_path)
else:
raise ValueError(f"Unsupported model type: {config.model_type}")
model.to(self.device)
if config.model_type in ["t5", "gpt2"]:
input_ids = tokenizer(text, return_tensors="pt").input_ids.to(self.device)
outputs = model.generate(input_ids, max_length=max_tokens)
generated_text = tokenizer.decode(outputs[0], skip_special_tokens=True)
return (generated_text,)
elif config.model_type == "bert":
return ("BERT model detected; specific task handling not implemented in this example.",)
NODE_CLASS_MAPPINGS = {
"LLM_Node": LLM_Node,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"LLM_Node": "LLM Node",
}