Add basic Gemma llm inference node for testing

This commit is contained in:
kijai
2024-06-19 15:35:49 +03:00
parent db51e52be2
commit 58d90daf1b
+77 -11
View File
@@ -13,7 +13,7 @@ sys.path.append(script_directory)
import lumina_models
from transport import ODE
from transformers import AutoModel, AutoTokenizer
from transformers import AutoModel, AutoTokenizer, GemmaForCausalLM
from contextlib import nullcontext
try:
@@ -101,6 +101,13 @@ class DownloadAndLoadGemmaModel:
"default": 'bf16'
}),
},
"optional": {
"mode": ([ 'text_encode','LLM'],
{
"default": 'text_encode'
}),
}
}
RETURN_TYPES = ("GEMMAODEL",)
@@ -108,7 +115,7 @@ class DownloadAndLoadGemmaModel:
FUNCTION = "loadmodel"
CATEGORY = "LuminaWrapper"
def loadmodel(self, precision):
def loadmodel(self, precision, mode='text_encode'):
device = mm.get_torch_device()
dtype = {"bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32}[precision]
@@ -127,16 +134,25 @@ class DownloadAndLoadGemmaModel:
attn_implementation = "flash_attention_2" if FLASH_ATTN_AVAILABLE and precision != "fp32" else "sdpa"
print(f"Gemma attention mode: {attn_implementation}")
text_encoder = AutoModel.from_pretrained(
gemma_path,
torch_dtype=dtype,
device_map=device,
attn_implementation=attn_implementation,
).eval()
if mode == 'text_encode':
text_encoder = AutoModel.from_pretrained(
gemma_path,
torch_dtype=dtype,
device_map=device,
attn_implementation=attn_implementation,
).eval()
elif mode == 'LLM':
text_encoder = GemmaForCausalLM.from_pretrained(
gemma_path,
torch_dtype=dtype,
device_map=device,
attn_implementation=attn_implementation
).eval()
gemma_model = {
'tokenizer': tokenizer,
'text_encoder': text_encoder
'text_encoder': text_encoder,
}
return (gemma_model,)
@@ -300,6 +316,54 @@ class LuminaGemmaTextEncodeArea:
}
return (lumina_embeds,)
class GemmaSampler:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"gemma_model": ("GEMMAODEL", ),
"prompt": ("STRING", {"multiline": True, "default": "",}),
"max_length": ("INT", {"default": 128, "min": 1, "max": 512, "step": 1}),
},
"optional": {
"keep_model_loaded": ("BOOLEAN", {"default": False}),
}
}
RETURN_TYPES = ("STRING",)
RETURN_NAMES =("string",)
FUNCTION = "process"
CATEGORY = "LuminaWrapper"
def process(self, gemma_model, prompt, max_length, keep_model_loaded=False):
device = mm.get_torch_device()
offload_device = mm.unet_offload_device()
tokenizer = gemma_model['tokenizer']
model = gemma_model['text_encoder']
model.to(device)
text_inputs = tokenizer(
prompt,
return_tensors="pt",
)
text_input_ids = text_inputs.input_ids.to(device)
result = model.generate(
text_input_ids,
max_length=max_length,
)
decoded = tokenizer.batch_decode(result, skip_special_tokens=True, clean_up_tokenization_spaces=False)[0]
print(decoded)
if not keep_model_loaded:
print("Offloading text encoder...")
model.to(offload_device)
return (decoded,)
class LuminaT2ISampler:
@classmethod
@@ -427,7 +491,8 @@ NODE_CLASS_MAPPINGS = {
"DownloadAndLoadGemmaModel": DownloadAndLoadGemmaModel,
"LuminaGemmaTextEncode": LuminaGemmaTextEncode,
"LuminaGemmaTextEncodeArea": LuminaGemmaTextEncodeArea,
"LuminaTextAreaAppend": LuminaTextAreaAppend
"LuminaTextAreaAppend": LuminaTextAreaAppend,
"GemmaSampler": GemmaSampler
}
NODE_DISPLAY_NAME_MAPPINGS = {
"LuminaT2ISampler": "Lumina T2I Sampler",
@@ -435,5 +500,6 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"DownloadAndLoadGemmaModel": "DownloadAndLoadGemmaModel",
"LuminaGemmaTextEncode": "Lumina Gemma Text Encode",
"LuminaGemmaTextEncodeArea": "Lumina Gemma Text Encode Area",
"LuminaTextAreaAppend": "Lumina Text Area Append"
"LuminaTextAreaAppend": "Lumina Text Area Append",
"GemmaSampler": "Gemma Sampler"
}