Add support for Llama 3.2 Vision models

This commit is contained in:
SeanScripts
2024-09-25 15:51:19 -07:00
parent 33de509524
commit 923b9385dd
4 changed files with 140 additions and 27 deletions
+16 -7
View File
@@ -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)
+118 -15
View File
@@ -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()}
+5 -5
View File
@@ -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 = ""
+1
View File
@@ -1 +1,2 @@
transformers >= 4.45.0
bitsandbytes