From 923b9385dda1bc8f07db72d5af89e99b39016433 Mon Sep 17 00:00:00 2001 From: SeanScripts Date: Wed, 25 Sep 2024 15:51:19 -0700 Subject: [PATCH] Add support for Llama 3.2 Vision models --- README.md | 23 +++++--- nodes.py | 133 +++++++++++++++++++++++++++++++++++++++++------ pyproject.toml | 10 ++-- requirements.txt | 1 + 4 files changed, 140 insertions(+), 27 deletions(-) diff --git a/README.md b/README.md index 85ce46d..d260c44 100644 --- a/README.md +++ b/README.md @@ -1,16 +1,25 @@ -# ComfyUI-Pixtral - For loading and running Pixtral models +# ComfyUI-PixtralLlamaVision + For loading and running Pixtral and Llama 3.2 Vision models -Includes two nodes, PixtralModelLoader and PixtralGenerateText. These should be self-explanatory. +Includes four nodes: +- PixtralModelLoader +- PixtralGenerateText +- LlamaVisionModelLoader +- LlamaVisionGenerateText + +These should be self-explanatory. -Install the latest version of transformers, which has support for Pixtral models: +Install the latest version of transformers, which has support for Pixtral/Llama Vision models: `python_embeded\python.exe -m pip install git+https://github.com/huggingface/transformers` -Requires transformers 4.45.0 + +Requires transformers 4.45.0 for Pixtral and 4.46.0 for Llama Vision. Also install bitsandbytes if you don't have it already: `python_embeded\python.exe -m pip install bitsandbytes` -You can get a 4-bit quantized version of Pixtral-12B which is compatible with these custom nodes here: -[https://huggingface.co/SeanScripts/pixtral-12b-nf4](https://huggingface.co/SeanScripts/pixtral-12b-nf4) +You can get a 4-bit quantized version of Pixtral-12B which is compatible with these custom nodes here: [https://huggingface.co/SeanScripts/pixtral-12b-nf4](https://huggingface.co/SeanScripts/pixtral-12b-nf4) + +You can get a 4-bit quantized version of Llama-3.2-11B-Vision-Instruct which is compatible with these custom nodes here: +[https://huggingface.co/SeanScripts/Llama-3.2-11B-Vision-Instruct-nf4](https://huggingface.co/SeanScripts/Llama-3.2-11B-Vision-Instruct-nf4) ![Example workflow](pixtral_caption_example.jpg) diff --git a/nodes.py b/nodes.py index b70b314..c2d75dc 100644 --- a/nodes.py +++ b/nodes.py @@ -2,7 +2,24 @@ import comfy.utils import comfy.model_management as mm import folder_paths -from transformers import LlavaForConditionalGeneration, AutoProcessor, BitsAndBytesConfig, set_seed +from transformers import AutoProcessor, BitsAndBytesConfig, set_seed + +pixtral = True +llama_vision = True +# transformers 4.45.0 +try: + from transformers import LlavaForConditionalGeneration +except ImportError: + print("[ComfyUI-PixtralLlamaVision] Can't load Pixtral, need to update transformers") + pixtral = False + +# transformers 4.46.0 +try: + from transformers import MllamaForConditionalGeneration +except ImportError: + print("[ComfyUI-PixtralLlamaVision] Can't load Llama Vision, need to update transformers") + llama_vision = False + from torchvision.transforms.functional import to_pil_image from PIL import Image import time @@ -38,6 +55,7 @@ class PixtralModelLoader: } return (pixtral_model,) + class PixtralGenerateText: @classmethod def INPUT_TYPES(s): @@ -59,10 +77,9 @@ class PixtralGenerateText: TITLE = "PixtralGenerateText" def generate_text(self, pixtral_model, images, prompt, max_new_tokens, do_sample, temperature, seed): - device = mm.get_torch_device() + device = pixtral_model['model'].device print(type(images)) - # How does batched input work? I really don't know - # Also I'm sure there is a way to do this without converting back to numpy and then PIL... + # I'm sure there is a way to do this without converting back to numpy and then PIL... # Pixtral requires PIL input for some reason, and the to_pil_image function requires channels to be the first dimension for a Tensor but the last dimension for a numpy array... Yeah idk print(f"Batch of {images.shape} images") image_list = [to_pil_image(image.numpy()) for image in images] @@ -87,14 +104,100 @@ class PixtralGenerateText: print(output) return (output,) -NODE_CLASS_MAPPINGS = { - "PixtralModelLoader": PixtralModelLoader, - "PixtralGenerateText": PixtralGenerateText, - # Not really much need to work with the image tokenization directly for something like image captioning, but might be interesting later... - #"PixtralImageEncode": PixtralImageEncode, - #"PixtralTextEncode": PixtralTextEncode, -} -NODE_DISPLAY_NAME_MAPPINGS = { - "PixtralModelLoader": "PixtralModelLoader", - "PixtralGenerateText": "PixtralGenerateText", -} + +class LlamaVisionModelLoader: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "model_name": ([item.name for item in Path(folder_paths.models_dir, "llama-vision").iterdir() if item.is_dir()],), + } + } + + RETURN_TYPES = ("LLAMA_VISION_MODEL",) + FUNCTION = "load_model" + CATEGORY = "LlamaVision" + TITLE = "LlamaVisionModelLoader" + + def load_model(self, model_name): + model_path = os.path.join(folder_paths.models_dir, "llama-vision", model_name) + device = mm.get_torch_device() + model = MllamaForConditionalGeneration.from_pretrained( + model_path, + use_safetensors=True, + device_map=device, + ) + processor = AutoProcessor.from_pretrained(model_path) + llama_vision_model = { + 'model': model, + 'processor': processor, + } + return (llama_vision_model,) + + +class LlamaVisionGenerateText: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "llama_vision_model": ("LLAMA_VISION_MODEL",), + "images": ("IMAGE",), + "prompt": ("STRING", {"default": "<|begin_of_text|><|start_header_id|>user<|end_header_id|>\nCaption this image:\n<|image|><|eot_id|><|start_header_id|>assistant<|end_header_id|>", "multiline": True}), + "max_new_tokens": ("INT", {"default": 256, "min": 1, "max": 4096}), + "do_sample": ("BOOLEAN", {"default": True}), + "temperature": ("FLOAT", {"default": 0.5}), + "seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffff}), + } + } + + RETURN_TYPES = ("STRING",) + FUNCTION = "generate_text" + CATEGORY = "LlamaVision" + TITLE = "LlamaVisionGenerateText" + + def generate_text(self, llama_vision_model, images, prompt, max_new_tokens, do_sample, temperature, seed): + device = llama_vision_model['model'].device + print(type(images)) + # I'm sure there is a way to do this without converting back to numpy and then PIL... + # Llama Vision also requires PIL input for some reason, and the to_pil_image function requires channels to be the first dimension for a Tensor but the last dimension for a numpy array... Yeah idk + print(f"Batch of {images.shape} images") + image_list = [to_pil_image(image.numpy()) for image in images] + inputs = llama_vision_model['processor'](images=image_list, text=prompt, return_tensors="pt").to(device) + prompt_tokens = len(inputs['input_ids'][0]) + print(f"Prompt tokens: {prompt_tokens}") + set_seed(seed) + t0 = time.time() + generate_ids = llama_vision_model['model'].generate( + **inputs, + max_new_tokens=max_new_tokens, + do_sample=do_sample, + temperature=temperature, + ) + t1 = time.time() + total_time = t1 - t0 + generated_tokens = len(generate_ids[0]) - prompt_tokens + time_per_token = generated_tokens/total_time + print(f"Generated {generated_tokens} tokens in {total_time:.3f} s ({time_per_token:.3f} tok/s)") + print(len(generate_ids[0][prompt_tokens:])) + output = pixtral_model['processor'].decode(generate_ids[0][prompt_tokens:], skip_special_tokens=True, clean_up_tokenization_spaces=False) + print(output) + return (output,) + +NODE_CLASS_MAPPINGS = {} + +if pixtral: + NODE_CLASS_MAPPINGS |= { + "PixtralModelLoader": PixtralModelLoader, + "PixtralGenerateText": PixtralGenerateText, + # Not really much need to work with the image tokenization directly for something like image captioning, but might be interesting later... + #"PixtralImageEncode": PixtralImageEncode, + #"PixtralTextEncode": PixtralTextEncode, + } + +if llama_vision: + NODE_CLASS_MAPPINGS |= { + "LlamaVisionModelLoader": LlamaVisionModelLoader, + "LlamaVisionGenerateText": LlamaVisionGenerateText, + } + +NODE_DISPLAY_NAME_MAPPINGS = {k:v.TITLE for k,v in NODE_CLASS_MAPPINGS.items()} \ No newline at end of file diff --git a/pyproject.toml b/pyproject.toml index 4c596f9..8f83a1f 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,14 +1,14 @@ [project] -name = "comfyui-pixtral" -description = "For loading and running Pixtral models" -version = "1.0.0" +name = "comfyui-pixtralllamavision" +description = "For loading and running Pixtral and Llama 3.2 Vision models" +version = "2.0.0" license = {file = "LICENSE"} [project.urls] -Repository = "https://github.com/SeanScripts/ComfyUI-Pixtral" +Repository = "https://github.com/SeanScripts/ComfyUI-PixtralLlamaVision" # Used by Comfy Registry https://comfyregistry.org [tool.comfy] PublisherId = "seanscripts" -DisplayName = "ComfyUI-Pixtral" +DisplayName = "ComfyUI-PixtralLlamaVision" Icon = "" diff --git a/requirements.txt b/requirements.txt index 38cb110..8f146d9 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1 +1,2 @@ +transformers >= 4.45.0 bitsandbytes