From 58d90daf1bc6e9c55a0f06cf1f0bb5cfc0ca14f2 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Wed, 19 Jun 2024 15:35:49 +0300 Subject: [PATCH] Add basic Gemma llm inference node for testing --- nodes.py | 88 +++++++++++++++++++++++++++++++++++++++++++++++++------- 1 file changed, 77 insertions(+), 11 deletions(-) diff --git a/nodes.py b/nodes.py index 91581eb..ac23de8 100644 --- a/nodes.py +++ b/nodes.py @@ -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" } \ No newline at end of file