support for batched input images
This commit is contained in:
+21
-32
@@ -7,7 +7,7 @@ from PIL import Image
|
||||
from openai import OpenAI
|
||||
|
||||
class OpenAILLMNode:
|
||||
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
@@ -24,66 +24,55 @@ class OpenAILLMNode:
|
||||
"image": ("IMAGE",),
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
RETURN_TYPES = ("STRING",)
|
||||
RETURN_NAMES = ("response",)
|
||||
FUNCTION = "generate"
|
||||
CATEGORY = "text/generation"
|
||||
|
||||
|
||||
def generate(self, prompt, model, temperature, reasoning_effort, api_key, max_output_tokens,
|
||||
system_instruction="", image=None):
|
||||
|
||||
|
||||
if not prompt.strip():
|
||||
raise ValueError("Prompt cannot be empty.")
|
||||
|
||||
|
||||
key = os.environ.get(api_key.strip(), api_key.strip()) or os.environ.get("OPENAI_API_KEY")
|
||||
if not key:
|
||||
raise ValueError("No API key provided.")
|
||||
|
||||
|
||||
client = OpenAI(api_key=key)
|
||||
|
||||
# Build input content
|
||||
|
||||
if image is not None:
|
||||
img_array = image.cpu().numpy() if isinstance(image, torch.Tensor) else image
|
||||
if len(img_array.shape) == 4:
|
||||
img_array = img_array[0]
|
||||
if img_array.dtype in [np.float32, np.float64]:
|
||||
img_array = (img_array * 255).astype(np.uint8)
|
||||
|
||||
buffered = io.BytesIO()
|
||||
Image.fromarray(img_array).save(buffered, format="PNG")
|
||||
base64_image = base64.b64encode(buffered.getvalue()).decode("utf-8")
|
||||
|
||||
input_content = [{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"type": "input_text", "text": prompt},
|
||||
{"type": "input_image", "image_url": f"data:image/png;base64,{base64_image}"}
|
||||
]
|
||||
}]
|
||||
content = [{"type": "input_text", "text": prompt}]
|
||||
for i in range(image.shape[0]):
|
||||
arr = (image[i].cpu().numpy() * 255).astype(np.uint8)
|
||||
buf = io.BytesIO()
|
||||
Image.fromarray(arr).save(buf, format="PNG")
|
||||
b64 = base64.b64encode(buf.getvalue()).decode("utf-8")
|
||||
content.append({"type": "input_image", "image_url": f"data:image/png;base64,{b64}"})
|
||||
input_content = [{"role": "user", "content": content}]
|
||||
else:
|
||||
input_content = prompt
|
||||
|
||||
# Build request
|
||||
|
||||
request_params = {
|
||||
"model": model,
|
||||
"input": input_content,
|
||||
"temperature": temperature,
|
||||
"max_output_tokens": max_output_tokens
|
||||
}
|
||||
|
||||
# Only add reasoning for o-series and gpt-5+ models
|
||||
|
||||
if not model.startswith("gpt-4"):
|
||||
request_params["reasoning"] = {"effort": reasoning_effort}
|
||||
else:
|
||||
print(f"Skipping reasoning parameter for {model} (not supported)")
|
||||
|
||||
|
||||
if system_instruction.strip():
|
||||
request_params["instructions"] = system_instruction
|
||||
|
||||
|
||||
response = client.responses.create(**request_params)
|
||||
|
||||
|
||||
return (response.output_text,)
|
||||
|
||||
|
||||
NODE_CLASS_MAPPINGS = {"OpenAILLMNode": OpenAILLMNode}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {"OpenAILLMNode": "OpenAI LLM"}
|
||||
+1
-1
@@ -1,7 +1,7 @@
|
||||
[project]
|
||||
name = "externalapi-helpers"
|
||||
description = "Various ComfyUI nodes for Gemini, Replicate and OpenAI"
|
||||
version = "1.1.3"
|
||||
version = "1.1.4"
|
||||
license = {file = "LICENSE"}
|
||||
# classifiers = [
|
||||
# # For OS-independent nodes (works on all operating systems)
|
||||
|
||||
Reference in New Issue
Block a user