From a5c70de645e0da66dfab32ea18cca1af93c2c8d4 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Fri, 9 Aug 2024 23:00:13 +0300 Subject: [PATCH] multi-image support --- nodes.py | 14 +++++++++----- 1 file changed, 9 insertions(+), 5 deletions(-) diff --git a/nodes.py b/nodes.py index 9b63fab..78f9faa 100644 --- a/nodes.py +++ b/nodes.py @@ -145,16 +145,20 @@ class LLaVA_OneVision_Run: image_processor = llava_model["image_processor"] device = llava_model["device"] dtype = llava_model["dtype"] + B, H, W, C = image.shape image = image.permute(0, 3, 1, 2) # Change shape to (B, C, H, W) transform = transforms.ToPILImage() - image_pils = [transform(image[i]) for i in range(B)] # Convert each image to PIL format image_sizes = [img.size for img in image_pils] # Get sizes for all images - image_tensors = process_images(image_pils, image_processor, model.config) # Process all images - image_tensors = [_image.to(dtype=dtype, device=device) for _image in image_tensors] # Move to appropriate device and dtype + + image_tensors = [] + for image_pil in image_pils: + processed_image = process_images([image_pil], image_processor, model.config) # Process individual image + processed_image = processed_image[0].to(dtype=dtype, device=device) # Move to appropriate device and dtype + image_tensors.append(processed_image) conv_template = "qwen_1_5" question = DEFAULT_IMAGE_TOKEN + prompt @@ -187,9 +191,9 @@ class LLaVA_OneVision_Run: mm.soft_empty_cache() text_outputs = tokenizer.batch_decode(result, skip_special_tokens=True) print(text_outputs) - + - return (text_outputs[0],) + return (text_outputs[0]) NODE_CLASS_MAPPINGS = { "DownloadAndLoadLLaVAOneVisionModel": DownloadAndLoadLLaVAOneVisionModel,