From c6d765298dd39a1c1fe8cf33e3d06ab7f2f949ac Mon Sep 17 00:00:00 2001 From: Aryan Date: Fri, 20 Jun 2025 11:32:47 +0530 Subject: [PATCH] First Commit --- __init__.py | 8 ++++ flux_kontext_max_node.py | 93 ++++++++++++++++++++++++++++++++++++++++ flux_kontext_pro_node.py | 93 ++++++++++++++++++++++++++++++++++++++++ requirements.txt | 4 ++ 4 files changed, 198 insertions(+) create mode 100644 __init__.py create mode 100644 flux_kontext_max_node.py create mode 100644 flux_kontext_pro_node.py create mode 100644 requirements.txt diff --git a/__init__.py b/__init__.py new file mode 100644 index 0000000..f0fdda8 --- /dev/null +++ b/__init__.py @@ -0,0 +1,8 @@ +from .flux_kontext_pro_node import NODE_CLASS_MAPPINGS as PRO_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as PRO_DISPLAY +from .flux_kontext_max_node import NODE_CLASS_MAPPINGS as MAX_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as MAX_DISPLAY + +# Combine both mappings +NODE_CLASS_MAPPINGS = {**PRO_MAPPINGS, **MAX_MAPPINGS} +NODE_DISPLAY_NAME_MAPPINGS = {**PRO_DISPLAY, **MAX_DISPLAY} + +__all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS'] \ No newline at end of file diff --git a/flux_kontext_max_node.py b/flux_kontext_max_node.py new file mode 100644 index 0000000..48439a4 --- /dev/null +++ b/flux_kontext_max_node.py @@ -0,0 +1,93 @@ +import replicate +import os +import requests +import torch +import numpy as np +from PIL import Image +import io + +class FluxKontextMaxNode: + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "image": ("IMAGE",), + "prompt": ("STRING", { + "multiline": True, + "default": "Make this a 90s cartoon" + }), + "replicate_api_token": ("STRING", { + "default": "your_replicate_api_token_here" + }), + "aspect_ratio": (["1:1", "16:9", "9:16", "4:3", "3:4", "3:2", "2:3", "5:4", "4:5", "21:9", "9:21", "2:1", "1:2", "match_input_image"], { + "default": "match_input_image" + }), + "output_format": (["jpg", "png"], { + "default": "jpg" + }), + "safety_tolerance": ("INT", { + "default": 2, + "min": 0, + "max": 6, + "step": 1 + }), + } + } + + RETURN_TYPES = ("IMAGE",) + RETURN_NAMES = ("image",) + FUNCTION = "generate_image" + CATEGORY = "image/generation" + + def generate_image(self, image, prompt, replicate_api_token, aspect_ratio, output_format, safety_tolerance): + try: + os.environ["REPLICATE_API_TOKEN"] = replicate_api_token + + # Convert tensor to PIL and save to buffer + tensor = image.squeeze(0) if len(image.shape) == 4 else image + if tensor.max() <= 1.0: + tensor = (tensor * 255).clamp(0, 255).byte() + pil_image = Image.fromarray(tensor.cpu().numpy(), 'RGB') + + img_buffer = io.BytesIO() + pil_image.save(img_buffer, format='PNG') + img_buffer.seek(0) + + # Run Replicate model + output = replicate.run( + "black-forest-labs/flux-kontext-max", + input={ + "prompt": prompt, + "input_image": img_buffer, + "aspect_ratio": aspect_ratio, + "output_format": output_format, + "safety_tolerance": safety_tolerance + } + ) + + # Get URL from output + output_url = output if isinstance(output, str) else (output[0] if isinstance(output, list) and output else str(output)) + + # Download and convert back to tensor + response = requests.get(output_url, timeout=30) + response.raise_for_status() + + downloaded_image = Image.open(io.BytesIO(response.content)) + if downloaded_image.mode != 'RGB': + downloaded_image = downloaded_image.convert('RGB') + + np_image = np.array(downloaded_image).astype(np.float32) / 255.0 + output_tensor = torch.from_numpy(np_image).unsqueeze(0) + + return (output_tensor,) + + except Exception as e: + return (torch.zeros((1, 512, 512, 3)),) + +NODE_CLASS_MAPPINGS = { + "FluxKontextMaxNode": FluxKontextMaxNode +} + +NODE_DISPLAY_NAME_MAPPINGS = { + "FluxKontextMaxNode": "Flux Kontext Max" +} \ No newline at end of file diff --git a/flux_kontext_pro_node.py b/flux_kontext_pro_node.py new file mode 100644 index 0000000..22466bc --- /dev/null +++ b/flux_kontext_pro_node.py @@ -0,0 +1,93 @@ +import replicate +import os +import requests +import torch +import numpy as np +from PIL import Image +import io + +class FluxKontextProNode: + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "image": ("IMAGE",), + "prompt": ("STRING", { + "multiline": True, + "default": "Make this a 90s cartoon" + }), + "replicate_api_token": ("STRING", { + "default": "your_replicate_api_token_here" + }), + "aspect_ratio": (["1:1", "16:9", "9:16", "4:3", "3:4", "3:2", "2:3", "5:4", "4:5", "21:9", "9:21", "2:1", "1:2", "match_input_image"], { + "default": "match_input_image" + }), + "output_format": (["jpg", "png"], { + "default": "jpg" + }), + "safety_tolerance": ("INT", { + "default": 2, + "min": 0, + "max": 6, + "step": 1 + }), + } + } + + RETURN_TYPES = ("IMAGE",) + RETURN_NAMES = ("image",) + FUNCTION = "generate_image" + CATEGORY = "image/generation" + + def generate_image(self, image, prompt, replicate_api_token, aspect_ratio, output_format, safety_tolerance): + try: + os.environ["REPLICATE_API_TOKEN"] = replicate_api_token + + # Convert tensor to PIL and save to buffer + tensor = image.squeeze(0) if len(image.shape) == 4 else image + if tensor.max() <= 1.0: + tensor = (tensor * 255).clamp(0, 255).byte() + pil_image = Image.fromarray(tensor.cpu().numpy(), 'RGB') + + img_buffer = io.BytesIO() + pil_image.save(img_buffer, format='PNG') + img_buffer.seek(0) + + # Run Replicate model + output = replicate.run( + "black-forest-labs/flux-kontext-pro", + input={ + "prompt": prompt, + "input_image": img_buffer, + "aspect_ratio": aspect_ratio, + "output_format": output_format, + "safety_tolerance": safety_tolerance + } + ) + + # Get URL from output + output_url = output if isinstance(output, str) else (output[0] if isinstance(output, list) and output else str(output)) + + # Download and convert back to tensor + response = requests.get(output_url, timeout=30) + response.raise_for_status() + + downloaded_image = Image.open(io.BytesIO(response.content)) + if downloaded_image.mode != 'RGB': + downloaded_image = downloaded_image.convert('RGB') + + np_image = np.array(downloaded_image).astype(np.float32) / 255.0 + output_tensor = torch.from_numpy(np_image).unsqueeze(0) + + return (output_tensor,) + + except Exception as e: + return (torch.zeros((1, 512, 512, 3)),) + +NODE_CLASS_MAPPINGS = { + "FluxKontextProNode": FluxKontextProNode +} + +NODE_DISPLAY_NAME_MAPPINGS = { + "FluxKontextProNode": "Flux Kontext Pro" +} \ No newline at end of file diff --git a/requirements.txt b/requirements.txt new file mode 100644 index 0000000..647b3e5 --- /dev/null +++ b/requirements.txt @@ -0,0 +1,4 @@ +replicate +pillow +numpy +torch \ No newline at end of file