From d4462d60228f91ccd8257ad1a1b37726b6ebf2ca Mon Sep 17 00:00:00 2001 From: pvlprk Date: Tue, 7 Oct 2025 12:51:52 +0200 Subject: [PATCH] Add files via upload --- __init__.py | 3 + pvl_fal_flux_with_lora.py | 302 +++++++-- pvl_fal_flux_with_lora_pulid.py | 314 +++++++++ pvl_fal_remove_bg_v2.py | 544 +++++++++++----- pvl_google_nano_banana.py | 740 ++++++++++++--------- pvl_google_nano_banana_mandatory_img.py | 814 ++++++++++++------------ pvl_google_nano_banana_multi_img.py | 393 ++++++------ 7 files changed, 1991 insertions(+), 1119 deletions(-) create mode 100644 pvl_fal_flux_with_lora_pulid.py diff --git a/__init__.py b/__init__.py index e676818..d2a6251 100644 --- a/__init__.py +++ b/__init__.py @@ -38,6 +38,7 @@ from .pvl_any2string import PVL_Any2String from .pvl_fal_remove_bg_v2 import PVL_fal_RemoveBackground_API from .pvl_fal_depth_anything_v2 import PVL_fal_DepthAnythingV2_API from .pvl_google_nano_banana_mandatory_img import PVL_Google_NanoBanana_API_mandatory_IMG +from .pvl_fal_flux_with_lora_pulid import PVL_fal_FluxWithLoraPulID_API NODE_CLASS_MAPPINGS = { "PVL Call OpenAI Assistant": CallAssistantNode, @@ -80,6 +81,7 @@ NODE_CLASS_MAPPINGS = { "PVL_fal_RemoveBackground_API": PVL_fal_RemoveBackground_API, "PVL_fal_DepthAnythingV2_API": PVL_fal_DepthAnythingV2_API, "PVL_Google_NanoBanana_API_mandatory_IMG": PVL_Google_NanoBanana_API_mandatory_IMG, + "PVL_fal_FluxWithLoraPulID_API": PVL_fal_FluxWithLoraPulID_API, } NODE_DISPLAY_NAME_MAPPINGS = { @@ -122,4 +124,5 @@ NODE_DISPLAY_NAME_MAPPINGS = { "PVL_fal_RemoveBackground_API": "PVL Remove Background V2 (fal.ai)", "PVL_fal_DepthAnythingV2_API": "PVL Depth Anything V2 (fal.ai)", "PVL_Google_NanoBanana_API_mandatory_IMG": "PVL Google Nano-Banana API mandatory IMG", + "PVL_fal_FluxWithLoraPulID_API": "PVL Flux Lora PulID (fal.ai)", } \ No newline at end of file diff --git a/pvl_fal_flux_with_lora.py b/pvl_fal_flux_with_lora.py index 7377e45..81f99a3 100644 --- a/pvl_fal_flux_with_lora.py +++ b/pvl_fal_flux_with_lora.py @@ -1,8 +1,14 @@ import os +import re import torch +import time +import requests +import io +from concurrent.futures import ThreadPoolExecutor, as_completed from .fal_utils import FalConfig, ImageUtils, ResultProcessor, ApiHandler class PVL_fal_FluxWithLora_API: + @classmethod def INPUT_TYPES(cls): return { @@ -19,77 +25,279 @@ class PVL_fal_FluxWithLora_API: "sync_mode": ("BOOLEAN", {"default": False}), }, "optional": { + "delimiter": ("STRING", {"default": "[*]", "multiline": False, "placeholder": "Delimiter for splitting prompts (e.g., [*], \\n, |)"}), "lora1_name": ("STRING", {"default": ""}), "lora1_scale": ("FLOAT", {"default": 1.0, "min": -2.0, "max": 2.0, "step": 0.1}), "lora2_name": ("STRING", {"default": ""}), "lora2_scale": ("FLOAT", {"default": 1.0, "min": -2.0, "max": 2.0, "step": 0.1}), "lora3_name": ("STRING", {"default": ""}), "lora3_scale": ("FLOAT", {"default": 1.0, "min": -2.0, "max": 2.0, "step": 0.1}), - } + }, } - + RETURN_TYPES = ("IMAGE",) FUNCTION = "generate_image" CATEGORY = "PVL_tools_FAL" - - def generate_image(self, prompt, width, height, steps, CFG, seed, + + def _build_call_prompts(self, base_prompts, num_images): + """ + Maps prompts to calls according to the rule: + - If len(prompts) >= num_images → take first num_images + - If len(prompts) < num_images → Use each prompt in order. For remaining calls, reuse the last prompt. + """ + N = max(1, int(num_images)) + if not base_prompts: + return [] + + if len(base_prompts) >= N: + call_prompts = base_prompts[:N] + else: + print(f"[PVL WARNING] Provided {len(base_prompts)} prompts but num_images={N}. " + f"Reusing the last prompt for remaining calls.") + call_prompts = base_prompts + [base_prompts[-1]] * (N - len(base_prompts)) + + return call_prompts + + # -------- FAL Queue API - TWO PHASE EXECUTION -------- + + def _fal_submit_only(self, prompt_text, width, height, steps, CFG, seed, + enable_safety_checker, output_format, sync_mode, + lora1_name, lora1_scale, lora2_name, lora2_scale, + lora3_name, lora3_scale): + """ + Phase 1: Submit request to FAL and return request info immediately. + Does NOT wait for completion. + """ + arguments = { + "prompt": prompt_text, + "num_inference_steps": steps, + "guidance_scale": CFG, + "num_images": 1, # Each call generates 1 image + "enable_safety_checker": enable_safety_checker, + "output_format": output_format, + "sync_mode": sync_mode, + "image_size": { + "width": width, + "height": height + } + } + + if seed != -1: + arguments["seed"] = seed + + # Handle LoRAs + loras = [] + if lora1_name.strip(): + loras.append({"path": lora1_name.strip(), "scale": lora1_scale}) + if lora2_name.strip(): + loras.append({"path": lora2_name.strip(), "scale": lora2_scale}) + if lora3_name.strip(): + loras.append({"path": lora3_name.strip(), "scale": lora3_scale}) + if loras: + arguments["loras"] = loras + + # Check if ApiHandler supports async submission + if hasattr(ApiHandler, 'submit_only'): + return ApiHandler.submit_only("fal-ai/flux-lora", arguments) + else: + # Fallback to direct FAL queue API + return self._direct_fal_submit("fal-ai/flux-lora", arguments) + + def _direct_fal_submit(self, endpoint, arguments): + """Direct FAL queue API submission when ApiHandler doesn't support async.""" + fal_key = os.getenv("FAL_KEY", "") + if not fal_key: + raise RuntimeError("FAL_KEY environment variable not set") + + base = "https://queue.fal.run" + submit_url = f"{base}/{endpoint}" + headers = {"Authorization": f"Key {fal_key}"} + + r = requests.post(submit_url, headers=headers, json=arguments, timeout=120) + if not r.ok: + raise RuntimeError(f"FAL submit error {r.status_code}: {r.text}") + + sub = r.json() + req_id = sub.get("request_id") + if not req_id: + raise RuntimeError("FAL did not return a request_id") + + status_url = sub.get("status_url") or f"{base}/{endpoint}/requests/{req_id}/status" + resp_url = sub.get("response_url") or f"{base}/{endpoint}/requests/{req_id}" + + return { + "request_id": req_id, + "status_url": status_url, + "response_url": resp_url, + } + + def _fal_poll_and_fetch(self, request_info, timeout=120): + """ + Phase 2: Poll a single FAL request until complete and fetch the result. + Returns image tensor. + """ + # Check if ApiHandler supports async polling + if hasattr(ApiHandler, 'poll_and_get_result'): + result = ApiHandler.poll_and_get_result(request_info, timeout) + else: + # Fallback to direct polling + fal_key = os.getenv("FAL_KEY", "") + headers = {"Authorization": f"Key {fal_key}"} + + status_url = request_info["status_url"] + resp_url = request_info["response_url"] + + # Poll for completion + deadline = time.time() + timeout + completed = False + while time.time() < deadline: + try: + sr = requests.get(status_url, headers=headers, timeout=10) + if sr.ok and sr.json().get("status") == "COMPLETED": + completed = True + break + except Exception: + pass + time.sleep(0.6) + + if not completed: + raise RuntimeError(f"FAL request timed out after {timeout}s") + + # Fetch result + rr = requests.get(resp_url, headers=headers, timeout=15) + if not rr.ok: + raise RuntimeError(f"FAL result fetch error {rr.status_code}: {rr.text}") + + result = rr.json().get("response", rr.json()) + + # Process result using ResultProcessor + return ResultProcessor.process_image_result(result) + + def generate_image(self, prompt, width, height, steps, CFG, seed, num_images, enable_safety_checker, output_format, sync_mode, + delimiter="[*]", lora1_name="", lora1_scale=1.0, lora2_name="", lora2_scale=1.0, lora3_name="", lora3_scale=1.0): + + _t0 = time.time() + try: - # Prepare the arguments for the API call - arguments = { - "prompt": prompt, - "num_inference_steps": steps, # Using the renamed parameter - "guidance_scale": CFG, # Using the renamed parameter - "num_images": num_images, - "enable_safety_checker": enable_safety_checker, - "output_format": output_format, - "sync_mode": sync_mode, - "image_size": { # Always use custom dimensions now - "width": width, - "height": height + # Split prompts using delimiter with regex support + try: + base_prompts = [p.strip() for p in re.split(delimiter, prompt) if str(p).strip()] + except re.error: + print(f"[PVL WARNING] Invalid regex pattern '{delimiter}', using literal split.") + base_prompts = [p.strip() for p in prompt.split(delimiter) if str(p).strip()] + + if not base_prompts: + raise RuntimeError("No valid prompts provided.") + + # Map prompts to num_images calls + call_prompts = self._build_call_prompts(base_prompts, num_images) + print(f"[PVL INFO] Processing {len(call_prompts)} prompts") + + # Single call: process directly (less overhead) + if len(call_prompts) == 1: + req_info = self._fal_submit_only( + call_prompts[0], width, height, steps, CFG, seed, + enable_safety_checker, output_format, sync_mode, + lora1_name, lora1_scale, lora2_name, lora2_scale, + lora3_name, lora3_scale + ) + result = self._fal_poll_and_fetch(req_info) + + img_tensor = result[0] if isinstance(result, tuple) else result + if img_tensor.ndim == 3: + img_tensor = img_tensor.unsqueeze(0) + + _t1 = time.time() + print(f"[PVL INFO] Successfully generated 1 image in {(_t1 - _t0):.2f}s") + return (img_tensor,) + + # Multiple calls: TRUE PARALLEL execution + print(f"[PVL INFO] Submitting {len(call_prompts)} requests in parallel...") + + # PHASE 1: Submit all requests in parallel + submit_results = [] + max_workers = min(len(call_prompts), 6) # Increased from 4 to 6 + + with ThreadPoolExecutor(max_workers=max_workers) as executor: + submit_futs = { + executor.submit( + self._fal_submit_only, + call_prompts[i], width, height, steps, CFG, seed, + enable_safety_checker, output_format, sync_mode, + lora1_name, lora1_scale, lora2_name, lora2_scale, + lora3_name, lora3_scale + ): i + for i in range(len(call_prompts)) } - } + + for fut in as_completed(submit_futs): + idx = submit_futs[fut] + try: + req_info = fut.result() + submit_results.append((idx, req_info)) + except Exception as e: + print(f"[PVL ERROR] Submit failed for prompt {idx}: {e}") - # Add seed if provided (not -1) - if seed != -1: - arguments["seed"] = seed + if not submit_results: + raise RuntimeError("All FAL submission requests failed") - # Handle LoRAs - loras = [] + print(f"[PVL INFO] {len(submit_results)} requests submitted. Polling for results...") - # Add LoRA 1 if provided - if lora1_name.strip(): - loras.append({ - "path": lora1_name.strip(), - "scale": lora1_scale - }) + # PHASE 2: Poll all requests in parallel + results = {} + failed_count = 0 - # Add LoRA 2 if provided - if lora2_name.strip(): - loras.append({ - "path": lora2_name.strip(), - "scale": lora2_scale - }) + with ThreadPoolExecutor(max_workers=max_workers) as executor: + poll_futs = { + executor.submit(self._fal_poll_and_fetch, req_info): idx + for idx, req_info in submit_results + } + + for fut in as_completed(poll_futs): + idx = poll_futs[fut] + try: + result = fut.result() + results[idx] = result + except Exception as e: + failed_count += 1 + print(f"[PVL ERROR] Poll failed for prompt {idx}: {e}") - # Add LoRA 3 if provided - if lora3_name.strip(): - loras.append({ - "path": lora3_name.strip(), - "scale": lora3_scale - }) + if not results: + raise RuntimeError(f"All FAL requests failed during polling ({failed_count} failures)") - if loras: - arguments["loras"] = loras + if failed_count > 0: + print(f"[PVL WARNING] {failed_count}/{len(call_prompts)} requests failed, continuing with {len(results)} successful results") - # Submit the request and get the result - result = ApiHandler.submit_and_get_result("fal-ai/flux-lora", arguments) + # Combine all image tensors in order + all_images = [] + for i in range(len(call_prompts)): + if i in results: + result = results[i] + img_tensor = result[0] if isinstance(result, tuple) else result + + if torch.is_tensor(img_tensor): + # Handle both 3D (H,W,C) and 4D (B,H,W,C) tensors + if img_tensor.ndim == 3: + img_tensor = img_tensor.unsqueeze(0) + all_images.append(img_tensor) - # Process the result and return the image tensor - return ResultProcessor.process_image_result(result) + if not all_images: + raise RuntimeError("No images were generated from API calls") + + # Stack all images into single batch + final_tensor = torch.cat(all_images, dim=0) + + _t1 = time.time() + print(f"[PVL INFO] Successfully generated {final_tensor.shape[0]} images in {(_t1 - _t0):.2f}s") + return (final_tensor,) except Exception as e: print(f"Error generating image with FLUX: {str(e)}") - return ApiHandler.handle_image_generation_error("FLUX", e) \ No newline at end of file + return ApiHandler.handle_image_generation_error("FLUX", e) + +NODE_CLASS_MAPPINGS = {"PVL_fal_FluxWithLora_API": PVL_fal_FluxWithLora_API} +NODE_DISPLAY_NAME_MAPPINGS = {"PVL_fal_FluxWithLora_API": "PVL FAL Flux with LoRA"} diff --git a/pvl_fal_flux_with_lora_pulid.py b/pvl_fal_flux_with_lora_pulid.py new file mode 100644 index 0000000..33c5c51 --- /dev/null +++ b/pvl_fal_flux_with_lora_pulid.py @@ -0,0 +1,314 @@ +import os +import re +import torch +import time +import requests +import io +import base64 +from concurrent.futures import ThreadPoolExecutor, as_completed +from .fal_utils import FalConfig, ImageUtils, ResultProcessor, ApiHandler + +class PVL_fal_FluxWithLoraPulID_API: + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "prompt": ("STRING", {"multiline": True, "default": "a woman holding sign with glowing green text 'PuLID for FLUX'"}), + "reference_image": ("IMAGE",), + "num_images": ("INT", {"default": 1, "min": 1, "max": 4}), + "num_inference_steps": ("INT", {"default": 20, "min": 1, "max": 100}), + "guidance_scale": ("FLOAT", {"default": 4.0, "min": 1.0, "max": 20.0, "step": 0.1}), + "seed": ("INT", {"default": -1, "min": -1, "max": 4294967295}), + "enable_safety_checker": ("BOOLEAN", {"default": True}), + "sync_mode": ("BOOLEAN", {"default": False}), + }, + "optional": { + "delimiter": ("STRING", {"default": "[*]", "multiline": False, "placeholder": "Delimiter for splitting prompts (e.g., [*], \\n, |)"}), + "image_size": (["square_hd", "square", "portrait_4_3", "portrait_16_9", "landscape_4_3", "landscape_16_9"], {"default": "landscape_4_3"}), + "custom_width": ("INT", {"default": 0, "min": 0, "max": 2048, "step": 64}), + "custom_height": ("INT", {"default": 0, "min": 0, "max": 2048, "step": 64}), + "negative_prompt": ("STRING", {"multiline": True, "default": "bad quality, worst quality, text, signature, watermark, extra limbs"}), + "lora_path": ("STRING", {"default": "", "placeholder": "Optional LoRA path"}), + "lora_strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 2.0, "step": 0.1}), + "start_step": ("INT", {"default": 0, "min": 0, "max": 100}), + "true_cfg": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.1}), + "id_weight": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 5.0, "step": 0.1}), + "max_sequence_length": (["128", "256", "512"], {"default": "128"}), + }, + } + + RETURN_TYPES = ("IMAGE",) + FUNCTION = "generate_image" + CATEGORY = "PVL_tools_FAL" + + def _build_call_prompts(self, base_prompts, num_images): + """ + Maps prompts to calls according to the rule: + - If len(prompts) >= num_images → take first num_images + - If len(prompts) < num_images → Use each prompt in order. For remaining calls, reuse the last prompt. + """ + N = max(1, int(num_images)) + if not base_prompts: + return [] + + if len(base_prompts) >= N: + call_prompts = base_prompts[:N] + else: + print(f"[PVL WARNING] Provided {len(base_prompts)} prompts but num_images={N}. " + f"Reusing the last prompt for remaining calls.") + call_prompts = base_prompts + [base_prompts[-1]] * (N - len(base_prompts)) + + return call_prompts + + # -------- FAL Queue API - TWO PHASE EXECUTION -------- + + def _fal_submit_only(self, prompt_text, reference_image_url, image_size, custom_width, custom_height, + num_inference_steps, seed, guidance_scale, negative_prompt, sync_mode, + enable_safety_checker, lora_path, lora_strength, start_step, true_cfg, + id_weight, max_sequence_length): + """ + Phase 1: Submit request to FAL and return request info immediately. + Does NOT wait for completion. + """ + # Build image_size argument + if custom_width > 0 and custom_height > 0: + image_size_arg = {"width": custom_width, "height": custom_height} + else: + image_size_arg = image_size + + arguments = { + "prompt": prompt_text, + "reference_image_url": reference_image_url, + "image_size": image_size_arg, + "num_inference_steps": num_inference_steps, + "guidance_scale": guidance_scale, + "negative_prompt": negative_prompt, + "sync_mode": sync_mode, + "enable_safety_checker": enable_safety_checker, + "true_cfg": true_cfg, + "id_weight": id_weight, + "max_sequence_length": max_sequence_length, + } + + if seed != -1: + arguments["seed"] = seed + if lora_path.strip(): + arguments["lora_path"] = lora_path.strip() + arguments["lora_strength"] = lora_strength + if start_step > 0: + arguments["start_step"] = start_step + + # Check if ApiHandler supports async submission + if hasattr(ApiHandler, 'submit_only'): + return ApiHandler.submit_only("fal-ai/flux-pulid-lora", arguments) + else: + # Fallback to direct FAL queue API + return self._direct_fal_submit("fal-ai/flux-pulid-lora", arguments) + + def _direct_fal_submit(self, endpoint, arguments): + """Direct FAL queue API submission when ApiHandler doesn't support async.""" + fal_key = os.getenv("FAL_KEY", "") + if not fal_key: + raise RuntimeError("FAL_KEY environment variable not set") + + base = "https://queue.fal.run" + submit_url = f"{base}/{endpoint}" + headers = {"Authorization": f"Key {fal_key}"} + + r = requests.post(submit_url, headers=headers, json=arguments, timeout=120) + if not r.ok: + raise RuntimeError(f"FAL submit error {r.status_code}: {r.text}") + + sub = r.json() + req_id = sub.get("request_id") + if not req_id: + raise RuntimeError("FAL did not return a request_id") + + status_url = sub.get("status_url") or f"{base}/{endpoint}/requests/{req_id}/status" + resp_url = sub.get("response_url") or f"{base}/{endpoint}/requests/{req_id}" + + return { + "request_id": req_id, + "status_url": status_url, + "response_url": resp_url, + } + + def _fal_poll_and_fetch(self, request_info, timeout=120): + """ + Phase 2: Poll a single FAL request until complete and fetch the result. + Returns image tensor. + """ + # Check if ApiHandler supports async polling + if hasattr(ApiHandler, 'poll_and_get_result'): + result = ApiHandler.poll_and_get_result(request_info, timeout) + else: + # Fallback to direct polling + fal_key = os.getenv("FAL_KEY", "") + headers = {"Authorization": f"Key {fal_key}"} + + status_url = request_info["status_url"] + resp_url = request_info["response_url"] + + # Poll for completion + deadline = time.time() + timeout + completed = False + while time.time() < deadline: + try: + sr = requests.get(status_url, headers=headers, timeout=10) + if sr.ok and sr.json().get("status") == "COMPLETED": + completed = True + break + except Exception: + pass + time.sleep(0.6) + + if not completed: + raise RuntimeError(f"FAL request timed out after {timeout}s") + + # Fetch result + rr = requests.get(resp_url, headers=headers, timeout=15) + if not rr.ok: + raise RuntimeError(f"FAL result fetch error {rr.status_code}: {rr.text}") + + result = rr.json().get("response", rr.json()) + + # Process result using ResultProcessor + return ResultProcessor.process_image_result(result) + + def generate_image(self, prompt, reference_image, num_images, num_inference_steps, + guidance_scale, seed, enable_safety_checker, sync_mode, + delimiter="[*]", + image_size="landscape_4_3", custom_width=0, custom_height=0, + negative_prompt="bad quality, worst quality, text, signature, watermark, extra limbs", + lora_path="", lora_strength=1.0, start_step=0, true_cfg=1.0, + id_weight=1.0, max_sequence_length="128"): + + _t0 = time.time() + + try: + # Upload reference image once (shared for all calls) + print("[PVL INFO] Uploading reference image to fal.ai storage...") + reference_image_url = ImageUtils.upload_image(reference_image) + print(f"[PVL INFO] Reference image uploaded: {reference_image_url}") + + # Split prompts using delimiter with regex support + try: + base_prompts = [p.strip() for p in re.split(delimiter, prompt) if str(p).strip()] + except re.error: + print(f"[PVL WARNING] Invalid regex pattern '{delimiter}', using literal split.") + base_prompts = [p.strip() for p in prompt.split(delimiter) if str(p).strip()] + + if not base_prompts: + raise RuntimeError("No valid prompts provided.") + + # Map prompts to num_images calls + call_prompts = self._build_call_prompts(base_prompts, num_images) + print(f"[PVL INFO] Processing {len(call_prompts)} prompts") + + # Single call: process directly (less overhead) + if len(call_prompts) == 1: + req_info = self._fal_submit_only( + call_prompts[0], reference_image_url, image_size, custom_width, custom_height, + num_inference_steps, seed, guidance_scale, negative_prompt, sync_mode, + enable_safety_checker, lora_path, lora_strength, start_step, true_cfg, + id_weight, max_sequence_length + ) + result = self._fal_poll_and_fetch(req_info) + + img_tensor = result[0] if isinstance(result, tuple) else result + if img_tensor.ndim == 3: + img_tensor = img_tensor.unsqueeze(0) + + _t1 = time.time() + print(f"[PVL INFO] Successfully generated 1 image in {(_t1 - _t0):.2f}s") + return (img_tensor,) + + # Multiple calls: TRUE PARALLEL execution + print(f"[PVL INFO] Submitting {len(call_prompts)} requests in parallel...") + + # PHASE 1: Submit all requests in parallel + submit_results = [] + max_workers = min(len(call_prompts), 6) + + with ThreadPoolExecutor(max_workers=max_workers) as executor: + submit_futs = { + executor.submit( + self._fal_submit_only, + call_prompts[i], reference_image_url, image_size, custom_width, custom_height, + num_inference_steps, seed, guidance_scale, negative_prompt, sync_mode, + enable_safety_checker, lora_path, lora_strength, start_step, true_cfg, + id_weight, max_sequence_length + ): i + for i in range(len(call_prompts)) + } + + for fut in as_completed(submit_futs): + idx = submit_futs[fut] + try: + req_info = fut.result() + submit_results.append((idx, req_info)) + except Exception as e: + print(f"[PVL ERROR] Submit failed for prompt {idx}: {e}") + + if not submit_results: + raise RuntimeError("All FAL submission requests failed") + + print(f"[PVL INFO] {len(submit_results)} requests submitted. Polling for results...") + + # PHASE 2: Poll all requests in parallel + results = {} + failed_count = 0 + + with ThreadPoolExecutor(max_workers=max_workers) as executor: + poll_futs = { + executor.submit(self._fal_poll_and_fetch, req_info): idx + for idx, req_info in submit_results + } + + for fut in as_completed(poll_futs): + idx = poll_futs[fut] + try: + result = fut.result() + results[idx] = result + except Exception as e: + failed_count += 1 + print(f"[PVL ERROR] Poll failed for prompt {idx}: {e}") + + if not results: + raise RuntimeError(f"All FAL requests failed during polling ({failed_count} failures)") + + if failed_count > 0: + print(f"[PVL WARNING] {failed_count}/{len(call_prompts)} requests failed, continuing with {len(results)} successful results") + + # Combine all image tensors in order + all_images = [] + for i in range(len(call_prompts)): + if i in results: + result = results[i] + img_tensor = result[0] if isinstance(result, tuple) else result + + if torch.is_tensor(img_tensor): + # Handle both 3D (H,W,C) and 4D (B,H,W,C) tensors + if img_tensor.ndim == 3: + img_tensor = img_tensor.unsqueeze(0) + all_images.append(img_tensor) + + if not all_images: + raise RuntimeError("No images were generated from API calls") + + # Stack all images into single batch + final_tensor = torch.cat(all_images, dim=0) + + _t1 = time.time() + print(f"[PVL INFO] Successfully generated {final_tensor.shape[0]} images in {(_t1 - _t0):.2f}s") + return (final_tensor,) + + except Exception as e: + print(f"Error generating image with FLUX PuLID LoRA: {str(e)}") + # Fallback error handling - return empty tensor + empty_tensor = torch.zeros((1, 64, 64, 3), dtype=torch.float32) + return (empty_tensor,) + +NODE_CLASS_MAPPINGS = {"PVL_fal_FluxWithLoraPulID_API": PVL_fal_FluxWithLoraPulID_API} +NODE_DISPLAY_NAME_MAPPINGS = {"PVL_fal_FluxWithLoraPulID_API": "PVL Flux Lora PulID (fal.ai)"} diff --git a/pvl_fal_remove_bg_v2.py b/pvl_fal_remove_bg_v2.py index a5d4507..55b50a9 100644 --- a/pvl_fal_remove_bg_v2.py +++ b/pvl_fal_remove_bg_v2.py @@ -1,168 +1,376 @@ -import torch -import numpy as np - -from .fal_utils import ImageUtils, ApiHandler - - -class PVL_fal_RemoveBackground_API: - """ - ComfyUI node for FAL 'fal-ai/birefnet/v2' — Remove Background V2. - - Inputs: - - image (IMAGE): One input image (batch NOT supported). - - model (CHOICE): Which model variant to use. - - operating_resolution (CHOICE): Resolution for inference ("1024x1024" or "2048x2048"). - - output_format (CHOICE): "png" or "webp". - - output_mask (BOOLEAN): Whether to also return the mask. - - refine_foreground (BOOLEAN): Whether to refine the foreground (default: True). - - Outputs: - - IMAGE: Foreground with background removed (RGB or RGBA if provided). - - MASK: Optional mask (1-channel, float, shape [1,H,W]) if output_mask=True. - """ - - @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "image": ("IMAGE",), - "model": ( - [ - "General Use (Light)", - "General Use (Light 2K)", - "General Use (Heavy)", - "Matting", - "Portrait", - ], - {"default": "General Use (Light)"}, - ), - "operating_resolution": ( - ["1024x1024", "2048x2048"], - {"default": "1024x1024"}, - ), - "output_format": ( - ["png", "webp"], - {"default": "png"}, - ), - "output_mask": ("BOOLEAN", {"default": False}), - "refine_foreground": ("BOOLEAN", {"default": True}), - } - } - - RETURN_TYPES = ("IMAGE", "MASK") - RETURN_NAMES = ("foreground", "mask") - FUNCTION = "remove_background" - CATEGORY = "PVL_tools" - - # ------------------------- helpers ------------------------- - def _raise(self, msg: str): - raise RuntimeError(msg) - - def _split_image_batch(self, image): - import numpy as np - if isinstance(image, torch.Tensor): - if image.ndim == 4: - return [image[i] for i in range(image.shape[0])] - elif image.ndim == 3: - return [image] - else: - self._raise("FAL: unsupported image tensor dimensionality.") - elif isinstance(image, np.ndarray): - t = torch.from_numpy(image) - return self._split_image_batch(t) - else: - self._raise("FAL: unsupported image type (expected torch Tensor or numpy.ndarray).") - - def _download_pil(self, url, mode=None): - import requests, io - from PIL import Image - resp = requests.get(url, timeout=120) - resp.raise_for_status() - pil = Image.open(io.BytesIO(resp.content)) - if mode is not None: - pil = pil.convert(mode) - return pil - - # ------------------------- main ------------------------- - def remove_background( - self, - image, - model, - operating_resolution, - output_format, - output_mask, - refine_foreground, - ): - # Split batch into frames - frames = self._split_image_batch(image) - if not frames: - self._raise("FAL: no input image frames provided.") - if len(frames) > 1: - self._raise("FAL: batch >1 not supported for Remove Background API.") - - # Inline image as base64 data URI (avoids storage upload) - image_url = ImageUtils.image_to_data_uri(frames[0]) - if not image_url: - self._raise("FAL: failed to convert input image.") - - arguments = { - "image_url": image_url, - "model": model, - "operating_resolution": operating_resolution, - "output_format": output_format, - "output_mask": bool(output_mask), - "refine_foreground": bool(refine_foreground), - } - - # Submit request - result = ApiHandler.submit_and_get_result("fal-ai/birefnet/v2", arguments) - if not isinstance(result, dict): - self._raise("FAL: unexpected response type (expected dict).") - - # ---- Foreground image (keep alpha if present) ---- - if "image" not in result or not isinstance(result["image"], dict): - self._raise("FAL: response missing foreground image.") - fg_url = result["image"].get("url") - if not fg_url: - self._raise("FAL: foreground image has no URL.") - - pil_fg = self._download_pil(fg_url, mode=None) # keep native mode - fg_arr = np.array(pil_fg).astype(np.float32) / 255.0 # (H,W,3) or (H,W,4) - if fg_arr.ndim == 2: # grayscale fallback - fg_arr = np.expand_dims(fg_arr, axis=-1) - fg_tensor = torch.from_numpy(fg_arr).unsqueeze(0) # (1,H,W,C) - - # ---- Mask (shape [1,H,W], float 0..1) ---- - mask_tensor = torch.zeros((1, fg_arr.shape[0], fg_arr.shape[1]), dtype=torch.float32) - - if bool(output_mask): - mask_url = None - if isinstance(result.get("mask_image"), dict): - mask_url = result["mask_image"].get("url") - - if mask_url: - pil_mask_rgba = self._download_pil(mask_url, mode=None) # keep native - if "A" in pil_mask_rgba.getbands(): - alpha = pil_mask_rgba.getchannel("A") - mask_arr = np.array(alpha).astype(np.float32) / 255.0 # (H,W) - else: - pil_mask_L = pil_mask_rgba.convert("L") - mask_arr = np.array(pil_mask_L).astype(np.float32) / 255.0 - mask_tensor = torch.from_numpy(mask_arr).unsqueeze(0) # (1,H,W) - else: - # fallback: try to read alpha channel from foreground if present - if "A" in pil_fg.getbands(): - alpha = pil_fg.getchannel("A") - mask_arr = np.array(alpha).astype(np.float32) / 255.0 - mask_tensor = torch.from_numpy(mask_arr).unsqueeze(0) - - return (fg_tensor, mask_tensor) - -# ---- ComfyUI discovery ---- -NODE_CLASS_MAPPINGS = { - "PVL_fal_RemoveBackground_API": PVL_fal_RemoveBackground_API, -} - -NODE_DISPLAY_NAME_MAPPINGS = { - "PVL_fal_RemoveBackground_API": "PVL Remove Background V2 (fal.ai)", -} +import torch +import numpy as np +import time +import requests +import io +from concurrent.futures import ThreadPoolExecutor, as_completed +from .fal_utils import ImageUtils, ApiHandler + +class PVL_fal_RemoveBackground_API: + """ + ComfyUI node for FAL 'fal-ai/birefnet/v2' — Remove Background V2. + + Features: + - TRUE PARALLEL execution: submits all requests first, then polls all in parallel + - Error handling for individual requests with partial results support + - Batch processing with worker cap to prevent thread explosion + - Optional sync_mode toggle + + Inputs: + - image (IMAGE): Input image(s). Batch processing is supported with parallel API calls. + - model (CHOICE): Which model variant to use. + - operating_resolution (CHOICE): Resolution for inference ("1024x1024" or "2048x2048"). + - output_format (CHOICE): "png" or "webp". + - output_mask (BOOLEAN): Whether to also return the mask. + - refine_foreground (BOOLEAN): Whether to refine the foreground (default: True). + - sync_mode (BOOLEAN): Use synchronous mode for FAL API (default: False). + + Outputs: + - IMAGE: Foreground with background removed (RGB or RGBA if provided). + - MASK: Optional mask (1-channel, float, shape [B,H,W]) if output_mask=True. + """ + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "image": ("IMAGE",), + "model": ( + [ + "General Use (Light)", + "General Use (Light 2K)", + "General Use (Heavy)", + "Matting", + "Portrait", + ], + {"default": "General Use (Light)"}, + ), + "operating_resolution": ( + ["1024x1024", "2048x2048"], + {"default": "1024x1024"}, + ), + "output_format": ( + ["png", "webp"], + {"default": "png"}, + ), + "output_mask": ("BOOLEAN", {"default": False}), + "refine_foreground": ("BOOLEAN", {"default": True}), + "sync_mode": ("BOOLEAN", {"default": False}), + }, + } + + RETURN_TYPES = ("IMAGE", "MASK") + RETURN_NAMES = ("foreground", "mask") + FUNCTION = "remove_background" + CATEGORY = "PVL_tools" + + # ------------------------- helpers ------------------------- + + def _raise(self, msg: str): + raise RuntimeError(msg) + + def _split_image_batch(self, image): + if isinstance(image, torch.Tensor): + if image.ndim == 4: + return [image[i] for i in range(image.shape[0])] + elif image.ndim == 3: + return [image] + else: + self._raise("FAL: unsupported image tensor dimensionality.") + elif isinstance(image, np.ndarray): + t = torch.from_numpy(image) + return self._split_image_batch(t) + else: + self._raise("FAL: unsupported image type (expected torch Tensor or numpy.ndarray).") + + def _download_pil(self, url, mode=None): + from PIL import Image + resp = requests.get(url, timeout=120) + resp.raise_for_status() + pil = Image.open(io.BytesIO(resp.content)) + if mode is not None: + pil = pil.convert(mode) + return pil + + # -------- FAL Queue API - TWO PHASE EXECUTION -------- + + def _fal_submit_only(self, frame, model, operating_resolution, output_format, + output_mask, refine_foreground, sync_mode): + """ + Phase 1: Submit request to FAL and return request info immediately. + Does NOT wait for completion. + """ + # Convert image to data URI + image_url = ImageUtils.image_to_data_uri(frame) + if not image_url: + raise RuntimeError("FAL: failed to convert input image.") + + arguments = { + "image_url": image_url, + "model": model, + "operating_resolution": operating_resolution, + "output_format": output_format, + "output_mask": bool(output_mask), + "refine_foreground": bool(refine_foreground), + } + + # Check if ApiHandler supports async submission + # If it has submit_only method, use it; otherwise use direct API call + if hasattr(ApiHandler, 'submit_only'): + return ApiHandler.submit_only("fal-ai/birefnet/v2", arguments, sync_mode) + else: + # Fallback to direct FAL queue API + return self._direct_fal_submit("fal-ai/birefnet/v2", arguments, sync_mode) + + def _direct_fal_submit(self, endpoint, arguments, sync_mode): + """Direct FAL queue API submission when ApiHandler doesn't support async.""" + import os + + fal_key = os.getenv("FAL_KEY", "") + if not fal_key: + raise RuntimeError("FAL_KEY environment variable not set") + + base = "https://queue.fal.run" + submit_url = f"{base}/{endpoint}" + headers = {"Authorization": f"Key {fal_key}"} + + payload = dict(arguments) + payload["sync_mode"] = sync_mode + + r = requests.post(submit_url, headers=headers, json=payload, timeout=120) + if not r.ok: + raise RuntimeError(f"FAL submit error {r.status_code}: {r.text}") + + sub = r.json() + req_id = sub.get("request_id") + if not req_id: + raise RuntimeError("FAL did not return a request_id") + + status_url = sub.get("status_url") or f"{base}/{endpoint}/requests/{req_id}/status" + resp_url = sub.get("response_url") or f"{base}/{endpoint}/requests/{req_id}" + + return { + "request_id": req_id, + "status_url": status_url, + "response_url": resp_url, + } + + def _fal_poll_and_fetch(self, request_info, timeout=120): + """ + Phase 2: Poll a single FAL request until complete and fetch the result. + Returns (fg_tensor, mask_tensor). + """ + import os + + fal_key = os.getenv("FAL_KEY", "") + headers = {"Authorization": f"Key {fal_key}"} + + # Check if ApiHandler supports async polling + if hasattr(ApiHandler, 'poll_and_get_result'): + result = ApiHandler.poll_and_get_result(request_info, timeout) + else: + # Fallback to direct polling + status_url = request_info["status_url"] + resp_url = request_info["response_url"] + + # Poll for completion + deadline = time.time() + timeout + completed = False + while time.time() < deadline: + try: + sr = requests.get(status_url, headers=headers, timeout=10) + if sr.ok and sr.json().get("status") == "COMPLETED": + completed = True + break + except Exception: + pass + time.sleep(0.6) + + if not completed: + raise RuntimeError(f"FAL request timed out after {timeout}s") + + # Fetch result + rr = requests.get(resp_url, headers=headers, timeout=15) + if not rr.ok: + raise RuntimeError(f"FAL result fetch error {rr.status_code}: {rr.text}") + + rdata = rr.json() + result = rdata.get("response", rdata) + + # Process result + return self._process_result(result) + + def _process_result(self, result): + """Process FAL API result and return tensors.""" + if not isinstance(result, dict): + raise RuntimeError("FAL: unexpected response type (expected dict).") + + # ---- Foreground image (keep alpha if present) ---- + if "image" not in result or not isinstance(result["image"], dict): + raise RuntimeError("FAL: response missing foreground image.") + + fg_url = result["image"].get("url") + if not fg_url: + raise RuntimeError("FAL: foreground image has no URL.") + + pil_fg = self._download_pil(fg_url, mode=None) # keep native mode + fg_arr = np.array(pil_fg).astype(np.float32) / 255.0 # (H,W,3) or (H,W,4) + + if fg_arr.ndim == 2: # grayscale fallback + fg_arr = np.expand_dims(fg_arr, axis=-1) + + fg_tensor = torch.from_numpy(fg_arr) # (H,W,C) + + # ---- Mask (shape [H,W], float 0..1) ---- + mask_tensor = torch.zeros((fg_arr.shape[0], fg_arr.shape[1]), dtype=torch.float32) + + # Try to get mask from response + mask_url = None + if isinstance(result.get("mask_image"), dict): + mask_url = result["mask_image"].get("url") + + if mask_url: + pil_mask_rgba = self._download_pil(mask_url, mode=None) + if "A" in pil_mask_rgba.getbands(): + alpha = pil_mask_rgba.getchannel("A") + mask_arr = np.array(alpha).astype(np.float32) / 255.0 + else: + pil_mask_L = pil_mask_rgba.convert("L") + mask_arr = np.array(pil_mask_L).astype(np.float32) / 255.0 + mask_tensor = torch.from_numpy(mask_arr) + else: + # Fallback: extract alpha channel from foreground if present + if "A" in pil_fg.getbands(): + alpha = pil_fg.getchannel("A") + mask_arr = np.array(alpha).astype(np.float32) / 255.0 + mask_tensor = torch.from_numpy(mask_arr) + + return fg_tensor, mask_tensor + + # ------------------------- main ------------------------- + + def remove_background( + self, + image, + model, + operating_resolution, + output_format, + output_mask, + refine_foreground, + sync_mode=False, + ): + _t0 = time.time() + + # Split batch into frames + frames = self._split_image_batch(image) + if not frames: + self._raise("FAL: no input image frames provided.") + + batch_size = len(frames) + print(f"[FAL RemoveBG] Processing {batch_size} images...") + + # Single image: process directly (less overhead) + if batch_size == 1: + try: + req_info = self._fal_submit_only( + frames[0], model, operating_resolution, output_format, + output_mask, refine_foreground, sync_mode + ) + fg_tensor, mask_tensor = self._fal_poll_and_fetch(req_info) + + _t1 = time.time() + print(f"[FAL RemoveBG] Completed in {(_t1 - _t0):.2f}s") + return (fg_tensor.unsqueeze(0), mask_tensor.unsqueeze(0)) + except Exception as e: + self._raise(f"FAL: processing failed: {e}") + + # Multiple images: TRUE PARALLEL execution + print(f"[FAL RemoveBG] Submitting {batch_size} requests in parallel...") + + # PHASE 1: Submit all requests in parallel + submit_results = [] + max_workers = min(batch_size, 6) # Cap at 6 workers + + with ThreadPoolExecutor(max_workers=max_workers) as executor: + submit_futs = { + executor.submit( + self._fal_submit_only, + frames[i], + model, + operating_resolution, + output_format, + output_mask, + refine_foreground, + sync_mode + ): i + for i in range(batch_size) + } + + for fut in as_completed(submit_futs): + idx = submit_futs[fut] + try: + req_info = fut.result() + submit_results.append((idx, req_info)) + except Exception as e: + print(f"[FAL RemoveBG] Submit failed for image {idx}: {e}") + + if not submit_results: + self._raise("All FAL submission requests failed") + + print(f"[FAL RemoveBG] {len(submit_results)} requests submitted. Polling for results...") + + # PHASE 2: Poll all requests in parallel + results = {} + failed_count = 0 + + with ThreadPoolExecutor(max_workers=max_workers) as executor: + poll_futs = { + executor.submit(self._fal_poll_and_fetch, req_info): idx + for idx, req_info in submit_results + } + + for fut in as_completed(poll_futs): + idx = poll_futs[fut] + try: + fg_tensor, mask_tensor = fut.result() + results[idx] = (fg_tensor, mask_tensor) + except Exception as e: + failed_count += 1 + print(f"[FAL RemoveBG] Poll failed for image {idx}: {e}") + + if not results: + self._raise(f"All FAL requests failed during polling ({failed_count} failures)") + + if failed_count > 0: + print(f"[FAL RemoveBG WARNING] {failed_count}/{batch_size} requests failed, continuing with {len(results)} successful results") + + # Stack results in original order (with None for failed images) + fg_list = [] + mask_list = [] + + for i in range(batch_size): + if i in results: + fg_tensor, mask_tensor = results[i] + fg_list.append(fg_tensor) + mask_list.append(mask_tensor) + + if not fg_list: + self._raise("No successful results to return") + + # Stack into batched tensors + fg_batch = torch.stack(fg_list, dim=0) # (B,H,W,C) + mask_batch = torch.stack(mask_list, dim=0) # (B,H,W) + + _t1 = time.time() + print(f"[FAL RemoveBG] Completed {len(fg_list)}/{batch_size} images in {(_t1 - _t0):.2f}s") + + return (fg_batch, mask_batch) + +# ---- ComfyUI discovery ---- +NODE_CLASS_MAPPINGS = { + "PVL_fal_RemoveBackground_API": PVL_fal_RemoveBackground_API, +} + +NODE_DISPLAY_NAME_MAPPINGS = { + "PVL_fal_RemoveBackground_API": "PVL Remove Background V2 (fal.ai) [Batch]", +} diff --git a/pvl_google_nano_banana.py b/pvl_google_nano_banana.py index 866301a..6b8533b 100644 --- a/pvl_google_nano_banana.py +++ b/pvl_google_nano_banana.py @@ -1,21 +1,33 @@ - # pvl_google_nano_banana.py -# Node: PVL Google Nano-Banana API (SDK + robust FAL fallback) + +# Node: PVL Google Nano-Banana API (Gemini + FAL) — with delimiter support + parallel prompts + # Author: PVL # License: MIT -# + # Requires: -# pip install google-genai -# -# Changes in this build: -# - FAL fallback now sends ONLY "image_urls" (list) to avoid duplicate outputs from some routes. -# - Enforces num_images: if num_images==1 returns the FIRST image only; otherwise caps to num_images. -# - Keeps sync_mode true and robust response parsing; prints raw FAL JSON when debug is ON. -# - Return order: ("STRING","IMAGE") => (text, images). +# pip install google-genai -import os, io, json, base64, typing as T, time +# Features: +# - Regex-based delimiter input for splitting prompts (default: [*]) +# - Parallel API calls for multiple prompts +# - TRUE PARALLEL FAL execution: submit all requests first, then poll for results +# - If prompts < num_images, reuses the last prompt to fill remaining calls +# - Single optional ComfyUI IMAGE input (can be batched). +# - aspect_ratio supported for both Google and FAL. +# - use_fal_fallback default True (fallback to FAL when Google returns no images). +# - force_fal toggle to always use FAL (bypass Google). +# - Dual FAL routes: img2img (edit) vs txt2img (generate) chosen automatically. +# - Individual request error handling with partial results support. + +import os +import io +import json +import base64 +import typing as T +import time +import re from concurrent.futures import ThreadPoolExecutor, as_completed - import requests import numpy as np from PIL import Image @@ -24,23 +36,24 @@ import torch NODE_NAME = "PVL Google Nano-Banana API" NODE_CATEGORY = "PVL/Google" DEFAULT_MODEL = "gemini-2.5-flash-image-preview" - _TOP_P = 0.95 _TOP_K = 64 _MAX_TOKENS = 4096 -_VALID_ASPECTS = {"21:9","1:1","4:3","3:2","2:3","5:4","4:5","3:4","16:9","9:16"} +_VALID_ASPECTS = {"21:9", "1:1", "4:3", "3:2", "2:3", "5:4", "4:5", "3:4", "16:9", "9:16"} + +# ----------------- image helpers ----------------- -# --- image helpers --- def pil_to_tensor(img: Image.Image) -> torch.Tensor: if img.mode != "RGB": img = img.convert("RGB") arr = np.asarray(img, dtype=np.float32) / 255.0 - return torch.from_numpy(arr)[None, ...] + t = torch.from_numpy(arr)[None, ...] + return t def tensor_to_pil(t: torch.Tensor) -> Image.Image: if t.ndim == 4: t = t[0] - arr = (t.clamp(0,1).cpu().numpy() * 255).astype("uint8") + arr = (t.clamp(0, 1).cpu().numpy() * 255).astype("uint8") return Image.fromarray(arr, "RGB") def encode_pil_bytes(img: Image.Image, mime: str) -> bytes: @@ -51,6 +64,24 @@ def encode_pil_bytes(img: Image.Image, mime: str) -> bytes: img.save(buf, format="PNG") return buf.getvalue() +def stack_images_same_size(tensors: T.List[torch.Tensor], debug: bool = False) -> torch.Tensor: + """Concatenate (B,H,W,C) batches along B. If shapes mismatch, resize to the first image size.""" + if not tensors: + raise RuntimeError("No images to stack.") + + try: + return torch.cat(tensors, dim=0) + except RuntimeError: + if debug: + print("[PVL NODE] Mismatched sizes, resizing to match first image.") + target_h, target_w = tensors[0].shape[1], tensors[0].shape[2] + fixed = [] + for t in tensors: + pil = tensor_to_pil(t) + rp = pil.resize((target_w, target_h), Image.LANCZOS) + fixed.append(pil_to_tensor(rp)) + return torch.cat(fixed, dim=0) + def _extract_image_bytes_from_part(part) -> T.Optional[bytes]: try: inline = getattr(part, "inline_data", None) or getattr(part, "inlineData", None) @@ -65,13 +96,14 @@ def _extract_image_bytes_from_part(part) -> T.Optional[bytes]: return None except Exception: pass + if isinstance(part, dict): + blob = None if "inline_data" in part and isinstance(part["inline_data"], dict): blob = part["inline_data"].get("data") elif "inlineData" in part and isinstance(part["inlineData"], dict): blob = part["inlineData"].get("data") - else: - blob = None + if isinstance(blob, (bytes, bytearray)): return bytes(blob) if isinstance(blob, str): @@ -79,6 +111,7 @@ def _extract_image_bytes_from_part(part) -> T.Optional[bytes]: return base64.b64decode(blob, validate=False) except Exception: return None + return None def _extract_text_from_part(part) -> T.Optional[str]: @@ -88,89 +121,99 @@ def _extract_text_from_part(part) -> T.Optional[str]: return str(txt) except Exception: pass + if isinstance(part, dict) and "text" in part and part["text"] is not None: return str(part["text"]) + return None def _data_url(mime: str, raw: bytes) -> str: return f"data:{mime};base64," + base64.b64encode(raw).decode("utf-8") class PVL_Google_NanoBanana_API: + @classmethod def INPUT_TYPES(cls): return { "required": { "prompt": ("STRING", {"multiline": True, "default": "A tiny banana spaceship over a neon city."}), + "delimiter": ("STRING", {"default": "[*]", "multiline": False, "placeholder": "Regex delimiter e.g. [*], \\n, |"}), }, "optional": { "images": ("IMAGE",), - "aspect_ratio": ("STRING", {"default": "1:1", "placeholder": "e.g. 16:9, 9:16, 3:2 — works for both Google & FAL"}), + "aspect_ratio": ("STRING", {"default": "1:1", "placeholder": "e.g. 16:9, 9:16, 3:2"}), "model": ("STRING", {"default": DEFAULT_MODEL}), "endpoint_override": ("STRING", {"default": ""}), - "api_key": ("STRING", {"default": "", "multiline": False, - "placeholder": "Leave empty to use GEMINI_API_KEY"}), + "api_key": ("STRING", {"default": "", "multiline": False, "placeholder": "Leave empty to use GEMINI_API_KEY"}), "temperature": ("FLOAT", {"default": 0.6, "min": 0.0, "max": 2.0, "step": 0.05}), - "output_format": (["png","jpeg"], {"default": "png"}), + "output_format": (["png", "jpeg"], {"default": "png"}), "capture_text_output": ("BOOLEAN", {"default": False}), "num_images": ("INT", {"default": 1, "min": 1, "max": 12, "step": 1}), "timeout_sec": ("INT", {"default": 120, "min": 5, "max": 600, "step": 5}), "request_id": ("STRING", {"default": ""}), "debug_log": ("BOOLEAN", {"default": False}), - - # --- FAL fallback --- - "use_fal_fallback": ("BOOLEAN", {"default": False}), - "fal_api_key": ("STRING", {"default": "", "multiline": False, - "placeholder": "Leave empty to use FAL_KEY"}), - "fal_route": ("STRING", {"default": "fal-ai/nano-banana/edit"}), + # FAL flags + "use_fal_fallback": ("BOOLEAN", {"default": True}), + "force_fal": ("BOOLEAN", {"default": False}), + "sync_mode": ("BOOLEAN", {"default": False}), + "fal_api_key": ("STRING", {"default": "", "multiline": False, "placeholder": "Leave empty to use FAL_KEY"}), + "fal_route_img2img": ("STRING", {"default": "fal-ai/nano-banana/edit"}), + "fal_route_txt2img": ("STRING", {"default": "fal-ai/nano-banana"}), } } - - RETURN_TYPES = ("STRING","IMAGE",) - RETURN_NAMES = ("text","images") + + RETURN_TYPES = ("STRING", "IMAGE",) + RETURN_NAMES = ("text", "images") FUNCTION = "run" CATEGORY = NODE_CATEGORY - - # ---- helpers ---- + + # -------- helpers -------- + def _make_client(self, api_key: str, endpoint_override: str): try: from google import genai from google.genai import types except Exception as e: raise RuntimeError("Google GenAI SDK not installed. Run: pip install google-genai") from e - + http_options = None if endpoint_override.strip(): try: http_options = types.HttpOptions(base_url=endpoint_override.strip()) except Exception: http_options = None - + if http_options is not None: client = genai.Client(api_key=api_key, http_options=http_options) else: client = genai.Client(api_key=api_key) + return client - + def _build_parts(self, prompt: str, images: T.Optional[torch.Tensor], mime: str): parts: T.List[dict] = [] + if prompt and prompt.strip(): parts.append({"text": prompt}) + if images is not None and torch.is_tensor(images): batch = images if images.ndim == 4 else images.unsqueeze(0) for i in range(batch.shape[0]): pil = tensor_to_pil(batch[i:i+1]) parts.append({"inline_data": {"mime_type": mime, "data": encode_pil_bytes(pil, mime)}}) + return parts - - def _build_config(self, temperature: float, want_text: bool, aspect_ratio: str = "1:1"): + + def _build_config(self, temperature: float, want_text: bool, aspect_ratio: str): try: from google.genai import types + cfg = types.GenerateContentConfig( temperature=float(temperature), top_p=float(_TOP_P), top_k=int(_TOP_K), max_output_tokens=int(_MAX_TOKENS), - response_modalities=["IMAGE","TEXT"] if want_text else ["IMAGE"], + response_modalities=["IMAGE", "TEXT"] if want_text else ["IMAGE"], image_config=types.ImageConfig(aspect_ratio=aspect_ratio), safety_settings=[ {"category": "HARM_CATEGORY_HARASSMENT", "threshold": "BLOCK_NONE"}, @@ -179,14 +222,17 @@ class PVL_Google_NanoBanana_API: {"category": "HARM_CATEGORY_DANGEROUS_CONTENT", "threshold": "BLOCK_NONE"}, ], ) + return cfg except Exception: + # Fallback dict config return { "temperature": float(temperature), "top_p": float(_TOP_P), "top_k": int(_TOP_K), "max_output_tokens": int(_MAX_TOKENS), - "response_modalities": ["IMAGE","TEXT"] if want_text else ["IMAGE"], + "response_modalities": ["IMAGE", "TEXT"] if want_text else ["IMAGE"], + "image_config": {"aspect_ratio": aspect_ratio}, "safety_settings": [ {"category": "HARM_CATEGORY_HARASSMENT", "threshold": "BLOCK_NONE"}, {"category": "HARM_CATEGORY_HATE_SPEECH", "threshold": "BLOCK_NONE"}, @@ -194,7 +240,35 @@ class PVL_Google_NanoBanana_API: {"category": "HARM_CATEGORY_DANGEROUS_CONTENT", "threshold": "BLOCK_NONE"}, ], } - + + def _build_call_prompts(self, base_prompts: T.List[str], num_images: int, debug: bool) -> T.List[str]: + """ + Maps prompts to calls according to the agreed rule: + - If len(prompts) >= num_images: take first num_images + - If len(prompts) < num_images: repeat the last prompt to fill N + """ + N = max(1, int(num_images)) + + if not base_prompts: + return [""] + + if len(base_prompts) >= N: + call_prompts = base_prompts[:N] + else: + if debug: + print(f"[PVL NODE] Provided {len(base_prompts)} prompts but num_images={N}. " + f"Reusing the last prompt for remaining calls.") + print(f"[PVL WARNING] prompt list shorter than num_images ({len(base_prompts)} < {N}). " + f"Last entry will be reused for the remaining {N - len(base_prompts)} calls.") + call_prompts = base_prompts + [base_prompts[-1]] * (N - len(base_prompts)) + + if debug: + for i, cp in enumerate(call_prompts, 1): + show = cp if len(cp) <= 160 else (cp[:157] + "...") + print(f"[PVL NODE] Call {i} prompt: {show}") + + return call_prompts + def _single_google_call(self, client, model: str, parts: list, cfg, request_id: str, debug: bool): kwargs = {} try: @@ -203,19 +277,21 @@ class PVL_Google_NanoBanana_API: kwargs["request_options"] = types.RequestOptions(request_id=request_id.strip()) except Exception: pass - + resp = client.models.generate_content( model=model, contents=[{"role": "user", "parts": parts}], config=cfg, **kwargs ) - + imgs, texts = [], [] cands = getattr(resp, "candidates", None) or [] + for cand in cands: content = getattr(cand, "content", None) finish_reason = getattr(cand, "finish_reason", None) + if debug: try: um = getattr(resp, "usage_metadata", None) @@ -224,10 +300,12 @@ class PVL_Google_NanoBanana_API: print(f"[PVL Debug] finish_reason: {finish_reason}") except Exception: pass + if content is None: continue - parts = getattr(content, "parts", []) or [] - for p in parts: + + pparts = getattr(content, "parts", []) or [] + for p in pparts: blob = _extract_image_bytes_from_part(p) if blob: try: @@ -240,298 +318,362 @@ class PVL_Google_NanoBanana_API: t = _extract_text_from_part(p) if t: texts.append(t) - + return imgs, texts, resp - - # ---- FAL fallback via Queue API ---- - def _fal_queue_call(self, route: str, prompt: str, image_tensor: T.Optional[torch.Tensor], - mime: str, fal_key: str, timeout: int, debug: bool, num_images: int, output_format: str): + + # -------- FAL Queue API - TWO PHASE EXECUTION -------- + + def _fal_submit_only(self, route: str, prompt: str, image_tensors: T.Optional[torch.Tensor], + mime: str, fal_key: str, timeout: int, debug: bool, + output_format: str, aspect_ratio: str = "1:1", sync_mode: bool = False): + """ + Phase 1: Submit request to FAL queue and return request info immediately. + Does NOT poll for completion. + """ if not fal_key: - raise RuntimeError("FAL fallback requested but FAL_KEY is missing. Provide fal_api_key or set env FAL_KEY.") - if image_tensor is None or not torch.is_tensor(image_tensor): - raise RuntimeError("FAL fallback requires an input image tensor.") - - batch = image_tensor if image_tensor.ndim == 4 else image_tensor.unsqueeze(0) - # Prepare data URIs (all frames) + raise RuntimeError("FAL requested but FAL_KEY is missing.") + + # Build data URLs only if image provided data_urls: T.List[str] = [] - for i in range(batch.shape[0]): - pil = tensor_to_pil(batch[i:i+1]) - raw = encode_pil_bytes(pil, mime) - data_urls.append(_data_url(mime, raw)) - + if image_tensors is not None and torch.is_tensor(image_tensors): + batch = image_tensors if image_tensors.ndim == 4 else image_tensors.unsqueeze(0) + for i in range(batch.shape[0]): + pil = tensor_to_pil(batch[i:i + 1]) + raw = encode_pil_bytes(pil, mime) + data_urls.append(_data_url(mime, raw)) + base = "https://queue.fal.run" submit_url = f"{base}/{route.strip()}" headers = {"Authorization": f"Key {fal_key}"} + payload = { "prompt": prompt or "", - # Only plural form to avoid duplicate outputs - "image_urls": data_urls, - "num_images": int(max(1, num_images)), - "output_format": ("png" if str(output_format).lower()=="png" else "jpeg"), - "sync_mode": True, + "num_images": 1, + "output_format": "png" if str(output_format).lower() == "png" else "jpeg", "aspect_ratio": aspect_ratio, + "sync_mode": sync_mode, } + + # Only add image_urls for img2img route + if data_urls: + payload["image_urls"] = data_urls + if debug: - print(f"[PVL FAL] QUEUE SUBMIT {submit_url} with {len(data_urls)} image(s) and num_images={payload['num_images']}") - - r = requests.post(submit_url, headers=headers, json=payload, timeout=timeout) - if r.status_code >= 400: - raise RuntimeError(f"FAL queue submit error {r.status_code}: {r.text}") - try: - sub = r.json() - except Exception: - raise RuntimeError("FAL queue submit returned non-JSON.") - + print(f"[FAL SUBMIT] prompt: {prompt[:60]}... sync_mode={sync_mode}") + + # Submit request + rr = requests.post(submit_url, headers=headers, json=payload, timeout=timeout) + if rr.status_code != 200: + raise RuntimeError(f"FAL submit error {rr.status_code}: {rr.text}") + + sub = rr.json() req_id = sub.get("request_id") - status_url = sub.get("status_url") or (f"{base}/{route.strip()}/requests/{req_id}/status" if req_id else None) - resp_url = sub.get("response_url") or (f"{base}/{route.strip()}/requests/{req_id}" if req_id else None) - if not req_id or not status_url or not resp_url: - raise RuntimeError("FAL queue submit missing request_id/status_url/response_url.") - - # Poll - deadline = time.time() + max(5, int(timeout)) - last_status = None - while time.time() < deadline: - sr = requests.get(status_url, headers=headers, timeout=10) - if not sr.ok: - time.sleep(0.5); continue - sdata = sr.json() - last_status = sdata.get("status") - if last_status == "COMPLETED": - break - time.sleep(0.6) - if last_status != "COMPLETED": - raise RuntimeError(f"FAL queue did not complete in time (last status={last_status})") - - # Fetch result - rr = requests.get(resp_url, headers=headers, timeout=15) - if rr.status_code >= 400: - raise RuntimeError(f"FAL result fetch error {rr.status_code}: {rr.text}") - try: - rdata = rr.json() - except Exception: - raise RuntimeError("FAL result returned non-JSON.") - + if not req_id: + raise RuntimeError("FAL did not return a request_id") + + # Get status and result URLs + status_url = sub.get("status_url") or f"{base}/{route.strip()}/requests/{req_id}/status" + resp_url = sub.get("response_url") or f"{base}/{route.strip()}/requests/{req_id}" + + return { + "request_id": req_id, + "status_url": status_url, + "response_url": resp_url, + "prompt": prompt + } + + def _fal_poll_and_fetch(self, request_info: dict, fal_key: str, timeout: int, debug: bool): + """ + Phase 2: Poll a single FAL request until complete and fetch the result. + Returns (image_tensor, description_text). + """ + headers = {"Authorization": f"Key {fal_key}"} + status_url = request_info["status_url"] + resp_url = request_info["response_url"] + req_id = request_info["request_id"] + if debug: + print(f"[FAL POLL] request_id={req_id[:16]}...") + + # Poll for completion with timeout check + deadline = time.time() + timeout + completed = False + while time.time() < deadline: try: - s = json.dumps(rdata)[:1200] - except Exception: - s = str(rdata)[:1200] - print("[PVL FAL] raw response:", s) - - # Normalize containers - resp = rdata.get("response") if isinstance(rdata, dict) else None - if not isinstance(resp, dict): - resp = rdata if isinstance(rdata, dict) else {} - - description = resp.get("description") or rdata.get("description") or resp.get("output_text") or "" - - # Collect potential image items + sr = requests.get(status_url, headers=headers, timeout=min(10, timeout)) + if sr.ok and sr.json().get("status") == "COMPLETED": + completed = True + break + except Exception as e: + if debug: + print(f"[FAL POLL] Status check error: {e}") + time.sleep(0.6) + + # Check if we timed out + if not completed: + raise RuntimeError(f"FAL request {req_id[:16]} timed out after {timeout}s") + + # Fetch result + rr = requests.get(resp_url, headers=headers, timeout=min(15, timeout)) + if not rr.ok: + raise RuntimeError(f"FAL result fetch error {rr.status_code}: {rr.text}") + + data = rr.json() + if debug: + print(f"[FAL RESULT] request_id={req_id[:16]}... status=COMPLETED") + + # Extract response data + resp = data.get("response") if isinstance(data, dict) else None + if resp is None and isinstance(data, dict): + resp = data + + # Parse images from various possible locations + images_out: T.List[torch.Tensor] = [] buckets: T.List[T.Union[str, dict]] = [] - for key in ("images", "outputs", "artifacts"): - val = resp.get(key) - if isinstance(val, list): - buckets.extend(val) - for key in ("image", "output", "result"): - val = resp.get(key) - if isinstance(val, (str, dict)): - buckets.append(val) - for key in ("images", "image", "output", "outputs", "artifacts"): - val = rdata.get(key) if isinstance(rdata, dict) else None - if isinstance(val, list): - buckets.extend(val) - elif isinstance(val, (str, dict)): - buckets.append(val) - - # Decode images, but respect num_images cap - out = [] - def add_image_from_item(item): + + if isinstance(resp, dict): + for key in ("images", "outputs", "artifacts"): + val = resp.get(key) + if isinstance(val, list): + buckets.extend(val) + + for key in ("image", "output", "result"): + val = resp.get(key) + if isinstance(val, (str, dict)): + buckets.append(val) + + for item in buckets: try: - if isinstance(item, str): - url_or_data = item - elif isinstance(item, dict): - url_or_data = item.get("url") or item.get("data") or item.get("image") or item.get("content") - else: - return + url_or_data = item if isinstance(item, str) else (item.get("url") or item.get("data") or item.get("image")) if not isinstance(url_or_data, str): - return + continue + if url_or_data.startswith("data:image/"): b64 = url_or_data.split(",", 1)[1] blob = base64.b64decode(b64) else: - ir = requests.get(url_or_data, timeout=timeout) + ir = requests.get(url_or_data, timeout=min(15, timeout)) if not ir.ok: - return + continue blob = ir.content + pil = Image.open(io.BytesIO(blob)).convert("RGB") - out.append(pil_to_tensor(pil)) + images_out.append(pil_to_tensor(pil)) except Exception as ex: if debug: - print("[PVL FAL] image decode failed:", ex) - - for item in buckets: - if len(out) >= int(max(1, num_images)): - break - add_image_from_item(item) - - if not out: - raise RuntimeError("FAL API returned no images.") - - if len(out) == 1: - images_tensor = out[0] - else: - images_tensor = torch.cat(out, dim=0) - - if debug: - print(f"[PVL FAL] returning {len(out)} image(s) (capped to num_images={int(max(1, num_images))})") - - return images_tensor, description - - # ---- main ---- - def run(self, prompt: str, images: T.Optional[torch.Tensor] = None, - aspect_ratio: str = "1:1", - model: str = DEFAULT_MODEL, endpoint_override: str = "", - api_key: str = "", - temperature: float = 0.6, output_format: str = "png", - capture_text_output: bool = False, num_images: int = 1, - timeout_sec: int = 120, request_id: str = "", - debug_log: bool = False, - use_fal_fallback: bool = False, fal_api_key: str = "", fal_route: str = "fal-ai/nano-banana/edit"): - + print("[FAL] image decode failed:", ex) + + description = resp.get("description", "") if isinstance(resp, dict) else "" + + if not images_out: + raise RuntimeError(f"FAL returned no images for request_id={req_id}") + + return images_out[0], description + + # -------- main -------- + + def run( + self, + prompt: str, + delimiter: str = "[*]", + images: T.Optional[torch.Tensor] = None, + aspect_ratio: str = "1:1", + model: str = DEFAULT_MODEL, + endpoint_override: str = "", + api_key: str = "", + temperature: float = 0.6, + output_format: str = "png", + capture_text_output: bool = False, + num_images: int = 1, + timeout_sec: int = 120, + request_id: str = "", + debug_log: bool = False, + use_fal_fallback: bool = True, + force_fal: bool = False, + sync_mode: bool = False, + fal_api_key: str = "", + fal_route_img2img: str = "fal-ai/nano-banana/edit", + fal_route_txt2img: str = "fal-ai/nano-banana", + ): + # Validate aspect_ratio if aspect_ratio.strip() not in _VALID_ASPECTS: print(f"[PVL WARNING] Invalid or missing aspect_ratio '{aspect_ratio}', defaulting to 1:1.") aspect_ratio = "1:1" - - key = (api_key or os.getenv("GEMINI_API_KEY","")).strip() + + # Split prompts using regex delimiter + try: + base_prompts = [p.strip() for p in re.split(delimiter, prompt) if str(p).strip()] + except re.error: + print(f"[PVL WARNING] Invalid regex pattern '{delimiter}', using literal split.") + base_prompts = [p.strip() for p in prompt.split(delimiter) if str(p).strip()] + + if not base_prompts: + raise RuntimeError("No valid prompts provided.") + + # Map prompts to num_images calls + call_prompts = self._build_call_prompts(base_prompts, num_images, debug_log) + + key = (api_key or os.getenv("GEMINI_API_KEY", "")).strip() input_mime = "image/png" if str(output_format).lower() == "png" else "image/jpeg" want_text = bool(capture_text_output) - - # If no Google key but fallback is enabled, attempt FAL directly + + # Decide FAL route based on presence of input images + route_to_use = fal_route_img2img if (images is not None and torch.is_tensor(images) and images.numel() > 0) else fal_route_txt2img + + # Helper function for parallel FAL submission + polling + def parallel_fal_execution(prompts_list, fal_key_str, route, debug): + """Submit all FAL requests in parallel, then poll all in parallel""" + if debug: + print(f"[FAL] Submitting {len(prompts_list)} requests in parallel...") + + # PHASE 1: Submit all requests IN PARALLEL + submit_results = [] + with ThreadPoolExecutor(max_workers=min(len(prompts_list), 6)) as ex: + submit_futs = { + ex.submit(self._fal_submit_only, route, p, images, input_mime, + fal_key_str, timeout_sec, debug, output_format, + aspect_ratio, sync_mode): p + for p in prompts_list + } + for fut in as_completed(submit_futs): + try: + req_info = fut.result() + submit_results.append(req_info) + except Exception as e: + if debug: + print(f"[FAL SUBMIT ERROR] {e}") + + if not submit_results: + raise RuntimeError("All FAL submission requests failed") + + if debug: + print(f"[FAL] {len(submit_results)} requests submitted successfully. Polling for results...") + + # PHASE 2: Poll all requests IN PARALLEL + results, texts = [], [] + failed_count = 0 + with ThreadPoolExecutor(max_workers=min(len(submit_results), 6)) as ex: + poll_futs = { + ex.submit(self._fal_poll_and_fetch, req_info, fal_key_str, + timeout_sec, debug): req_info + for req_info in submit_results + } + for fut in as_completed(poll_futs): + try: + img, t = fut.result() + results.append(img) + if t: + texts.append(t) + except Exception as e: + failed_count += 1 + if debug: + print(f"[FAL POLL ERROR] {e}") + + if not results: + raise RuntimeError(f"All FAL requests failed during polling ({failed_count} failures)") + + if failed_count > 0: + print(f"[PVL WARNING] {failed_count}/{len(submit_results)} FAL requests failed, continuing with {len(results)} successful results") + + return results, texts + + # ---- CASE: FAL only (TRUE PARALLEL) ---- + if force_fal: + fal_key = (fal_api_key or os.getenv("FAL_KEY", "")).strip() + if not fal_key: + raise RuntimeError("force_fal=True but FAL_KEY missing.") + + results, texts = parallel_fal_execution(call_prompts, fal_key, route_to_use, debug_log) + + images_tensor = stack_images_same_size(results, debug_log) + text_out = "\n".join(texts) if want_text else "" + return text_out, images_tensor + + # ---- Google path ---- if not key: if use_fal_fallback: - fal_key = (fal_api_key or os.getenv("FAL_KEY","")).strip() - try: - img_tensor, fal_text = self._fal_queue_call(fal_route, prompt, images, input_mime, fal_key, int(timeout_sec), debug_log, num_images, output_format) - text_out = (fal_text or "") if want_text else "" - if text_out: - print("[PVL FAL Text]:\n" + text_out) - return (text_out, img_tensor) - except Exception as fe: - print(f"[PVL FAL Fallback] FAL call failed without Google key: {fe}") - raise RuntimeError(f"Gemini image generation failed and FAL fallback also failed: {fe}") + # no Google key, try FAL directly (TRUE PARALLEL) + fal_key = (fal_api_key or os.getenv("FAL_KEY", "")).strip() + if not fal_key: + raise RuntimeError("GEMINI_API_KEY missing and FAL_KEY missing.") + + results, texts = parallel_fal_execution(call_prompts, fal_key, route_to_use, debug_log) + + images_tensor = stack_images_same_size(results, debug_log) + text_out = "\n".join(texts) if want_text else "" + return text_out, images_tensor + raise RuntimeError("Gemini API key missing. Pass api_key or set GEMINI_API_KEY.") - - # Build client & request + client = self._make_client(key, endpoint_override) - parts = self._build_parts(prompt, images, input_mime) cfg = self._build_config(temperature, want_text, aspect_ratio) - + if debug_log: - p_preview = (prompt or "")[:180].replace("\n"," ") - img_count = (images.shape[0] if (isinstance(images, torch.Tensor) and images.ndim==4) else (1 if isinstance(images, torch.Tensor) else 0)) - print(f"[PVL Debug] prompt chars={len(prompt or '')} preview='{p_preview}...'") - print(f"[PVL Debug] parts: text={1 if (prompt and prompt.strip()) else 0}, images={img_count}") - safe_parts = [] - for pr in parts: - if "text" in pr: - safe_parts.append({"text": pr["text"][:120]}) - elif "inline_data" in pr: - di = pr["inline_data"] - safe_parts.append({"inline_data": {"mime_type": di.get("mime_type","image/*"), "data": ""}}) - try: - temp = getattr(cfg,'temperature',None) if hasattr(cfg,'temperature') else cfg.get('temperature') - top_p = getattr(cfg,'top_p',None) if hasattr(cfg,'top_p') else cfg.get('top_p') - top_k = getattr(cfg,'top_k',None) if hasattr(cfg,'top_k') else cfg.get('top_k') - mot = getattr(cfg,'max_output_tokens',None) if hasattr(cfg,'max_output_tokens') else cfg.get('max_output_tokens') - mods = getattr(cfg,'response_modalities',None) if hasattr(cfg,'response_modalities') else cfg.get('response_modalities') - except Exception: - temp=top_p=top_k=mot=mods=None - print("[PVL Debug] config:", {"temperature": temp, "top_p": top_p, "top_k": top_k, "max_output_tokens": mot, "modalities": mods}) - print("[PVL Debug] contents:", [{"role":"user","parts": safe_parts}]) - - # Parallel Google calls - N = max(1, int(num_images)) - results = [None] * N - errors = [] - - def call_i(i: int): - rid = (request_id.strip() + f"-{i}") if request_id.strip() else f"pvl-nb-{int(time.time()*1000)}-{i}" - return self._single_google_call(client, model, parts, cfg, rid, debug_log) - - max_workers = min(N, 6) - if N == 1: - try: - results[0] = call_i(0) - except Exception as e: - errors.append(f"google call 0 failed: {e}") - else: - with ThreadPoolExecutor(max_workers=max_workers) as ex: - futmap = {ex.submit(call_i, i): i for i in range(N)} - for fut in as_completed(futmap): - i = futmap[fut] - try: - results[i] = fut.result() - except Exception as e: - errors.append(f"google call {i} failed: {e}") - - # Parse all Google results - out_imgs, out_texts = [], [] - for idx, tup in enumerate(results): - if tup is None: - errors.append(f"google call {idx} returned no response") - continue - imgs_i, texts_i, resp = tup - if not imgs_i: + print(f"[PVL Debug] Processing {len(call_prompts)} prompts in parallel") + + def google_call(p: str, debug: bool): + parts = self._build_parts(p, images, input_mime) + g_imgs, g_texts, _resp = self._single_google_call(client, model, parts, cfg, request_id, debug) + return g_imgs, g_texts + + out_imgs: T.List[torch.Tensor] = [] + out_texts: T.List[str] = [] + failed_google = 0 + + # Google API calls with error handling + with ThreadPoolExecutor(max_workers=min(len(call_prompts), 6)) as ex: + futs = {ex.submit(google_call, p, debug_log): p for p in call_prompts} + for fut in as_completed(futs): try: - cands = getattr(resp, "candidates", None) or [] - fin = getattr(cands[0], "finish_reason", None) if cands else None - except Exception: - fin = None - errors.append(f"google call {idx} returned no images (finish_reason={fin})") - else: - out_imgs.extend(imgs_i) - if texts_i: - out_texts.append("\n".join(texts_i)) - - # If Google failed and fallback enabled -> try FAL - if errors and use_fal_fallback: - print("[PVL Fallback] Google call failed; attempting FAL.ai fallback...") - for e in errors: - print("[PVL Google Error]", e) - try: - fal_key = (fal_api_key or os.getenv("FAL_KEY","")).strip() - img_tensor, fal_text = self._fal_queue_call(fal_route, prompt, images, input_mime, fal_key, int(timeout_sec), debug_log, num_images, output_format) - final_text = "" # Start with empty; then merge Google texts if requested - if bool(capture_text_output): - pieces = [] - if out_texts: - pieces.append("\n\n--- Google ---\n\n" + ("\n".join(out_texts))) - if fal_text: - pieces.append("\n\n--- FAL ---\n\n" + fal_text) - final_text = "".join(pieces) - if final_text: - print("[PVL Fallback Note] Combined text:") - print(final_text) - return (final_text, img_tensor) - except Exception as fe: - print("[PVL FAL Error]", fe) - raise RuntimeError("Both Google and FAL failed. See console for details.") - - # If Google produced errors and fallback not used -> raise - if errors: - if out_texts: - print("[PVL Google Text]:\n" + ("\n\n---\n\n".join(out_texts))) - raise RuntimeError("Gemini image generation failed: " + " | ".join(errors[:5])) - + imgs, texts = fut.result() + if imgs: + out_imgs.append(imgs[0]) + out_texts.extend(texts) + except Exception as e: + failed_google += 1 + if debug_log: + print(f"[GOOGLE ERROR] {e}") + + if failed_google > 0: + print(f"[PVL WARNING] {failed_google}/{len(call_prompts)} Google requests failed") + + # If Google failed to produce images, optionally fallback to FAL (TRUE PARALLEL) if not out_imgs: - raise RuntimeError("Gemini image generation failed: no images across all Google calls.") - - images_tensor = torch.cat(out_imgs, dim=0) if len(out_imgs) > 1 else out_imgs[0] - final_text = ("\n\n---\n\n".join(out_texts)) if (bool(capture_text_output) and out_texts) else "" - if final_text: - print(f"[PVL Google NanoBanana Output]:\n{final_text}\n") - - return (final_text, images_tensor,) + if use_fal_fallback: + fal_key = (fal_api_key or os.getenv("FAL_KEY", "")).strip() + if not fal_key: + raise RuntimeError("Google returned no images and FAL_KEY is missing for fallback.") + + results, texts = parallel_fal_execution(call_prompts, fal_key, route_to_use, debug_log) + + images_tensor = stack_images_same_size(results, debug_log) + + if want_text: + combined_text = "" + if out_texts: + combined_text += "\n\n--- Google ---\n\n" + "\n".join(out_texts) + if texts: + combined_text += "\n\n--- FAL ---\n\n" + "\n".join(texts) + text_out = combined_text + else: + text_out = "" + + return text_out, images_tensor + else: + raise RuntimeError("Gemini returned no images") + + # Merge Google images to a 4D tensor + images_tensor = stack_images_same_size(out_imgs, debug_log) + + if images_tensor.ndim == 3: + images_tensor = images_tensor.unsqueeze(0) + + text_out = "\n".join(out_texts) if (want_text and out_texts) else "" + + if text_out and debug_log: + print("[PVL Google NanoBanana Output]:\n" + text_out) + + return text_out, images_tensor NODE_CLASS_MAPPINGS = {"PVL_Google_NanoBanana_API": PVL_Google_NanoBanana_API} NODE_DISPLAY_NAME_MAPPINGS = {"PVL_Google_NanoBanana_API": NODE_NAME} diff --git a/pvl_google_nano_banana_mandatory_img.py b/pvl_google_nano_banana_mandatory_img.py index 9de8e5c..d040128 100644 --- a/pvl_google_nano_banana_mandatory_img.py +++ b/pvl_google_nano_banana_mandatory_img.py @@ -1,46 +1,47 @@ +# pvl_google_nano_banana_mandatory_img.py + +# Node: PVL Google Nano-Banana API (Mandatory IMG + Delimiter + Aspect Ratio + Parallel + FAL) -# pvl_google_nano_banana.py -# Node: PVL Google Nano-Banana API (SDK + robust FAL fallback) # Author: PVL # License: MIT -# -# Requires: -# pip install google-genai -# -# Changes in this build: -# - FAL fallback now sends ONLY "image_urls" (list) to avoid duplicate outputs from some routes. -# - Enforces num_images: if num_images==1 returns the FIRST image only; otherwise caps to num_images. -# - Keeps sync_mode true and robust response parsing; prints raw FAL JSON when debug is ON. -# - Return order: ("STRING","IMAGE") => (text, images). + +# Features in this build: +# - Mandatory image input. +# - Literal delimiter input (default "[*]") to split a long prompt into multiple prompts. +# - Prompt reuse rule to match num_images +# - Aspect ratio support (STRING) via new SDK +# - TRUE PARALLEL execution for both Google and FAL +# - FAL fallback and Force-FAL modes with sync_mode toggle +# - Error handling for individual requests with partial results support +# - Execution time is printed at the end import os, io, json, base64, typing as T, time from concurrent.futures import ThreadPoolExecutor, as_completed - import requests import numpy as np from PIL import Image import torch +from google.genai import types NODE_NAME = "PVL Google Nano-Banana API mandatory IMG" NODE_CATEGORY = "PVL/Google" DEFAULT_MODEL = "gemini-2.5-flash-image-preview" - _TOP_P = 0.95 _TOP_K = 64 _MAX_TOKENS = 4096 -_VALID_ASPECTS = {"21:9","1:1","4:3","3:2","2:3","5:4","4:5","3:4","16:9","9:16"} -# --- image helpers --- +# --------------------------- Image/Tensor helpers --------------------------- + def pil_to_tensor(img: Image.Image) -> torch.Tensor: if img.mode != "RGB": img = img.convert("RGB") arr = np.asarray(img, dtype=np.float32) / 255.0 - return torch.from_numpy(arr)[None, ...] + return torch.from_numpy(arr)[None, ...] # (1,H,W,3) def tensor_to_pil(t: torch.Tensor) -> Image.Image: if t.ndim == 4: t = t[0] - arr = (t.clamp(0,1).cpu().numpy() * 255).astype("uint8") + arr = (t.clamp(0, 1).cpu().numpy() * 255).astype("uint8") return Image.fromarray(arr, "RGB") def encode_pil_bytes(img: Image.Image, mime: str) -> bytes: @@ -65,475 +66,476 @@ def _extract_image_bytes_from_part(part) -> T.Optional[bytes]: return None except Exception: pass + if isinstance(part, dict): - if "inline_data" in part and isinstance(part["inline_data"], dict): - blob = part["inline_data"].get("data") - elif "inlineData" in part and isinstance(part["inlineData"], dict): - blob = part["inlineData"].get("data") - else: - blob = None - if isinstance(blob, (bytes, bytearray)): - return bytes(blob) - if isinstance(blob, str): - try: - return base64.b64decode(blob, validate=False) - except Exception: - return None + inline = part.get("inline_data") or part.get("inlineData") + if isinstance(inline, dict): + data = inline.get("data") + if isinstance(data, (bytes, bytearray)): + return bytes(data) + if isinstance(data, str): + try: + return base64.b64decode(data, validate=False) + except Exception: + return None + return None def _extract_text_from_part(part) -> T.Optional[str]: - try: - txt = getattr(part, "text", None) - if txt is not None: - return str(txt) - except Exception: - pass - if isinstance(part, dict) and "text" in part and part["text"] is not None: - return str(part["text"]) - return None + if isinstance(part, dict): + if "text" in part and part["text"] is not None: + return str(part["text"]) + txt = getattr(part, "text", None) + return str(txt) if txt is not None else None def _data_url(mime: str, raw: bytes) -> str: return f"data:{mime};base64," + base64.b64encode(raw).decode("utf-8") +def _stack_images_same_size(tensors: T.List[torch.Tensor], debug: bool = False) -> torch.Tensor: + if not tensors: + raise RuntimeError("No images to stack.") + try: + return torch.cat(tensors, dim=0) + except RuntimeError: + if debug: + print("[PVL NODE] Mismatched sizes, resizing to match first image.") + target_h, target_w = tensors[0].shape[1], tensors[0].shape[2] + fixed = [] + for t in tensors: + pil = tensor_to_pil(t) + rp = pil.resize((target_w, target_h), Image.LANCZOS) + fixed.append(pil_to_tensor(rp)) + return torch.cat(fixed, dim=0) + +# --------------------------- Main Node --------------------------- + class PVL_Google_NanoBanana_API_mandatory_IMG: + @classmethod def INPUT_TYPES(cls): return { "required": { - "prompt": ("STRING", {"multiline": True, "default": "A tiny banana spaceship over a neon city."}), + "prompt": ("STRING", {"multiline": True, "default": "", "placeholder": "Enter prompts separated by delimiter"}), "images": ("IMAGE",), }, "optional": { - + "delimiter": ("STRING", {"default": "[*]", "placeholder": "Delimiter string"}), + "aspect_ratio": ("STRING", {"default": "1:1"}), "model": ("STRING", {"default": DEFAULT_MODEL}), "endpoint_override": ("STRING", {"default": ""}), - "api_key": ("STRING", {"default": "", "multiline": False, - "placeholder": "Leave empty to use GEMINI_API_KEY"}), + "api_key": ("STRING", {"default": "", "multiline": False, "placeholder": "Leave empty to use GEMINI_API_KEY"}), "temperature": ("FLOAT", {"default": 0.6, "min": 0.0, "max": 2.0, "step": 0.05}), - "output_format": (["png","jpeg"], {"default": "png"}), + "output_format": (["png", "jpeg"], {"default": "png"}), "capture_text_output": ("BOOLEAN", {"default": False}), "num_images": ("INT", {"default": 1, "min": 1, "max": 12, "step": 1}), "timeout_sec": ("INT", {"default": 120, "min": 5, "max": 600, "step": 5}), - "request_id": ("STRING", {"default": ""}), "debug_log": ("BOOLEAN", {"default": False}), - - "aspect_ratio": ("STRING", {"default": "1:1", "placeholder": "e.g. 16:9, 9:16, 3:2 — works for both Google & FAL"}), - - # --- FAL fallback --- - "use_fal_fallback": ("BOOLEAN", {"default": False}), - "fal_api_key": ("STRING", {"default": "", "multiline": False, - "placeholder": "Leave empty to use FAL_KEY"}), + "use_fal_fallback": ("BOOLEAN", {"default": True}), + "force_fal": ("BOOLEAN", {"default": False}), + "sync_mode": ("BOOLEAN", {"default": False}), + "fal_api_key": ("STRING", {"default": "", "multiline": False, "placeholder": "Leave empty to use FAL_KEY"}), "fal_route": ("STRING", {"default": "fal-ai/nano-banana/edit"}), } } - - RETURN_TYPES = ("STRING","IMAGE",) - RETURN_NAMES = ("text","images") + + RETURN_TYPES = ("STRING", "IMAGE",) + RETURN_NAMES = ("text", "images") FUNCTION = "run" CATEGORY = NODE_CATEGORY - - # ---- helpers ---- + + # ------------------------- INTERNAL HELPERS ------------------------- + def _make_client(self, api_key: str, endpoint_override: str): - try: - from google import genai - from google.genai import types - except Exception as e: - raise RuntimeError("Google GenAI SDK not installed. Run: pip install google-genai") from e - + from google import genai + from google.genai import types as gtypes + http_options = None if endpoint_override.strip(): try: - http_options = types.HttpOptions(base_url=endpoint_override.strip()) + http_options = gtypes.HttpOptions(base_url=endpoint_override.strip()) except Exception: http_options = None - + if http_options is not None: - client = genai.Client(api_key=api_key, http_options=http_options) - else: - client = genai.Client(api_key=api_key) - return client - - def _build_parts(self, prompt: str, images: T.Optional[torch.Tensor], mime: str): + return genai.Client(api_key=api_key, http_options=http_options) + return genai.Client(api_key=api_key) + + def _build_parts(self, prompt: str, images: torch.Tensor, mime: str): parts: T.List[dict] = [] + if prompt and prompt.strip(): parts.append({"text": prompt}) - if images is not None and torch.is_tensor(images): - batch = images if images.ndim == 4 else images.unsqueeze(0) - for i in range(batch.shape[0]): - pil = tensor_to_pil(batch[i:i+1]) - parts.append({"inline_data": {"mime_type": mime, "data": encode_pil_bytes(pil, mime)}}) + + batch = images if images.ndim == 4 else images.unsqueeze(0) + for i in range(batch.shape[0]): + pil = tensor_to_pil(batch[i:i+1]) + parts.append({"inline_data": {"mime_type": mime, "data": encode_pil_bytes(pil, mime)}}) + return parts - - def _build_config(self, temperature: float, want_text: bool, aspect_ratio: str = "1:1"): - try: - from google.genai import types - cfg = types.GenerateContentConfig( - temperature=float(temperature), - top_p=float(_TOP_P), - top_k=int(_TOP_K), - max_output_tokens=int(_MAX_TOKENS), - response_modalities=["IMAGE","TEXT"] if want_text else ["IMAGE"], - image_config=types.ImageConfig(aspect_ratio=aspect_ratio), - safety_settings=[ - {"category": "HARM_CATEGORY_HARASSMENT", "threshold": "BLOCK_NONE"}, - {"category": "HARM_CATEGORY_HATE_SPEECH", "threshold": "BLOCK_NONE"}, - {"category": "HARM_CATEGORY_SEXUALLY_EXPLICIT", "threshold": "BLOCK_NONE"}, - {"category": "HARM_CATEGORY_DANGEROUS_CONTENT", "threshold": "BLOCK_NONE"}, - ], - ) - return cfg - except Exception: - return { - "temperature": float(temperature), - "top_p": float(_TOP_P), - "top_k": int(_TOP_K), - "max_output_tokens": int(_MAX_TOKENS), - "response_modalities": ["IMAGE","TEXT"] if want_text else ["IMAGE"], - "safety_settings": [ - {"category": "HARM_CATEGORY_HARASSMENT", "threshold": "BLOCK_NONE"}, - {"category": "HARM_CATEGORY_HATE_SPEECH", "threshold": "BLOCK_NONE"}, - {"category": "HARM_CATEGORY_SEXUALLY_EXPLICIT", "threshold": "BLOCK_NONE"}, - {"category": "HARM_CATEGORY_DANGEROUS_CONTENT", "threshold": "BLOCK_NONE"}, - ], - } - - def _single_google_call(self, client, model: str, parts: list, cfg, request_id: str, debug: bool): - kwargs = {} - try: - from google.genai import types - if request_id.strip(): - kwargs["request_options"] = types.RequestOptions(request_id=request_id.strip()) - except Exception: - pass - - resp = client.models.generate_content( - model=model, - contents=[{"role": "user", "parts": parts}], - config=cfg, - **kwargs - ) - - imgs, texts = [], [] - cands = getattr(resp, "candidates", None) or [] - for cand in cands: - content = getattr(cand, "content", None) - finish_reason = getattr(cand, "finish_reason", None) - if debug: - try: - um = getattr(resp, "usage_metadata", None) - if um is not None: - print("[PVL Debug] usage_metadata:", getattr(um, "__dict__", str(um))) - print(f"[PVL Debug] finish_reason: {finish_reason}") - except Exception: - pass - if content is None: - continue - parts = getattr(content, "parts", []) or [] - for p in parts: - blob = _extract_image_bytes_from_part(p) - if blob: - try: - pil = Image.open(io.BytesIO(blob)).convert("RGB") - imgs.append(pil_to_tensor(pil)) - except Exception as ex: - if debug: - print("[PVL Debug] image decode error:", ex) - else: - t = _extract_text_from_part(p) - if t: - texts.append(t) - - return imgs, texts, resp - - # ---- FAL fallback via Queue API ---- - def _fal_queue_call(self, route: str, prompt: str, image_tensor: T.Optional[torch.Tensor], - mime: str, fal_key: str, timeout: int, debug: bool, num_images: int, output_format: str): + + # ---- FAL Queue API - TWO PHASE EXECUTION ---- + + def _fal_submit_only(self, route: str, prompt: str, image_tensor: torch.Tensor, + mime: str, fal_key: str, timeout: int, debug: bool, + output_format: str, aspect_ratio: str = "1:1", sync_mode: bool = False): + """ + Phase 1: Submit request to FAL queue and return request info immediately. + Does NOT poll for completion. + """ if not fal_key: - raise RuntimeError("FAL fallback requested but FAL_KEY is missing. Provide fal_api_key or set env FAL_KEY.") - if image_tensor is None or not torch.is_tensor(image_tensor): - raise RuntimeError("FAL fallback requires an input image tensor.") - - batch = image_tensor if image_tensor.ndim == 4 else image_tensor.unsqueeze(0) - # Prepare data URIs (all frames) + raise RuntimeError("FAL requested but FAL_KEY is missing.") + + # Build data URLs for all frames data_urls: T.List[str] = [] + batch = image_tensor if image_tensor.ndim == 4 else image_tensor.unsqueeze(0) for i in range(batch.shape[0]): pil = tensor_to_pil(batch[i:i+1]) raw = encode_pil_bytes(pil, mime) data_urls.append(_data_url(mime, raw)) - + base = "https://queue.fal.run" submit_url = f"{base}/{route.strip()}" headers = {"Authorization": f"Key {fal_key}"} + payload = { "prompt": prompt or "", - # Only plural form to avoid duplicate outputs "image_urls": data_urls, - "num_images": int(max(1, num_images)), - "output_format": ("png" if str(output_format).lower()=="png" else "jpeg"), - "sync_mode": True, + "num_images": 1, + "output_format": ("png" if str(output_format).lower() == "png" else "jpeg"), "aspect_ratio": aspect_ratio, + "sync_mode": sync_mode, } + if debug: - print(f"[PVL FAL] QUEUE SUBMIT {submit_url} with {len(data_urls)} image(s) and num_images={payload['num_images']}") - + print(f"[FAL SUBMIT] prompt: {prompt[:60]}... sync_mode={sync_mode}") + + # Submit request r = requests.post(submit_url, headers=headers, json=payload, timeout=timeout) - if r.status_code >= 400: - raise RuntimeError(f"FAL queue submit error {r.status_code}: {r.text}") - try: - sub = r.json() - except Exception: - raise RuntimeError("FAL queue submit returned non-JSON.") - + if not r.ok: + raise RuntimeError(f"FAL submit error {r.status_code}: {r.text}") + + sub = r.json() req_id = sub.get("request_id") - status_url = sub.get("status_url") or (f"{base}/{route.strip()}/requests/{req_id}/status" if req_id else None) - resp_url = sub.get("response_url") or (f"{base}/{route.strip()}/requests/{req_id}" if req_id else None) - if not req_id or not status_url or not resp_url: - raise RuntimeError("FAL queue submit missing request_id/status_url/response_url.") - - # Poll - deadline = time.time() + max(5, int(timeout)) - last_status = None + if not req_id: + raise RuntimeError("FAL did not return a request_id") + + # Get status and result URLs + status_url = sub.get("status_url") or f"{base}/{route.strip()}/requests/{req_id}/status" + resp_url = sub.get("response_url") or f"{base}/{route.strip()}/requests/{req_id}" + + return { + "request_id": req_id, + "status_url": status_url, + "response_url": resp_url, + "prompt": prompt + } + + def _fal_poll_and_fetch(self, request_info: dict, fal_key: str, timeout: int, debug: bool): + """ + Phase 2: Poll a single FAL request until complete and fetch the result. + Returns (image_tensor, description_text). + """ + headers = {"Authorization": f"Key {fal_key}"} + status_url = request_info["status_url"] + resp_url = request_info["response_url"] + req_id = request_info["request_id"] + + if debug: + print(f"[FAL POLL] request_id={req_id[:16]}...") + + # Poll for completion with timeout check + deadline = time.time() + timeout + completed = False while time.time() < deadline: - sr = requests.get(status_url, headers=headers, timeout=10) - if not sr.ok: - time.sleep(0.5); continue - sdata = sr.json() - last_status = sdata.get("status") - if last_status == "COMPLETED": - break + try: + sr = requests.get(status_url, headers=headers, timeout=10) + if sr.ok and sr.json().get("status") == "COMPLETED": + completed = True + break + except Exception as e: + if debug: + print(f"[FAL POLL] Status check error: {e}") time.sleep(0.6) - if last_status != "COMPLETED": - raise RuntimeError(f"FAL queue did not complete in time (last status={last_status})") - + + # Check if we timed out + if not completed: + raise RuntimeError(f"FAL request {req_id[:16]} timed out after {timeout}s") + # Fetch result rr = requests.get(resp_url, headers=headers, timeout=15) - if rr.status_code >= 400: + if not rr.ok: raise RuntimeError(f"FAL result fetch error {rr.status_code}: {rr.text}") - try: - rdata = rr.json() - except Exception: - raise RuntimeError("FAL result returned non-JSON.") - + + rdata = rr.json() if debug: + print(f"[FAL RESULT] request_id={req_id[:16]}... status=COMPLETED") + + # Extract response data + resp = rdata.get("response", rdata) if isinstance(rdata, dict) else rdata + + # Parse images from various possible locations + buckets = [] + if isinstance(resp, dict): + for key in ("images", "outputs", "artifacts"): + val = resp.get(key) + if isinstance(val, list): + buckets.extend(val) + + for key in ("image", "output", "result"): + val = resp.get(key) + if isinstance(val, (str, dict)): + buckets.append(val) + + out: T.List[torch.Tensor] = [] + for item in buckets: try: - s = json.dumps(rdata)[:1200] - except Exception: - s = str(rdata)[:1200] - print("[PVL FAL] raw response:", s) - - # Normalize containers - resp = rdata.get("response") if isinstance(rdata, dict) else None - if not isinstance(resp, dict): - resp = rdata if isinstance(rdata, dict) else {} - - description = resp.get("description") or rdata.get("description") or resp.get("output_text") or "" - - # Collect potential image items - buckets: T.List[T.Union[str, dict]] = [] - for key in ("images", "outputs", "artifacts"): - val = resp.get(key) - if isinstance(val, list): - buckets.extend(val) - for key in ("image", "output", "result"): - val = resp.get(key) - if isinstance(val, (str, dict)): - buckets.append(val) - for key in ("images", "image", "output", "outputs", "artifacts"): - val = rdata.get(key) if isinstance(rdata, dict) else None - if isinstance(val, list): - buckets.extend(val) - elif isinstance(val, (str, dict)): - buckets.append(val) - - # Decode images, but respect num_images cap - out = [] - def add_image_from_item(item): - try: - if isinstance(item, str): - url_or_data = item - elif isinstance(item, dict): - url_or_data = item.get("url") or item.get("data") or item.get("image") or item.get("content") + url = item if isinstance(item, str) else (item.get("url") or item.get("data") or item.get("image")) + if not url: + continue + + if url.startswith("data:image/"): + blob = base64.b64decode(url.split(",", 1)[1]) else: - return - if not isinstance(url_or_data, str): - return - if url_or_data.startswith("data:image/"): - b64 = url_or_data.split(",", 1)[1] - blob = base64.b64decode(b64) - else: - ir = requests.get(url_or_data, timeout=timeout) + ir = requests.get(url, timeout=timeout) if not ir.ok: - return + continue blob = ir.content + pil = Image.open(io.BytesIO(blob)).convert("RGB") out.append(pil_to_tensor(pil)) except Exception as ex: if debug: - print("[PVL FAL] image decode failed:", ex) - - for item in buckets: - if len(out) >= int(max(1, num_images)): - break - add_image_from_item(item) - + print("[FAL decode fail]", ex) + if not out: - raise RuntimeError("FAL API returned no images.") - - if len(out) == 1: - images_tensor = out[0] - else: - images_tensor = torch.cat(out, dim=0) - - if debug: - print(f"[PVL FAL] returning {len(out)} image(s) (capped to num_images={int(max(1, num_images))})") - + raise RuntimeError(f"FAL returned no images for request_id={req_id}") + + description = "" + if isinstance(resp, dict): + description = resp.get("description") or resp.get("output_text") or "" + + images_tensor = _stack_images_same_size(out, debug) return images_tensor, description - - # ---- main ---- - def run(self, prompt: str, images: T.Optional[torch.Tensor] = None, - aspect_ratio: str = "1:1", + + # --------------------------- RUN MAIN ----------------------------- + + def run(self, prompt: str, images: torch.Tensor, + delimiter: str = "[*]", aspect_ratio: str = "1:1", model: str = DEFAULT_MODEL, endpoint_override: str = "", - api_key: str = "", - temperature: float = 0.6, output_format: str = "png", - capture_text_output: bool = False, num_images: int = 1, - timeout_sec: int = 120, request_id: str = "", - debug_log: bool = False, - use_fal_fallback: bool = False, fal_api_key: str = "", fal_route: str = "fal-ai/nano-banana/edit"): - - if aspect_ratio.strip() not in _VALID_ASPECTS: - print(f"[PVL WARNING] Invalid or missing aspect_ratio '{aspect_ratio}', defaulting to 1:1.") - aspect_ratio = "1:1" - - key = (api_key or os.getenv("GEMINI_API_KEY","")).strip() - input_mime = "image/png" if str(output_format).lower() == "png" else "image/jpeg" + api_key: str = "", temperature: float = 0.6, + output_format: str = "png", capture_text_output: bool = False, + num_images: int = 1, timeout_sec: int = 120, + debug_log: bool = False, use_fal_fallback: bool = True, + force_fal: bool = False, sync_mode: bool = False, + fal_api_key: str = "", fal_route: str = "fal-ai/nano-banana/edit"): + + _t0 = time.time() + + key = (api_key or os.getenv("GEMINI_API_KEY", "")).strip() + input_mime = "image/png" if output_format.lower() == "png" else "image/jpeg" want_text = bool(capture_text_output) - - # If no Google key but fallback is enabled, attempt FAL directly + N = max(1, int(num_images)) + + # Split prompts by literal delimiter and expand to N using reuse rule + raw_parts = [p.strip() for p in str(prompt).split(delimiter) if p.strip()] + if not raw_parts: + raise RuntimeError("Prompt is empty after splitting by delimiter.") + + if len(raw_parts) >= N: + prompts = raw_parts[:N] + else: + prompts = raw_parts + [raw_parts[-1]] * (N - len(raw_parts)) + + if debug_log: + print(f"[PVL Debug] delimiter='{delimiter}' aspect_ratio='{aspect_ratio}' num_images={N} sync_mode={sync_mode}") + for i, pr in enumerate(prompts, 1): + preview = pr if len(pr) <= 160 else (pr[:157] + "...") + print(f"[PVL Debug] Call #{i} prompt: {preview}") + + # Helper function for parallel FAL submission + polling + def parallel_fal_execution(prompts_list, fal_key_str, debug): + """Submit all FAL requests in parallel, then poll all in parallel""" + if debug: + print(f"[FAL] Submitting {len(prompts_list)} requests in parallel...") + + # PHASE 1: Submit all requests IN PARALLEL + submit_results = [] + with ThreadPoolExecutor(max_workers=min(len(prompts_list), 6)) as ex: + submit_futs = { + ex.submit(self._fal_submit_only, fal_route, p, images, input_mime, + fal_key_str, int(timeout_sec), debug, output_format, + aspect_ratio, sync_mode): p + for p in prompts_list + } + for fut in as_completed(submit_futs): + try: + req_info = fut.result() + submit_results.append(req_info) + except Exception as e: + if debug: + print(f"[FAL SUBMIT ERROR] {e}") + + if not submit_results: + raise RuntimeError("All FAL submission requests failed") + + if debug: + print(f"[FAL] {len(submit_results)} requests submitted successfully. Polling for results...") + + # PHASE 2: Poll all requests IN PARALLEL + images_list, texts = [], [] + failed_count = 0 + with ThreadPoolExecutor(max_workers=min(len(submit_results), 6)) as ex: + poll_futs = { + ex.submit(self._fal_poll_and_fetch, req_info, fal_key_str, + int(timeout_sec), debug): req_info + for req_info in submit_results + } + for fut in as_completed(poll_futs): + try: + it, txt = fut.result() + images_list.append(it) + if want_text and txt: + texts.append(txt) + except Exception as e: + failed_count += 1 + if debug: + print(f"[FAL POLL ERROR] {e}") + + if not images_list: + raise RuntimeError(f"All FAL requests failed during polling ({failed_count} failures)") + + if failed_count > 0: + print(f"[PVL WARNING] {failed_count}/{len(submit_results)} FAL requests failed, continuing with {len(images_list)} successful results") + + return images_list, texts + + # Force-FAL path (TRUE PARALLEL) + if force_fal: + fal_key = (fal_api_key or os.getenv("FAL_KEY", "")).strip() + if not fal_key: + raise RuntimeError("force_fal=True but FAL_KEY missing.") + + images_list, texts = parallel_fal_execution(prompts, fal_key, debug_log) + + images_tensor = _stack_images_same_size(images_list, debug_log) + text_out = "\n".join(texts) if (want_text and texts) else "" + + _t1 = time.time() + print(f"[PVL Google NanoBanana mandatory IMG] Execution time: {(_t1 - _t0):.2f}s") + return text_out, images_tensor + + # No API key? Try fallback if enabled (TRUE PARALLEL) if not key: if use_fal_fallback: - fal_key = (fal_api_key or os.getenv("FAL_KEY","")).strip() - try: - img_tensor, fal_text = self._fal_queue_call(fal_route, prompt, images, input_mime, fal_key, int(timeout_sec), debug_log, num_images, output_format) - text_out = (fal_text or "") if want_text else "" - if text_out: - print("[PVL FAL Text]:\n" + text_out) - return (text_out, img_tensor) - except Exception as fe: - print(f"[PVL FAL Fallback] FAL call failed without Google key: {fe}") - raise RuntimeError(f"Gemini image generation failed and FAL fallback also failed: {fe}") + fal_key = (fal_api_key or os.getenv("FAL_KEY", "")).strip() + if not fal_key: + raise RuntimeError("GEMINI_API_KEY missing and FAL_KEY missing.") + + images_list, texts = parallel_fal_execution(prompts, fal_key, debug_log) + + images_tensor = _stack_images_same_size(images_list, debug_log) + text_out = "\n".join(texts) if (want_text and texts) else "" + + _t1 = time.time() + print(f"[PVL Google NanoBanana mandatory IMG] Execution time: {(_t1 - _t0):.2f}s") + return text_out, images_tensor + raise RuntimeError("Gemini API key missing. Pass api_key or set GEMINI_API_KEY.") - - # Build client & request + + # Google SDK path + from google import genai client = self._make_client(key, endpoint_override) - parts = self._build_parts(prompt, images, input_mime) - cfg = self._build_config(temperature, want_text, aspect_ratio) - - if debug_log: - p_preview = (prompt or "")[:180].replace("\n"," ") - img_count = (images.shape[0] if (isinstance(images, torch.Tensor) and images.ndim==4) else (1 if isinstance(images, torch.Tensor) else 0)) - print(f"[PVL Debug] prompt chars={len(prompt or '')} preview='{p_preview}...'") - print(f"[PVL Debug] parts: text={1 if (prompt and prompt.strip()) else 0}, images={img_count}") - safe_parts = [] - for pr in parts: - if "text" in pr: - safe_parts.append({"text": pr["text"][:120]}) - elif "inline_data" in pr: - di = pr["inline_data"] - safe_parts.append({"inline_data": {"mime_type": di.get("mime_type","image/*"), "data": ""}}) - try: - temp = getattr(cfg,'temperature',None) if hasattr(cfg,'temperature') else cfg.get('temperature') - top_p = getattr(cfg,'top_p',None) if hasattr(cfg,'top_p') else cfg.get('top_p') - top_k = getattr(cfg,'top_k',None) if hasattr(cfg,'top_k') else cfg.get('top_k') - mot = getattr(cfg,'max_output_tokens',None) if hasattr(cfg,'max_output_tokens') else cfg.get('max_output_tokens') - mods = getattr(cfg,'response_modalities',None) if hasattr(cfg,'response_modalities') else cfg.get('response_modalities') - except Exception: - temp=top_p=top_k=mot=mods=None - print("[PVL Debug] config:", {"temperature": temp, "top_p": top_p, "top_k": top_k, "max_output_tokens": mot, "modalities": mods}) - print("[PVL Debug] contents:", [{"role":"user","parts": safe_parts}]) - - # Parallel Google calls - N = max(1, int(num_images)) - results = [None] * N - errors = [] - - def call_i(i: int): - rid = (request_id.strip() + f"-{i}") if request_id.strip() else f"pvl-nb-{int(time.time()*1000)}-{i}" - return self._single_google_call(client, model, parts, cfg, rid, debug_log) - - max_workers = min(N, 6) - if N == 1: - try: - results[0] = call_i(0) - except Exception as e: - errors.append(f"google call 0 failed: {e}") - else: - with ThreadPoolExecutor(max_workers=max_workers) as ex: - futmap = {ex.submit(call_i, i): i for i in range(N)} - for fut in as_completed(futmap): - i = futmap[fut] - try: - results[i] = fut.result() - except Exception as e: - errors.append(f"google call {i} failed: {e}") - - # Parse all Google results - out_imgs, out_texts = [], [] - for idx, tup in enumerate(results): - if tup is None: - errors.append(f"google call {idx} returned no response") - continue - imgs_i, texts_i, resp = tup - if not imgs_i: + + def google_call(p: str, idx: int): + parts = self._build_parts(p, images, input_mime) + + cfg = types.GenerateContentConfig( + temperature=float(temperature), + top_p=_TOP_P, + top_k=_TOP_K, + max_output_tokens=_MAX_TOKENS, + response_modalities=["Image"], + image_config=types.ImageConfig(aspect_ratio=aspect_ratio), + ) + + if debug_log: + print(f"[GOOGLE SUBMIT] Call #{idx+1} model={model}") + + resp = client.models.generate_content( + model=model, + contents=[{"role": "user", "parts": parts}], + config=cfg, + ) + + imgs, texts = [], [] + for cand in getattr(resp, "candidates", []) or []: + content = getattr(cand, "content", None) + parts_out = getattr(content, "parts", []) if content else [] + + for prt in parts_out: + blob = _extract_image_bytes_from_part(prt) + if blob: + try: + pil = Image.open(io.BytesIO(blob)).convert("RGB") + imgs.append(pil_to_tensor(pil)) + except Exception as e: + if debug_log: + print("[Decode fail]", e) + else: + t = _extract_text_from_part(prt) + if t: + texts.append(t) + + return imgs, texts + + out_imgs: T.List[torch.Tensor] = [] + out_texts: T.List[str] = [] + failed_google = 0 + + with ThreadPoolExecutor(max_workers=min(N, 6)) as ex: + futs = {ex.submit(google_call, p, i): i for i, p in enumerate(prompts)} + for fut in as_completed(futs): try: - cands = getattr(resp, "candidates", None) or [] - fin = getattr(cands[0], "finish_reason", None) if cands else None - except Exception: - fin = None - errors.append(f"google call {idx} returned no images (finish_reason={fin})") - else: - out_imgs.extend(imgs_i) - if texts_i: - out_texts.append("\n".join(texts_i)) - - # If Google failed and fallback enabled -> try FAL - if errors and use_fal_fallback: - print("[PVL Fallback] Google call failed; attempting FAL.ai fallback...") - for e in errors: - print("[PVL Google Error]", e) - try: - fal_key = (fal_api_key or os.getenv("FAL_KEY","")).strip() - img_tensor, fal_text = self._fal_queue_call(fal_route, prompt, images, input_mime, fal_key, int(timeout_sec), debug_log, num_images, output_format) - final_text = "" # Start with empty; then merge Google texts if requested - if bool(capture_text_output): - pieces = [] - if out_texts: - pieces.append("\n\n--- Google ---\n\n" + ("\n".join(out_texts))) - if fal_text: - pieces.append("\n\n--- FAL ---\n\n" + fal_text) - final_text = "".join(pieces) - if final_text: - print("[PVL Fallback Note] Combined text:") - print(final_text) - return (final_text, img_tensor) - except Exception as fe: - print("[PVL FAL Error]", fe) - raise RuntimeError("Both Google and FAL failed. See console for details.") - - # If Google produced errors and fallback not used -> raise - if errors: - if out_texts: - print("[PVL Google Text]:\n" + ("\n\n---\n\n".join(out_texts))) - raise RuntimeError("Gemini image generation failed: " + " | ".join(errors[:5])) - + imgs, texts = fut.result() + if imgs: + out_imgs.extend(imgs) + if texts: + out_texts.extend(texts) + except Exception as e: + failed_google += 1 + if debug_log: + print(f"[GOOGLE ERROR] {e}") + + if failed_google > 0: + print(f"[PVL WARNING] {failed_google}/{N} Google requests failed") + if not out_imgs: - raise RuntimeError("Gemini image generation failed: no images across all Google calls.") - - images_tensor = torch.cat(out_imgs, dim=0) if len(out_imgs) > 1 else out_imgs[0] - final_text = ("\n\n---\n\n".join(out_texts)) if (bool(capture_text_output) and out_texts) else "" - if final_text: - print(f"[PVL Google NanoBanana Output]:\n{final_text}\n") - - return (final_text, images_tensor,) + if use_fal_fallback: + # per-prompt FAL fallback (TRUE PARALLEL) + fal_key = (fal_api_key or os.getenv("FAL_KEY", "")).strip() + if not fal_key: + raise RuntimeError("Google returned no images and FAL_KEY missing for fallback.") + + images_list, texts = parallel_fal_execution(prompts, fal_key, debug_log) + + images_tensor = _stack_images_same_size(images_list, debug_log) + text_out = "\n".join(out_texts + texts) if want_text else "" + + _t1 = time.time() + print(f"[PVL Google NanoBanana mandatory IMG] Execution time: {(_t1 - _t0):.2f}s") + return text_out, images_tensor + + raise RuntimeError("Gemini returned no images") + + images_tensor = _stack_images_same_size(out_imgs, debug_log) + text_out = "\n".join(out_texts) if (want_text and out_texts) else "" + + if text_out: + print(f"[PVL Google NanoBanana Output]:\n{text_out}\n") + + _t1 = time.time() + print(f"[PVL Google NanoBanana mandatory IMG] Execution time: {(_t1 - _t0):.2f}s") + return text_out, images_tensor NODE_CLASS_MAPPINGS = {"PVL_Google_NanoBanana_API_mandatory_IMG": PVL_Google_NanoBanana_API_mandatory_IMG} NODE_DISPLAY_NAME_MAPPINGS = {"PVL_Google_NanoBanana_API_mandatory_IMG": NODE_NAME} diff --git a/pvl_google_nano_banana_multi_img.py b/pvl_google_nano_banana_multi_img.py index 6511c2c..38470a1 100644 --- a/pvl_google_nano_banana_multi_img.py +++ b/pvl_google_nano_banana_multi_img.py @@ -1,16 +1,20 @@ # pvl_google_nano_banana_multi_img.py + # PVL Google Nano-Banana Multi API — with regex delimiter support + # Author: PVL # License: MIT -# + # Features: # - Text box for prompt input ("STRING"). # - Regex-based delimiter input (e.g., \n|\| or ;+). # - num_images controls number of parallel API calls. # - If prompts < num_images, reuses the *last* prompt to fill the rest, with a warning. # - Parallel calls for Google; optional FAL-only mode and Google→FAL fallback. +# - TRUE PARALLEL FAL execution: submit all requests first, then poll for results. # - Optionally capture text output and print to console. # - Ensures IMAGE output is a 4D tensor (B, H, W, C). +# - sync_mode toggle for FAL API (default: False) import os, io, json, base64, typing as T, time, re from concurrent.futures import ThreadPoolExecutor, as_completed @@ -23,30 +27,24 @@ import torch NODE_NAME = "PVL Google Nano-Banana Multi API" NODE_CATEGORY = "PVL/Google" DEFAULT_MODEL = "gemini-2.5-flash-image-preview" - _TOP_P = 0.95 _TOP_K = 64 _MAX_TOKENS = 4096 - # --------------------------- Image helpers --------------------------- def pil_to_tensor(img: Image.Image) -> torch.Tensor: if img.mode != "RGB": img = img.convert("RGB") arr = np.asarray(img, dtype=np.float32) / 255.0 - # Shape (H,W,C) -> (1,H,W,C) as ComfyUI image tensor return torch.from_numpy(arr)[None, ...] - def tensor_to_pil(t: torch.Tensor) -> Image.Image: - # Accept (B,H,W,C) or (H,W,C) if t.ndim == 4: t = t[0] arr = (t.clamp(0, 1).cpu().numpy() * 255).astype("uint8") return Image.fromarray(arr, "RGB") - def encode_pil_bytes(img: Image.Image, mime: str) -> bytes: buf = io.BytesIO() if mime == "image/jpeg": @@ -55,118 +53,49 @@ def encode_pil_bytes(img: Image.Image, mime: str) -> bytes: img.save(buf, format="PNG") return buf.getvalue() - def _extract_image_bytes_from_part(part) -> T.Optional[bytes]: - """ - Enhanced to handle more image data formats from Gemini API response - """ - # Debug: Print the part structure - print(f"[EXTRACT DEBUG] Part type: {type(part)}") - + """Extract image bytes from Gemini API response part.""" try: - # Try to get inline_data in various formats - inline = None - if hasattr(part, 'inline_data'): - inline = getattr(part, 'inline_data') - print(f"[EXTRACT DEBUG] Found inline_data attribute") - elif hasattr(part, 'inlineData'): - inline = getattr(part, 'inlineData') - print(f"[EXTRACT DEBUG] Found inlineData attribute") - elif isinstance(part, dict) and 'inline_data' in part: - inline = part['inline_data'] - print(f"[EXTRACT DEBUG] Found inline_data dict key") - elif isinstance(part, dict) and 'inlineData' in part: - inline = part['inlineData'] - print(f"[EXTRACT DEBUG] Found inlineData dict key") - + inline = getattr(part, "inline_data", None) or getattr(part, "inlineData", None) if inline is not None: - # Try to get data from inline - data = None - if hasattr(inline, 'data'): - data = getattr(inline, 'data') - print(f"[EXTRACT DEBUG] Found data attribute, type: {type(data)}") - elif isinstance(inline, dict) and 'data' in inline: - data = inline['data'] - print(f"[EXTRACT DEBUG] Found data dict key, type: {type(data)}") - - if data is not None: - if isinstance(data, str): - try: - result = base64.b64decode(data, validate=False) - print(f"[EXTRACT DEBUG] Successfully decoded base64 string of length {len(data)}") - return result - except Exception as e: - print(f"[EXTRACT DEBUG] Failed to decode base64: {e}") - elif isinstance(data, bytes): - print(f"[EXTRACT DEBUG] Found raw bytes of length {len(data)}") - return data - except Exception as e: - print(f"[EXTRACT DEBUG] Exception during extraction: {e}") - - # Try alternative approaches - try: - # Check if the part itself has image data - if hasattr(part, 'image_bytes'): - image_bytes = getattr(part, 'image_bytes') - print(f"[EXTRACT DEBUG] Found image_bytes attribute") - if isinstance(image_bytes, bytes): - return image_bytes + data = getattr(inline, "data", None) + if isinstance(data, (bytes, bytearray)): + return bytes(data) + if isinstance(data, str): + try: + return base64.b64decode(data, validate=False) + except Exception: + return None + except Exception: + pass + + if isinstance(part, dict): + blob = None + if "inline_data" in part and isinstance(part["inline_data"], dict): + blob = part["inline_data"].get("data") + elif "inlineData" in part and isinstance(part["inlineData"], dict): + blob = part["inlineData"].get("data") - # Check if the part is a dict with image data - if isinstance(part, dict): - for key in ['image_bytes', 'image_data', 'binary_data']: - if key in part: - data = part[key] - print(f"[EXTRACT DEBUG] Found {key} dict key") - if isinstance(data, str): - try: - return base64.b64decode(data, validate=False) - except Exception: - pass - elif isinstance(data, bytes): - return data - except Exception as e: - print(f"[EXTRACT DEBUG] Exception during alternative extraction: {e}") + if isinstance(blob, (bytes, bytearray)): + return bytes(blob) + if isinstance(blob, str): + try: + return base64.b64decode(blob, validate=False) + except Exception: + return None - # Last resort: try to find any base64 string in the part - try: - if isinstance(part, dict): - for key, value in part.items(): - if isinstance(value, str) and len(value) > 100: # Likely a base64 encoded image - try: - # Check if it looks like base64 - if value.startswith('data:image/') or re.match(r'^[A-Za-z0-9+/]+={0,2}$', value): - # If it's a data URL, extract the base64 part - if value.startswith('data:image/'): - base64_data = value.split(',', 1)[1] - else: - base64_data = value - - result = base64.b64decode(base64_data, validate=False) - print(f"[EXTRACT DEBUG] Found and decoded base64 in dict key '{key}'") - return result - except Exception: - pass - except Exception as e: - print(f"[EXTRACT DEBUG] Exception during base64 search: {e}") - - print("[EXTRACT DEBUG] No image data found") return None - def _extract_text_from_part(part) -> T.Optional[str]: - # Prefer dict, fallback to attribute getattr if isinstance(part, dict): if "text" in part and part["text"] is not None: return str(part["text"]) txt = getattr(part, "text", None) return str(txt) if txt is not None else None - def _data_url(mime: str, raw: bytes) -> str: return f"data:{mime};base64," + base64.b64encode(raw).decode("utf-8") - def _stack_images_same_size(tensors: T.List[torch.Tensor], debug: bool = False) -> torch.Tensor: """ Concatenate (B,H,W,C) batches along B. If shapes mismatch, resize to the first image size. @@ -186,16 +115,16 @@ def _stack_images_same_size(tensors: T.List[torch.Tensor], debug: bool = False) fixed.append(pil_to_tensor(rp)) return torch.cat(fixed, dim=0) - # --------------------------- Main Node --------------------------- class PVL_Google_NanoBanana_Multi_API: + @classmethod def INPUT_TYPES(cls): return { "required": { "prompt": ("STRING", {"multiline": True, "default": "", "placeholder": "Enter prompts separated by regex delimiter"}), - "delimiter": ("STRING", {"default": "\\n-----\\n", "multiline": False, "placeholder": "Regex (e.g. \\n|\\| or ;+)"}), + "delimiter": ("STRING", {"default": "[*]", "multiline": False, "placeholder": "Regex (e.g. \\n|\\| or ;+)"}), }, "optional": { "image_1": ("IMAGE",), @@ -210,45 +139,50 @@ class PVL_Google_NanoBanana_Multi_API: "model": ("STRING", {"default": DEFAULT_MODEL}), "endpoint_override": ("STRING", {"default": ""}), "api_key": ("STRING", {"default": "", "multiline": False, - "placeholder": "Leave empty to use GEMINI_API_KEY"}), + "placeholder": "Leave empty to use GEMINI_API_KEY"}), "temperature": ("FLOAT", {"default": 0.6, "min": 0.0, "max": 2.0, "step": 0.05}), "output_format": (["png", "jpeg"], {"default": "png"}), "capture_text_output": ("BOOLEAN", {"default": False}), "num_images": ("INT", {"default": 1, "min": 1, "max": 12, "step": 1}), "timeout_sec": ("INT", {"default": 120, "min": 5, "max": 600, "step": 5}), "debug_log": ("BOOLEAN", {"default": False}), - "use_fal_fallback": ("BOOLEAN", {"default": False}), + "use_fal_fallback": ("BOOLEAN", {"default": True}), "force_fal": ("BOOLEAN", {"default": False}), + "sync_mode": ("BOOLEAN", {"default": False}), "fal_api_key": ("STRING", {"default": "", "multiline": False, - "placeholder": "Leave empty to use FAL_KEY"}), + "placeholder": "Leave empty to use FAL_KEY"}), "fal_route": ("STRING", {"default": "fal-ai/nano-banana/edit"}), } } - + RETURN_TYPES = ("STRING", "IMAGE",) RETURN_NAMES = ("text", "images") FUNCTION = "run" CATEGORY = NODE_CATEGORY - + # ------------------------- INTERNAL HELPERS ------------------------- - + def _make_client(self, api_key: str, endpoint_override: str): from google import genai from google.genai import types + http_options = None if endpoint_override.strip(): try: http_options = types.HttpOptions(base_url=endpoint_override.strip()) except Exception: http_options = None + if http_options is not None: return genai.Client(api_key=api_key, http_options=http_options) return genai.Client(api_key=api_key) - + def _build_parts(self, prompt: str, image_tensors: T.List[torch.Tensor], mime: str): parts: T.List[dict] = [] + if prompt and prompt.strip(): parts.append({"text": prompt}) + for img_tensor in image_tensors: if img_tensor.ndim == 4: for i in range(img_tensor.shape[0]): @@ -258,41 +192,49 @@ class PVL_Google_NanoBanana_Multi_API: else: pil = tensor_to_pil(img_tensor) parts.append({"inline_data": {"mime_type": mime, "data": encode_pil_bytes(pil, mime)}}) + return parts - + def _build_call_prompts(self, base_prompts: T.List[str], num_images: int, debug: bool) -> T.List[str]: """ Maps prompts to calls according to the agreed rule: - - If len(prompts) >= num_images → take first num_images - - If len(prompts) < num_images → repeat the *last* prompt to fill + - If len(prompts) >= num_images → take first num_images + - If len(prompts) < num_images → repeat the *last* prompt to fill """ N = max(1, int(num_images)) + if not base_prompts: return [] + if len(base_prompts) >= N: call_prompts = base_prompts[:N] else: if debug: print(f"[PVL NODE] Provided {len(base_prompts)} prompts but num_images={N}. " f"Reusing the last prompt for remaining calls.") - print(f"[PVL WARNING] prompt list shorter than num_images: {len(base_prompts)} < {N}. " - f"Last entry will be reused for the remaining {N - len(base_prompts)} calls.") + print(f"[PVL WARNING] prompt list shorter than num_images: {len(base_prompts)} < {N}. " + f"Last entry will be reused for the remaining {N - len(base_prompts)} calls.") call_prompts = base_prompts + [base_prompts[-1]] * (N - len(base_prompts)) + if debug: for i, cp in enumerate(call_prompts, 1): show = cp if len(cp) <= 160 else (cp[:157] + "...") print(f"[PVL NODE] Call #{i} prompt: {show}") + return call_prompts - - def _fal_queue_call(self, route: str, prompt: str, image_tensors: T.List[torch.Tensor], - mime: str, fal_key: str, timeout: int, debug: bool, - output_format: str, aspect_ratio: str = "1:1"): + + # -------- FAL API - TWO PHASE EXECUTION -------- + + def _fal_submit_only(self, route: str, prompt: str, image_tensors: T.List[torch.Tensor], + mime: str, fal_key: str, timeout: int, debug: bool, + output_format: str, aspect_ratio: str = "1:1", sync_mode: bool = False): """ - Calls FAL queue endpoint in sync_mode, returns (image_tensor_batched1, description_text). + Phase 1: Submit request to FAL queue and return request info immediately. + Does NOT poll for completion. """ if not fal_key: raise RuntimeError("FAL requested but FAL_KEY is missing.") - + # Convert any input images to data URLs data_urls = [] for t in image_tensors: @@ -306,59 +248,89 @@ class PVL_Google_NanoBanana_Multi_API: pil = tensor_to_pil(t) raw = encode_pil_bytes(pil, mime) data_urls.append(_data_url(mime, raw)) - + base = "https://queue.fal.run" submit_url = f"{base}/{route.strip()}" headers = {"Authorization": f"Key {fal_key}"} + payload = { "prompt": prompt or "", "image_urls": data_urls, "num_images": 1, "output_format": ("png" if str(output_format).lower() == "png" else "jpeg"), "aspect_ratio": aspect_ratio, - "sync_mode": True, + "sync_mode": sync_mode, } + if debug: - print("[FAL SUBMIT]", json.dumps(payload)[:1000]) - + print(f"[FAL SUBMIT] prompt: {prompt[:60]}... sync_mode={sync_mode}") + r = requests.post(submit_url, headers=headers, json=payload, timeout=timeout) if not r.ok: raise RuntimeError(f"FAL submit error {r.status_code}: {r.text}") + sub = r.json() req_id = sub.get("request_id") - status_url = sub.get("status_url") or (f"{base}/{route.strip()}/requests/{req_id}/status" if req_id else None) - resp_url = sub.get("response_url") or (f"{base}/{route.strip()}/requests/{req_id}" if req_id else None) - if not req_id or not status_url or not resp_url: - raise RuntimeError("FAL queue missing request_id/status/response") - + if not req_id: + raise RuntimeError("FAL did not return a request_id") + + status_url = sub.get("status_url") or f"{base}/{route.strip()}/requests/{req_id}/status" + resp_url = sub.get("response_url") or f"{base}/{route.strip()}/requests/{req_id}" + + return { + "request_id": req_id, + "status_url": status_url, + "response_url": resp_url, + "prompt": prompt + } + + def _fal_poll_and_fetch(self, request_info: dict, fal_key: str, timeout: int, debug: bool): + """ + Phase 2: Poll a single FAL request until complete and fetch the result. + Returns (image_tensor, description_text). + """ + headers = {"Authorization": f"Key {fal_key}"} + status_url = request_info["status_url"] + resp_url = request_info["response_url"] + req_id = request_info["request_id"] + + if debug: + print(f"[FAL POLL] request_id={req_id[:16]}...") + + # Poll for completion deadline = time.time() + timeout while time.time() < deadline: sr = requests.get(status_url, headers=headers, timeout=10) if sr.ok and sr.json().get("status") == "COMPLETED": break time.sleep(0.6) - + + # Fetch result rr = requests.get(resp_url, headers=headers, timeout=15) if not rr.ok: raise RuntimeError(f"FAL result fetch error {rr.status_code}: {rr.text}") + rdata = rr.json() if debug: - print("[FAL RAW]", json.dumps(rdata)[:2000]) - + print(f"[FAL RESULT] request_id={req_id[:16]}... status=COMPLETED") + + # Extract response data buckets = [] resp = rdata.get("response") if isinstance(rdata, dict) else None if resp is None and isinstance(rdata, dict): resp = rdata + if isinstance(resp, dict): for key in ("images", "outputs", "artifacts"): val = resp.get(key) if isinstance(val, list): buckets.extend(val) + for key in ("image", "output", "result"): val = resp.get(key) if isinstance(val, (str, dict)): buckets.append(val) - + out = [] for item in buckets: try: @@ -366,6 +338,7 @@ class PVL_Google_NanoBanana_Multi_API: else (item.get("url") or item.get("data") or item.get("image")) if not url: continue + if url.startswith("data:image/"): blob = base64.b64decode(url.split(",", 1)[1]) else: @@ -373,19 +346,21 @@ class PVL_Google_NanoBanana_Multi_API: if not ir.ok: continue blob = ir.content + pil = Image.open(io.BytesIO(blob)).convert("RGB") out.append(pil_to_tensor(pil)) except Exception as ex: if debug: - print("[FAL decode fail]", ex) - + print(f"[FAL decode fail] {ex}") + if not out: - raise RuntimeError("FAL returned no images") - - return out[0], (resp.get("description", "") if isinstance(resp, dict) else "") - + raise RuntimeError(f"FAL returned no images for request_id={req_id}") + + description = resp.get("description", "") if isinstance(resp, dict) else "" + return out[0], description + # --------------------------- RUN MAIN ----------------------------- - + def run(self, prompt: str, delimiter: str, image_1=None, image_2=None, image_3=None, image_4=None, image_5=None, image_6=None, image_7=None, image_8=None, @@ -395,7 +370,9 @@ class PVL_Google_NanoBanana_Multi_API: capture_text_output: bool = False, num_images: int = 1, timeout_sec: int = 120, debug_log: bool = False, use_fal_fallback: bool = False, force_fal: bool = False, + sync_mode: bool = False, fal_api_key: str = "", fal_route: str = "fal-ai/nano-banana/edit"): + # --- Validate aspect ratio (for Google & FAL) --- ar_original = aspect_ratio.strip() valid_ratios = {"21:9","1:1","4:3","3:2","2:3","5:4","4:5","3:4","16:9","9:16"} @@ -404,163 +381,181 @@ class PVL_Google_NanoBanana_Multi_API: aspect_ratio = "1:1" else: aspect_ratio = ar_original - + # --- Regex-based prompt splitting (with safe fallback to literal split) --- try: base_prompts = [p.strip() for p in re.split(delimiter, prompt) if str(p).strip()] except re.error: print(f"[PVL WARNING] Invalid regex pattern '{delimiter}', using literal split.") base_prompts = [p.strip() for p in prompt.split(delimiter) if str(p).strip()] - + if not base_prompts: raise RuntimeError("No valid prompts provided.") - + # Collect any provided images (tensors) in order image_tensors: T.List[torch.Tensor] = [] for img in [image_1, image_2, image_3, image_4, image_5, image_6, image_7, image_8]: if img is not None and torch.is_tensor(img): image_tensors.append(img) - + # Map provided prompts to num_images calls call_prompts = self._build_call_prompts(base_prompts, num_images, debug_log) - + input_mime = "image/png" if output_format.lower() == "png" else "image/jpeg" want_text = bool(capture_text_output) - - # ---------------- FAL-only path ---------------- + + # ---------------- FAL-only path (TRUE PARALLEL) ---------------- if force_fal: fal_key = (fal_api_key or os.getenv("FAL_KEY", "")).strip() if not fal_key: raise RuntimeError("force_fal=True but no FAL_KEY provided.") + + if debug_log: + print(f"[FAL] Submitting {len(call_prompts)} requests...") + + # PHASE 1: Submit all requests (fast, non-blocking) + request_infos = [] + for p in call_prompts: + req_info = self._fal_submit_only(fal_route, p, image_tensors, + input_mime, fal_key, timeout_sec, + debug_log, output_format, + aspect_ratio, sync_mode) + request_infos.append(req_info) + + if debug_log: + print(f"[FAL] All {len(request_infos)} requests submitted. Polling for results...") + + # PHASE 2: Poll all requests in parallel results, texts = [], [] - with ThreadPoolExecutor(max_workers=min(len(call_prompts), 6)) as ex: + with ThreadPoolExecutor(max_workers=min(len(request_infos), 6)) as ex: futs = { - ex.submit(self._fal_queue_call, fal_route, p, image_tensors, - input_mime, fal_key, timeout_sec, debug_log, output_format, aspect_ratio): p - for p in call_prompts + ex.submit(self._fal_poll_and_fetch, req_info, fal_key, + timeout_sec, debug_log): req_info + for req_info in request_infos } for fut in as_completed(futs): img, t = fut.result() - results.append(img) # Fixed: removed .unsqueeze(0) + results.append(img) if t: texts.append(t) + images_tensor = _stack_images_same_size(results, debug_log) text_out = "\n".join(texts) if want_text else "" return text_out, images_tensor - + # ---------------- Google path (with optional FAL fallback) ---------------- key = (api_key or os.getenv("GEMINI_API_KEY", "")).strip() if not key: raise RuntimeError("Gemini API key missing. Pass api_key or set GEMINI_API_KEY.") + client = self._make_client(key, endpoint_override) - + def google_call(p: str, debug: bool): - # Prepend explicit image generation instruction - image_prompt = f"{p}" - if debug: - print(f"[GEMINI PROMPT] {image_prompt}") + parts = self._build_parts(p, image_tensors, input_mime) - parts = self._build_parts(image_prompt, image_tensors, input_mime) - - # Correct configuration for image generation cfg = types.GenerateContentConfig( temperature=float(temperature), top_p=_TOP_P, top_k=_TOP_K, max_output_tokens=_MAX_TOKENS, - response_modalities=["Image"], # Request image generation - # No response_mime_type - model returns PNG by default + response_modalities=["Image"], image_config=types.ImageConfig(aspect_ratio=aspect_ratio), ) - + + if debug: + print(f"[GOOGLE SUBMIT] prompt: {p[:100]}...") + resp = client.models.generate_content( model=model, contents=[{"role": "user", "parts": parts}], config=cfg, ) - if debug: - print("[GEMINI RESPONSE STRUCTURE]") - print(f" Candidates: {len(getattr(resp, 'candidates', []))}") - for i, cand in enumerate(getattr(resp, 'candidates', []) or []): - content = getattr(cand, 'content', None) - parts_out = getattr(content, 'parts', []) if content else [] - print(f" Candidate {i}: {len(parts_out)} parts") - for j, prt in enumerate(parts_out): - # Enhanced debugging - print(f" Part {j}: Type = {type(prt)}") - if hasattr(prt, '__dict__'): - print(f" Attributes: {list(prt.__dict__.keys())}") - elif isinstance(prt, dict): - print(f" Keys: {list(prt.keys())}") - imgs, texts = [], [] for cand in getattr(resp, "candidates", []) or []: content = getattr(cand, "content", None) parts_out = getattr(content, "parts", []) if content else [] + for prt in parts_out: blob = _extract_image_bytes_from_part(prt) if blob: - if debug: - print(f" Extracted image blob of size {len(blob)} bytes") try: pil = Image.open(io.BytesIO(blob)).convert("RGB") imgs.append(pil_to_tensor(pil)) except Exception as e: - print(f" Failed to convert blob to image: {e}") + if debug: + print(f"[Decode fail] {e}") else: t = _extract_text_from_part(prt) if t: - if debug: - print(f" Extracted text: {t[:100]}...") texts.append(t) if debug: - print(f"[GEMINI RESULT] Found {len(imgs)} images and {len(texts)} text parts") + print(f"[GOOGLE RESULT] Found {len(imgs)} images") return imgs, texts - + out_imgs: T.List[torch.Tensor] = [] out_texts: T.List[str] = [] + with ThreadPoolExecutor(max_workers=min(len(call_prompts), 6)) as ex: futs = {ex.submit(google_call, p, debug_log): p for p in call_prompts} for fut in as_completed(futs): imgs, texts = fut.result() if imgs: - # Fixed: removed .unsqueeze(0) since pil_to_tensor already adds batch dimension out_imgs.append(imgs[0]) if texts: out_texts.extend(texts) - + if not out_imgs: if use_fal_fallback: fal_key = (fal_api_key or os.getenv("FAL_KEY", "")).strip() if not fal_key: raise RuntimeError("Google returned no images and FAL fallback has no FAL_KEY.") + + if debug_log: + print(f"[FAL FALLBACK] Submitting {len(call_prompts)} requests...") + + # PHASE 1: Submit all requests + request_infos = [] + for p in call_prompts: + req_info = self._fal_submit_only(fal_route, p, image_tensors, + input_mime, fal_key, timeout_sec, + debug_log, output_format, + aspect_ratio, sync_mode) + request_infos.append(req_info) + + if debug_log: + print(f"[FAL FALLBACK] All requests submitted. Polling for results...") + + # PHASE 2: Poll all requests in parallel results, texts = [], [] - with ThreadPoolExecutor(max_workers=min(len(call_prompts), 6)) as ex: + with ThreadPoolExecutor(max_workers=min(len(request_infos), 6)) as ex: futs = { - ex.submit(self._fal_queue_call, fal_route, p, image_tensors, - input_mime, fal_key, timeout_sec, debug_log, output_format, aspect_ratio): p - for p in call_prompts + ex.submit(self._fal_poll_and_fetch, req_info, fal_key, + timeout_sec, debug_log): req_info + for req_info in request_infos } for fut in as_completed(futs): img, t = fut.result() - results.append(img) # Fixed: removed .unsqueeze(0) + results.append(img) if t: texts.append(t) + images_tensor = _stack_images_same_size(results, debug_log) combo_texts = out_texts + texts text_out = "\n".join(combo_texts) if want_text else "" return text_out, images_tensor + raise RuntimeError("Gemini returned no images") - + images_tensor = _stack_images_same_size(out_imgs, debug_log) text_out = "\n".join(out_texts) if (want_text and out_texts) else "" - if text_out: + + if text_out and debug_log: print(f"[PVL Google NanoBanana Output]:\n{text_out}\n") + return text_out, images_tensor - NODE_CLASS_MAPPINGS = {"PVL_Google_NanoBanana_Multi_API": PVL_Google_NanoBanana_Multi_API} -NODE_DISPLAY_NAME_MAPPINGS = {"PVL_Google_NanoBanana_Multi_API": NODE_NAME} \ No newline at end of file +NODE_DISPLAY_NAME_MAPPINGS = {"PVL_Google_NanoBanana_Multi_API": NODE_NAME}