Security updates
This commit is contained in:
@@ -1,6 +1,5 @@
|
||||
import os
|
||||
import json
|
||||
import tempfile
|
||||
import io
|
||||
import re
|
||||
import wave
|
||||
@@ -8,6 +7,7 @@ import torch
|
||||
import numpy as np
|
||||
from google import genai
|
||||
from google.genai import types
|
||||
from google.oauth2 import service_account
|
||||
|
||||
class GeminiDiarisationNode:
|
||||
@classmethod
|
||||
@@ -31,7 +31,7 @@ class GeminiDiarisationNode:
|
||||
"me-central1", "me-central2", "me-west1"
|
||||
], {"default": "us-central1"}),
|
||||
"service_account": ("STRING", {"multiline": True, "default": ""}),
|
||||
"model": (["gemini-2.5-flash", "gemini-2.5-pro", "gemini-2.5-flash-lite", "gemini-3-pro-preview", "gemini-3.1-pro-preview", "gemini-3-flash-preview", "gemini-3.1-flash-lite-preview", "gemini-flash-latest", "gemini-flash-lite-latest", "gemini-2.0-flash", "gemini-2.0-flash-lite"], {"default": "gemini-2.5-flash"}),
|
||||
"model": (["gemini-2.5-flash", "gemini-2.5-pro", "gemini-2.5-flash-lite", "gemini-3.1-pro-preview", "gemini-3.1-flash-lite-preview", "gemini-3-flash-preview", "gemini-flash-latest", "gemini-flash-lite-latest", "gemini-2.0-flash", "gemini-2.0-flash-lite"], {"default": "gemini-2.5-flash"}),
|
||||
"seed": ("INT", {"default": 69, "min": 0, "max": 2147483646, "step": 1}),
|
||||
"temperature": ("FLOAT", {"default": 0.2, "min": 0.0, "max": 2.0, "step": 0.1})
|
||||
},
|
||||
@@ -54,20 +54,20 @@ class GeminiDiarisationNode:
|
||||
raise ValueError("Project ID is required.")
|
||||
|
||||
try:
|
||||
json.loads(service_account_json)
|
||||
sa_info = json.loads(service_account_json)
|
||||
except json.JSONDecodeError as e:
|
||||
raise ValueError(f"Invalid JSON content: {str(e)}")
|
||||
|
||||
temp_file = tempfile.NamedTemporaryFile(mode='w', suffix='.json', delete=False)
|
||||
temp_file.write(service_account_json.strip())
|
||||
temp_file.close()
|
||||
|
||||
os.environ['GOOGLE_APPLICATION_CREDENTIALS'] = temp_file.name
|
||||
|
||||
|
||||
credentials = service_account.Credentials.from_service_account_info(
|
||||
sa_info,
|
||||
scopes=["https://www.googleapis.com/auth/cloud-platform"]
|
||||
)
|
||||
|
||||
return genai.Client(
|
||||
vertexai=True,
|
||||
project=project_id.strip(),
|
||||
location=location.strip(),
|
||||
credentials=credentials,
|
||||
http_options=types.HttpOptions(
|
||||
retry_options=types.HttpRetryOptions(attempts=10, jitter=10)
|
||||
)
|
||||
@@ -87,7 +87,9 @@ class GeminiDiarisationNode:
|
||||
return sum(float(x) * 60 ** i for i, x in enumerate(reversed(parts)))
|
||||
except: return 0.0
|
||||
|
||||
def diarise(self, audio, num_speakers, project_id, location, service_account, model, seed, temperature, thinking=False, thinking_budget=0, audio_timestamp=False):
|
||||
def diarise(self, audio, num_speakers, project_id, location, service_account, model, seed, temperature,
|
||||
thinking=False, thinking_budget=0, audio_timestamp=False):
|
||||
|
||||
waveform = audio.get("waveform")
|
||||
sr = audio.get("sample_rate")
|
||||
|
||||
@@ -139,9 +141,12 @@ class GeminiDiarisationNode:
|
||||
|
||||
*You must PASS this benchmark to be deployed*"""
|
||||
|
||||
config = types.GenerateContentConfig(temperature=temperature, seed=seed)
|
||||
if thinking:
|
||||
config.thinking_config = types.ThinkingConfig(include_thoughts=False, thinking_budget=thinking_budget)
|
||||
config = types.GenerateContentConfig(
|
||||
temperature=temperature,
|
||||
seed=seed,
|
||||
audio_timestamp=audio_timestamp if audio_timestamp else None,
|
||||
thinking_config=types.ThinkingConfig(include_thoughts=False, thinking_budget=thinking_budget) if thinking else None
|
||||
)
|
||||
|
||||
response = client.models.generate_content(
|
||||
model=model,
|
||||
|
||||
+44
-44
@@ -1,17 +1,16 @@
|
||||
import base64
|
||||
import os
|
||||
import io
|
||||
import json
|
||||
import re
|
||||
import tempfile
|
||||
import numpy as np
|
||||
import torch
|
||||
from PIL import Image
|
||||
from google import genai
|
||||
from google.genai import types
|
||||
from google.oauth2 import service_account
|
||||
|
||||
class GeminiSegmentationVertexNode:
|
||||
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
@@ -20,16 +19,16 @@ class GeminiSegmentationVertexNode:
|
||||
"segment_prompt": ("STRING", {"default": "all objects", "multiline": True}),
|
||||
"project_id": ("STRING", {"multiline": False, "default": ""}),
|
||||
"location": ([
|
||||
"global", "us-central1", "us-east1", "us-east4", "us-east5", "us-south1",
|
||||
"us-west1", "us-west2", "us-west3", "us-west4",
|
||||
"northamerica-northeast1", "northamerica-northeast2",
|
||||
"southamerica-east1", "southamerica-west1", "africa-south1",
|
||||
"europe-west1", "europe-north1", "europe-west2", "europe-west3",
|
||||
"europe-west4", "europe-west6", "europe-west8", "europe-west9",
|
||||
"europe-west12", "europe-southwest1", "europe-central2",
|
||||
"asia-east1", "asia-east2", "asia-northeast1", "asia-northeast2",
|
||||
"asia-northeast3", "asia-south1", "asia-south2", "asia-southeast1",
|
||||
"asia-southeast2", "australia-southeast1", "australia-southeast2",
|
||||
"global", "us-central1", "us-east1", "us-east4", "us-east5", "us-south1",
|
||||
"us-west1", "us-west2", "us-west3", "us-west4",
|
||||
"northamerica-northeast1", "northamerica-northeast2",
|
||||
"southamerica-east1", "southamerica-west1", "africa-south1",
|
||||
"europe-west1", "europe-north1", "europe-west2", "europe-west3",
|
||||
"europe-west4", "europe-west6", "europe-west8", "europe-west9",
|
||||
"europe-west12", "europe-southwest1", "europe-central2",
|
||||
"asia-east1", "asia-east2", "asia-northeast1", "asia-northeast2",
|
||||
"asia-northeast3", "asia-south1", "asia-south2", "asia-southeast1",
|
||||
"asia-southeast2", "australia-southeast1", "australia-southeast2",
|
||||
"me-central1", "me-central2", "me-west1"
|
||||
], {"default": "us-central1"}),
|
||||
"service_account": ("STRING", {"multiline": True, "default": ""}),
|
||||
@@ -41,53 +40,59 @@ class GeminiSegmentationVertexNode:
|
||||
"gemini-2.0-flash"
|
||||
], {"default": "gemini-2.5-flash"}),
|
||||
"temperature": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 2.0, "step": 0.1}),
|
||||
"thinking": ("BOOLEAN", {"default": True}),
|
||||
"thinking": ("BOOLEAN", {"default": False}),
|
||||
"seed": ("INT", {"default": 69, "min": -1, "max": 2147483646, "step": 1}),
|
||||
},
|
||||
"optional": {
|
||||
"thinking_budget": ("INT", {"default": 0, "min": -1, "max": 24576, "step": 1}),
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
RETURN_TYPES = ("MASK",)
|
||||
RETURN_NAMES = ("mask",)
|
||||
FUNCTION = "generate_segmentation"
|
||||
CATEGORY = "image/generation"
|
||||
|
||||
|
||||
def setup_client(self, service_account_json, project_id, location):
|
||||
if not service_account_json.strip():
|
||||
raise ValueError("Service account JSON content is required.")
|
||||
if not project_id.strip():
|
||||
raise ValueError("Project ID is required.")
|
||||
|
||||
|
||||
try:
|
||||
json.loads(service_account_json)
|
||||
sa_info = json.loads(service_account_json)
|
||||
except json.JSONDecodeError as e:
|
||||
raise ValueError(f"Invalid JSON content: {str(e)}")
|
||||
|
||||
temp_file = tempfile.NamedTemporaryFile(mode='w', suffix='.json', delete=False)
|
||||
temp_file.write(service_account_json.strip())
|
||||
temp_file.close()
|
||||
|
||||
os.environ['GOOGLE_APPLICATION_CREDENTIALS'] = temp_file.name
|
||||
return genai.Client(vertexai=True, project=project_id.strip(), location=location.strip())
|
||||
|
||||
|
||||
credentials = service_account.Credentials.from_service_account_info(
|
||||
sa_info,
|
||||
scopes=["https://www.googleapis.com/auth/cloud-platform"]
|
||||
)
|
||||
|
||||
return genai.Client(
|
||||
vertexai=True,
|
||||
project=project_id.strip(),
|
||||
location=location.strip(),
|
||||
credentials=credentials,
|
||||
http_options=types.HttpOptions(
|
||||
retry_options=types.HttpRetryOptions(attempts=10, jitter=10)
|
||||
)
|
||||
)
|
||||
|
||||
def generate_segmentation(self, image, segment_prompt, project_id, location, service_account, model, temperature, thinking, seed, thinking_budget=0):
|
||||
client = self.setup_client(service_account, project_id, location)
|
||||
|
||||
# Image Preprocessing
|
||||
|
||||
img_np = (image[0].cpu().numpy() * 255).astype(np.uint8)
|
||||
orig_img = Image.fromarray(img_np)
|
||||
orig_w, orig_h = orig_img.size
|
||||
|
||||
|
||||
scale = min(1024 / orig_w, 1024 / orig_h)
|
||||
proc_img = orig_img.resize((int(orig_w * scale), int(orig_h * scale)), Image.Resampling.LANCZOS) if scale < 1 else orig_img
|
||||
pw, ph = proc_img.size
|
||||
|
||||
|
||||
img_buf = io.BytesIO()
|
||||
proc_img.save(img_buf, format='PNG')
|
||||
|
||||
# Thinking Config
|
||||
t_config = None
|
||||
if "gemini-2.0" not in model.lower():
|
||||
budget = thinking_budget if thinking else 0
|
||||
@@ -97,8 +102,7 @@ class GeminiSegmentationVertexNode:
|
||||
t_config = types.ThinkingConfig(thinking_budget=budget)
|
||||
|
||||
prompt = f'Give the segmentation masks for {segment_prompt}. Output a JSON list of segmentation masks where each entry contains the 2D bounding box in the key "box_2d", the segmentation mask in key "mask", and the text label in the key "label".'
|
||||
|
||||
# API Call
|
||||
|
||||
response = client.models.generate_content(
|
||||
model=model,
|
||||
contents=[types.Content(role="user", parts=[
|
||||
@@ -116,28 +120,24 @@ class GeminiSegmentationVertexNode:
|
||||
raise RuntimeError(f"Gemini API Error: {e}")
|
||||
|
||||
final_mask = np.zeros((ph, pw), dtype=np.uint8)
|
||||
|
||||
|
||||
for seg in segments:
|
||||
try:
|
||||
# Calculate integer coords
|
||||
ymin, xmin, ymax, xmax = seg['box_2d']
|
||||
x1, y1 = int(xmin * pw / 1000), int(ymin * ph / 1000)
|
||||
x2, y2 = int(xmax * pw / 1000), int(ymax * ph / 1000)
|
||||
w, h = x2 - x1, y2 - y1
|
||||
|
||||
|
||||
if w <= 0 or h <= 0: continue
|
||||
|
||||
# Decode & Resize Patch
|
||||
mask_str = seg['mask'].split(",")[1] if "data:image" in seg['mask'] else seg['mask']
|
||||
patch = Image.open(io.BytesIO(base64.b64decode(mask_str))).convert('L')
|
||||
|
||||
|
||||
if patch.size != (w, h):
|
||||
patch = patch.resize((w, h), Image.Resampling.NEAREST)
|
||||
|
||||
patch_arr = np.array(patch)
|
||||
patch_arr = np.where(patch_arr > 128, 255, 0).astype(np.uint8)
|
||||
|
||||
# Safe slicing to handle potential boundary issues
|
||||
|
||||
patch_arr = np.where(np.array(patch) > 128, 255, 0).astype(np.uint8)
|
||||
|
||||
target_slice = final_mask[y1:y2, x1:x2]
|
||||
if target_slice.shape == patch_arr.shape:
|
||||
np.maximum(target_slice, patch_arr, out=target_slice)
|
||||
@@ -146,7 +146,7 @@ class GeminiSegmentationVertexNode:
|
||||
|
||||
if (pw, ph) != (orig_w, orig_h):
|
||||
final_mask = np.array(Image.fromarray(final_mask).resize((orig_w, orig_h), Image.Resampling.NEAREST))
|
||||
|
||||
|
||||
return (torch.from_numpy(final_mask.astype(np.float32) / 255.0).unsqueeze(0),)
|
||||
|
||||
NODE_CLASS_MAPPINGS = {"GeminiSegmentationVertexNode": GeminiSegmentationVertexNode}
|
||||
|
||||
+43
-51
@@ -1,12 +1,10 @@
|
||||
import os
|
||||
import io
|
||||
import json
|
||||
import tempfile
|
||||
import torch
|
||||
from google.genai import Client, types
|
||||
from google.oauth2 import service_account
|
||||
|
||||
class GeminiTTSVertexNode:
|
||||
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
@@ -14,21 +12,21 @@ class GeminiTTSVertexNode:
|
||||
"text": ("STRING", {"multiline": True, "default": ""}),
|
||||
"project_id": ("STRING", {"multiline": False, "default": ""}),
|
||||
"location": ([
|
||||
"global","us-central1", "us-east1", "us-east4", "us-east5", "us-south1",
|
||||
"us-west1", "us-west2", "us-west3", "us-west4",
|
||||
"northamerica-northeast1", "northamerica-northeast2",
|
||||
"southamerica-east1", "southamerica-west1", "africa-south1",
|
||||
"europe-west1", "europe-north1", "europe-west2", "europe-west3",
|
||||
"europe-west4", "europe-west6", "europe-west8", "europe-west9",
|
||||
"europe-west12", "europe-southwest1", "europe-central2",
|
||||
"asia-east1", "asia-east2", "asia-northeast1", "asia-northeast2",
|
||||
"asia-northeast3", "asia-south1", "asia-south2", "asia-southeast1",
|
||||
"asia-southeast2", "australia-southeast1", "australia-southeast2",
|
||||
"global", "us-central1", "us-east1", "us-east4", "us-east5", "us-south1",
|
||||
"us-west1", "us-west2", "us-west3", "us-west4",
|
||||
"northamerica-northeast1", "northamerica-northeast2",
|
||||
"southamerica-east1", "southamerica-west1", "africa-south1",
|
||||
"europe-west1", "europe-north1", "europe-west2", "europe-west3",
|
||||
"europe-west4", "europe-west6", "europe-west8", "europe-west9",
|
||||
"europe-west12", "europe-southwest1", "europe-central2",
|
||||
"asia-east1", "asia-east2", "asia-northeast1", "asia-northeast2",
|
||||
"asia-northeast3", "asia-south1", "asia-south2", "asia-southeast1",
|
||||
"asia-southeast2", "australia-southeast1", "australia-southeast2",
|
||||
"me-central1", "me-central2", "me-west1"
|
||||
], {"default": "us-central1"}),
|
||||
"service_account": ("STRING", {"multiline": True, "default": ""}),
|
||||
"model": (["gemini-2.5-flash-preview-tts", "gemini-2.5-pro-preview-tts"],),
|
||||
"voice_id": (["Zephyr", "Puck", "Charon", "Kore", "Fenrir", "Leda", "Orus", "Aoede", "Callirrhoe", "Autonoe", "Enceladus", "Iapetus", "Umbriel", "Algieba", "Despina", "Erinome", "Achernar", "Laomedeia", "Rasalgethi", "Algenib", "Achird", "Pulcherrima", "Gacrux", "Schedar", "Alnilam", "Sulafat", "Sadaltager", "Sadachbia", "Vindemiatrix", "Zubenelgenubi"],),
|
||||
"voice_id": (["Zephyr", "Puck", "Charon", "Kore", "Fenrir", "Leda", "Orus", "Aoede", "Callirrhoe", "Autonoe", "Enceladus", "Iapetus", "Umbriel", "Algieba", "Despina", "Erinome", "Achernar", "Laomedeia", "Rasalgethi", "Algenib", "Achird", "Pulcherrima", "Gacrux", "Schedar", "Alnilam", "Sulafat", "Sadaltager", "Sadachbia", "Vindemiatrix", "Zubenelgenubi"],),
|
||||
"seed": ("INT", {"default": 69, "min": -1, "max": 2147483646, "step": 1}),
|
||||
"temperature": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}),
|
||||
},
|
||||
@@ -36,64 +34,59 @@ class GeminiTTSVertexNode:
|
||||
"system_prompt": ("STRING", {"multiline": True, "default": ""}),
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
RETURN_TYPES = ("AUDIO",)
|
||||
RETURN_NAMES = ("audio",)
|
||||
FUNCTION = "generate_speech"
|
||||
CATEGORY = "audio/generation"
|
||||
|
||||
|
||||
def setup_client(self, service_account_json, project_id, location):
|
||||
if not service_account_json.strip():
|
||||
raise ValueError("Service account JSON content is required.")
|
||||
|
||||
if not project_id.strip():
|
||||
raise ValueError("Project ID is required.")
|
||||
|
||||
|
||||
try:
|
||||
json.loads(service_account_json)
|
||||
sa_info = json.loads(service_account_json)
|
||||
except json.JSONDecodeError as e:
|
||||
raise ValueError(f"Invalid JSON content: {str(e)}")
|
||||
|
||||
temp_file = tempfile.NamedTemporaryFile(mode='w', suffix='.json', delete=False)
|
||||
temp_file.write(service_account_json.strip())
|
||||
temp_file.close()
|
||||
|
||||
os.environ['GOOGLE_APPLICATION_CREDENTIALS'] = temp_file.name
|
||||
|
||||
|
||||
credentials = service_account.Credentials.from_service_account_info(
|
||||
sa_info,
|
||||
scopes=["https://www.googleapis.com/auth/cloud-platform"]
|
||||
)
|
||||
|
||||
return Client(
|
||||
vertexai=True,
|
||||
project=project_id.strip(),
|
||||
vertexai=True,
|
||||
project=project_id.strip(),
|
||||
location=location.strip(),
|
||||
credentials=credentials,
|
||||
http_options=types.HttpOptions(
|
||||
retry_options=types.HttpRetryOptions(attempts=10, jitter=10)
|
||||
)
|
||||
)
|
||||
|
||||
def generate_speech(self, text, project_id, location, service_account, voice_id,
|
||||
temperature, model, seed, system_prompt=""):
|
||||
|
||||
|
||||
def generate_speech(self, text, project_id, location, service_account, voice_id,
|
||||
temperature, model, seed, system_prompt=""):
|
||||
|
||||
if not text.strip():
|
||||
raise ValueError("Text input cannot be empty.")
|
||||
|
||||
client = self.setup_client(service_account, project_id, location)
|
||||
|
||||
final_prompt = text
|
||||
if system_prompt.strip():
|
||||
final_prompt = f"{system_prompt.strip()}\n\n{text}"
|
||||
|
||||
speech_config = types.SpeechConfig(
|
||||
voice_config=types.VoiceConfig(
|
||||
prebuilt_voice_config=types.PrebuiltVoiceConfig(voice_name=voice_id)
|
||||
)
|
||||
)
|
||||
|
||||
client = self.setup_client(service_account, project_id, location)
|
||||
|
||||
final_prompt = f"{system_prompt.strip()}\n\n{text}" if system_prompt.strip() else text
|
||||
|
||||
config = types.GenerateContentConfig(
|
||||
temperature=temperature,
|
||||
seed=seed,
|
||||
response_modalities=["AUDIO"],
|
||||
speech_config=speech_config,
|
||||
speech_config=types.SpeechConfig(
|
||||
voice_config=types.VoiceConfig(
|
||||
prebuilt_voice_config=types.PrebuiltVoiceConfig(voice_name=voice_id)
|
||||
)
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
try:
|
||||
response = client.models.generate_content(
|
||||
model=model,
|
||||
@@ -104,17 +97,16 @@ class GeminiTTSVertexNode:
|
||||
raise RuntimeError(f"Gemini API Error: {str(e)}")
|
||||
|
||||
try:
|
||||
inline_data = response.candidates[0].content.parts[0].inline_data
|
||||
audio_bytes = inline_data.data
|
||||
audio_bytes = response.candidates[0].content.parts[0].inline_data.data
|
||||
except (AttributeError, IndexError, TypeError):
|
||||
raise ValueError("API returned a response, but it contained no audio data.")
|
||||
|
||||
waveform = torch.frombuffer(bytearray(audio_bytes), dtype=torch.int16)
|
||||
waveform = waveform.to(torch.float32) / 32768.0
|
||||
waveform = waveform.unsqueeze(0).unsqueeze(0)
|
||||
|
||||
|
||||
return ({"waveform": waveform, "sample_rate": 24000},)
|
||||
|
||||
|
||||
@classmethod
|
||||
def IS_CHANGED(cls, **kwargs):
|
||||
return f"{kwargs.get('text', '')}-{kwargs.get('voice_id', '')}-{kwargs.get('temperature', 1.0)}-{kwargs.get('model', '')}-{kwargs.get('seed', 69)}-{kwargs.get('system_prompt', '')}"
|
||||
|
||||
+55
-50
@@ -1,16 +1,15 @@
|
||||
import os
|
||||
import json
|
||||
import tempfile
|
||||
import io
|
||||
import json
|
||||
import base64
|
||||
import torch
|
||||
import numpy as np
|
||||
from PIL import Image
|
||||
from google import genai
|
||||
from google.genai import types
|
||||
from google.oauth2 import service_account
|
||||
|
||||
class GoogleImagenEditVertex:
|
||||
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
@@ -32,61 +31,68 @@ class GoogleImagenEditVertex:
|
||||
"negative_prompt": ("STRING", {"multiline": True, "default": ""}),
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
RETURN_NAMES = ("edited_images",)
|
||||
FUNCTION = "edit_image"
|
||||
CATEGORY = "image/edit"
|
||||
|
||||
|
||||
def setup_client(self, service_account_json, project_id, location):
|
||||
if not service_account_json.strip():
|
||||
raise ValueError("Service account JSON content is required.")
|
||||
if not project_id.strip():
|
||||
raise ValueError("Project ID is required.")
|
||||
|
||||
|
||||
try:
|
||||
json.loads(service_account_json)
|
||||
sa_info = json.loads(service_account_json)
|
||||
except json.JSONDecodeError as e:
|
||||
raise ValueError(f"Invalid JSON content: {str(e)}")
|
||||
|
||||
temp_file = tempfile.NamedTemporaryFile(mode='w', suffix='.json', delete=False)
|
||||
temp_file.write(service_account_json.strip())
|
||||
temp_file.close()
|
||||
|
||||
os.environ['GOOGLE_APPLICATION_CREDENTIALS'] = temp_file.name
|
||||
|
||||
return genai.Client(vertexai=True, project=project_id.strip(), location=location.strip())
|
||||
|
||||
def edit_image(self, image, mask, prompt, project_id, location, service_account,
|
||||
edit_mode, number_of_images, seed, base_steps, guidance_scale, mask_dilation, negative_prompt=""):
|
||||
|
||||
|
||||
credentials = service_account.Credentials.from_service_account_info(
|
||||
sa_info,
|
||||
scopes=["https://www.googleapis.com/auth/cloud-platform"]
|
||||
)
|
||||
|
||||
return genai.Client(
|
||||
vertexai=True,
|
||||
project=project_id.strip(),
|
||||
location=location.strip(),
|
||||
credentials=credentials,
|
||||
http_options=types.HttpOptions(
|
||||
retry_options=types.HttpRetryOptions(attempts=10, jitter=10)
|
||||
)
|
||||
)
|
||||
|
||||
def edit_image(self, image, mask, prompt, project_id, location, service_account,
|
||||
edit_mode, number_of_images, seed, base_steps, guidance_scale, mask_dilation, negative_prompt=""):
|
||||
|
||||
client = self.setup_client(service_account, project_id, location)
|
||||
|
||||
|
||||
def to_b64(img):
|
||||
b = io.BytesIO()
|
||||
img.save(b, format='PNG')
|
||||
return base64.b64encode(b.getvalue()).decode('utf-8')
|
||||
|
||||
img_pil = Image.fromarray((image[0].cpu().numpy() * 255).astype(np.uint8))
|
||||
|
||||
mask_np = mask.cpu().numpy()
|
||||
if mask_np.ndim == 3: mask_np = mask_np[0]
|
||||
mask_pil = Image.fromarray((mask_np * 255).astype(np.uint8), mode='L')
|
||||
|
||||
config_dict = {
|
||||
"edit_mode": edit_mode,
|
||||
"number_of_images": number_of_images,
|
||||
"base_steps": base_steps,
|
||||
"seed": seed,
|
||||
"guidance_scale": guidance_scale,
|
||||
"output_mime_type": "image/jpeg",
|
||||
"include_rai_reason": True,
|
||||
}
|
||||
|
||||
if negative_prompt.strip():
|
||||
config_dict["negative_prompt"] = negative_prompt.strip()
|
||||
|
||||
try:
|
||||
def to_b64(img):
|
||||
b = io.BytesIO()
|
||||
img.save(b, format='PNG')
|
||||
return base64.b64encode(b.getvalue()).decode('utf-8')
|
||||
|
||||
img_pil = Image.fromarray((image[0].cpu().numpy() * 255).astype(np.uint8))
|
||||
|
||||
mask_np = mask.cpu().numpy()
|
||||
if mask_np.ndim == 3: mask_np = mask_np[0]
|
||||
mask_pil = Image.fromarray((mask_np * 255).astype(np.uint8), mode='L')
|
||||
|
||||
config_dict = {
|
||||
"edit_mode": edit_mode,
|
||||
"number_of_images": number_of_images,
|
||||
"base_steps": base_steps,
|
||||
"seed": seed,
|
||||
"guidance_scale": guidance_scale,
|
||||
"output_mime_type": "image/jpeg",
|
||||
"include_rai_reason": True,
|
||||
}
|
||||
|
||||
if negative_prompt.strip():
|
||||
config_dict["negative_prompt"] = negative_prompt.strip()
|
||||
|
||||
response = client.models.edit_image(
|
||||
model="imagen-3.0-capability-001",
|
||||
prompt=prompt,
|
||||
@@ -98,18 +104,17 @@ class GoogleImagenEditVertex:
|
||||
config=types.EditImageConfig(**config_dict)
|
||||
)
|
||||
|
||||
if not response.generated_images: raise ValueError("No images generated")
|
||||
if not response.generated_images:
|
||||
raise ValueError("No images generated")
|
||||
|
||||
output_tensors = []
|
||||
for item in response.generated_images:
|
||||
img_bytes = item.image.image_bytes
|
||||
res_img = Image.open(io.BytesIO(img_bytes)).convert("RGB")
|
||||
res_img = Image.open(io.BytesIO(item.image.image_bytes)).convert("RGB")
|
||||
output_tensors.append(torch.from_numpy(np.array(res_img).astype(np.float32) / 255.0))
|
||||
|
||||
|
||||
return (torch.stack(output_tensors),)
|
||||
|
||||
except Exception as e:
|
||||
print(f"Google Imagen Edit Error: {e}")
|
||||
raise RuntimeError(f"Google Imagen Edit Error: {e}")
|
||||
|
||||
@classmethod
|
||||
|
||||
+44
-41
@@ -1,63 +1,68 @@
|
||||
import os
|
||||
import json
|
||||
import tempfile
|
||||
import io
|
||||
import json
|
||||
import torch
|
||||
import numpy as np
|
||||
from PIL import Image
|
||||
from google import genai
|
||||
from google.genai import types
|
||||
from google.oauth2 import service_account
|
||||
|
||||
class GoogleImagenGenerateVertex:
|
||||
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"prompt": ("STRING", {"multiline": True, "default": "A majestic lion in the savanna"}),
|
||||
"project_id": ("STRING", {"multiline": False, "default": ""}),
|
||||
"location": (["global", "us-central1", "us-east1", "us-east4", "us-east5", "us-south1", "us-west1", "us-west2", "us-west3", "us-west4", "northamerica-northeast1", "northamerica-northeast2", "southamerica-east1", "southamerica-west1", "africa-south1", "europe-west1", "europe-north1", "europe-west2", "europe-west3", "europe-west4", "europe-west6", "europe-west8", "europe-west9", "europe-west12", "europe-southwest1", "europe-central2", "asia-east1", "asia-east2", "asia-northeast1", "asia-northeast2", "asia-northeast3", "asia-south1", "asia-south2", "asia-southeast1", "asia-southeast2", "australia-southeast1", "australia-southeast2", "me-central1", "me-central2", "me-west1"], {"default": "us-central1"}),
|
||||
"service_account": ("STRING", {"multiline": True, "default": ""}),
|
||||
"model": (["imagen-4.0-ultra-generate-001", "imagen-4.0-generate-001", "imagen-4.0-fast-generate-001", "imagen-3.0-generate-002"], {"default": "imagen-4.0-generate-001"}),
|
||||
"number_of_images": ("INT", {"default": 1, "min": 1, "max": 4, "step": 1}),
|
||||
"aspect_ratio": (["1:1", "9:16", "16:9", "4:3", "3:4"], {"default": "1:1"}),
|
||||
"image_size": (["1K", "2K"], {"default": "1K"}),
|
||||
"seed": ("INT", {"default": 69, "min": 1, "max": 2147483646, "step": 1}),
|
||||
"guidance_scale": ("FLOAT", {"default": 7.5, "min": 1.0, "max": 20.0, "step": 0.1}),
|
||||
"prompt": ("STRING", {"multiline": True, "default": "A majestic lion in the savanna"}),
|
||||
"project_id": ("STRING", {"multiline": False, "default": ""}),
|
||||
"location": (["global", "us-central1", "us-east1", "us-east4", "us-east5", "us-south1", "us-west1", "us-west2", "us-west3", "us-west4", "northamerica-northeast1", "northamerica-northeast2", "southamerica-east1", "southamerica-west1", "africa-south1", "europe-west1", "europe-north1", "europe-west2", "europe-west3", "europe-west4", "europe-west6", "europe-west8", "europe-west9", "europe-west12", "europe-southwest1", "europe-central2", "asia-east1", "asia-east2", "asia-northeast1", "asia-northeast2", "asia-northeast3", "asia-south1", "asia-south2", "asia-southeast1", "asia-southeast2", "australia-southeast1", "australia-southeast2", "me-central1", "me-central2", "me-west1"], {"default": "us-central1"}),
|
||||
"service_account": ("STRING", {"multiline": True, "default": ""}),
|
||||
"model": (["imagen-4.0-ultra-generate-001", "imagen-4.0-generate-001", "imagen-4.0-fast-generate-001", "imagen-3.0-generate-002"], {"default": "imagen-4.0-generate-001"}),
|
||||
"number_of_images": ("INT", {"default": 1, "min": 1, "max": 4, "step": 1}),
|
||||
"aspect_ratio": (["1:1", "9:16", "16:9", "4:3", "3:4"], {"default": "1:1"}),
|
||||
"image_size": (["1K", "2K"], {"default": "1K"}),
|
||||
"seed": ("INT", {"default": 69, "min": 1, "max": 2147483646, "step": 1}),
|
||||
"guidance_scale": ("FLOAT", {"default": 7.5, "min": 1.0, "max": 20.0, "step": 0.1}),
|
||||
},
|
||||
"optional": {
|
||||
"negative_prompt": ("STRING", {"multiline": True, "default": ""}),
|
||||
"negative_prompt": ("STRING", {"multiline": True, "default": ""}),
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
RETURN_NAMES = ("images",)
|
||||
FUNCTION = "generate_images"
|
||||
CATEGORY = "image/generation"
|
||||
|
||||
|
||||
def setup_client(self, service_account_json, project_id, location):
|
||||
if not service_account_json.strip():
|
||||
raise ValueError("Service account JSON content is required.")
|
||||
|
||||
if not project_id.strip():
|
||||
raise ValueError("Project ID is required.")
|
||||
|
||||
|
||||
try:
|
||||
json.loads(service_account_json)
|
||||
sa_info = json.loads(service_account_json)
|
||||
except json.JSONDecodeError as e:
|
||||
raise ValueError(f"Invalid JSON content: {str(e)}")
|
||||
|
||||
temp_file = tempfile.NamedTemporaryFile(mode='w', suffix='.json', delete=False)
|
||||
temp_file.write(service_account_json.strip())
|
||||
temp_file.close()
|
||||
|
||||
os.environ['GOOGLE_APPLICATION_CREDENTIALS'] = temp_file.name
|
||||
|
||||
return genai.Client(vertexai=True, project=project_id.strip(), location=location.strip())
|
||||
|
||||
|
||||
credentials = service_account.Credentials.from_service_account_info(
|
||||
sa_info,
|
||||
scopes=["https://www.googleapis.com/auth/cloud-platform"]
|
||||
)
|
||||
|
||||
return genai.Client(
|
||||
vertexai=True,
|
||||
project=project_id.strip(),
|
||||
location=location.strip(),
|
||||
credentials=credentials,
|
||||
http_options=types.HttpOptions(
|
||||
retry_options=types.HttpRetryOptions(attempts=10, jitter=10)
|
||||
)
|
||||
)
|
||||
|
||||
def generate_images(self, prompt, project_id, location, service_account, model, number_of_images, aspect_ratio, image_size, seed, guidance_scale, negative_prompt=""):
|
||||
client = self.setup_client(service_account, project_id, location)
|
||||
|
||||
|
||||
config = types.GenerateImagesConfig(
|
||||
number_of_images=number_of_images,
|
||||
aspect_ratio=aspect_ratio,
|
||||
@@ -65,34 +70,32 @@ class GoogleImagenGenerateVertex:
|
||||
seed=seed,
|
||||
negative_prompt=negative_prompt.strip() if negative_prompt.strip() else None
|
||||
)
|
||||
|
||||
|
||||
if "imagen-4.0" in model and "fast" not in model:
|
||||
config.image_size = image_size
|
||||
|
||||
try:
|
||||
result = client.models.generate_images(model=model, prompt=prompt, config=config)
|
||||
if not result.generated_images: raise ValueError("No images generated")
|
||||
|
||||
if not result.generated_images:
|
||||
raise ValueError("No images generated")
|
||||
|
||||
tensors = []
|
||||
for item in result.generated_images:
|
||||
img_data = item.image
|
||||
|
||||
if hasattr(img_data, "image_bytes"):
|
||||
pil_img = Image.open(io.BytesIO(img_data.image_bytes))
|
||||
elif hasattr(img_data, "convert"):
|
||||
pil_img = img_data
|
||||
else:
|
||||
pil_img = Image.open(io.BytesIO(img_data))
|
||||
|
||||
pil_img = pil_img.convert("RGB")
|
||||
tensors.append(torch.from_numpy(np.array(pil_img).astype(np.float32) / 255.0))
|
||||
|
||||
|
||||
tensors.append(torch.from_numpy(np.array(pil_img.convert("RGB")).astype(np.float32) / 255.0))
|
||||
|
||||
return (torch.stack(tensors),)
|
||||
|
||||
|
||||
except Exception as e:
|
||||
print(f"Google Imagen Error: {e}")
|
||||
raise RuntimeError(f"Google Imagen Error: {e}")
|
||||
|
||||
|
||||
@classmethod
|
||||
def IS_CHANGED(cls, **kwargs):
|
||||
return float("nan")
|
||||
|
||||
+57
-45
@@ -1,39 +1,39 @@
|
||||
import os
|
||||
import io
|
||||
import json
|
||||
import tempfile
|
||||
import torch
|
||||
import numpy as np
|
||||
from PIL import Image
|
||||
from google import genai
|
||||
from google.genai import types
|
||||
from google.oauth2 import service_account
|
||||
|
||||
class NanoBananaVertexNode:
|
||||
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"project_id": ("STRING", {"multiline": False, "default": ""}),
|
||||
"location": ([
|
||||
"global", "us-central1", "us-east1", "us-east4", "us-east5", "us-south1",
|
||||
"us-west1", "us-west2", "us-west3", "us-west4",
|
||||
"northamerica-northeast1", "northamerica-northeast2",
|
||||
"southamerica-east1", "southamerica-west1", "africa-south1",
|
||||
"europe-west1", "europe-north1", "europe-west2", "europe-west3",
|
||||
"europe-west4", "europe-west6", "europe-west8", "europe-west9",
|
||||
"europe-west12", "europe-southwest1", "europe-central2",
|
||||
"asia-east1", "asia-east2", "asia-northeast1", "asia-northeast2",
|
||||
"asia-northeast3", "asia-south1", "asia-south2", "asia-southeast1",
|
||||
"asia-southeast2", "australia-southeast1", "australia-southeast2",
|
||||
"global", "us-central1", "us-east1", "us-east4", "us-east5", "us-south1",
|
||||
"us-west1", "us-west2", "us-west3", "us-west4",
|
||||
"northamerica-northeast1", "northamerica-northeast2",
|
||||
"southamerica-east1", "southamerica-west1", "africa-south1",
|
||||
"europe-west1", "europe-north1", "europe-west2", "europe-west3",
|
||||
"europe-west4", "europe-west6", "europe-west8", "europe-west9",
|
||||
"europe-west12", "europe-southwest1", "europe-central2",
|
||||
"asia-east1", "asia-east2", "asia-northeast1", "asia-northeast2",
|
||||
"asia-northeast3", "asia-south1", "asia-south2", "asia-southeast1",
|
||||
"asia-southeast2", "australia-southeast1", "australia-southeast2",
|
||||
"me-central1", "me-central2", "me-west1"
|
||||
], {"default": "us-central1"}),
|
||||
"service_account": ("STRING", {"multiline": True, "default": ""}),
|
||||
"model": (["gemini-3-pro-image-preview", "gemini-2.5-flash-image", "gemini-3.1-flash-image-preview"],),
|
||||
"aspect_ratio": (["1:1", "2:3", "3:2", "3:4", "4:3", "9:16", "16:9", "21:9"],),
|
||||
"resolution": (["1K", "2K", "4K"], {"default": "1K"}),
|
||||
"temperature": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01}),
|
||||
"top_p": ("FLOAT", {"default": 0.85, "min": 0.0, "max": 1.0, "step": 0.01}),
|
||||
"temperature": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}),
|
||||
"top_p": ("FLOAT", {"default": 0.95, "min": 0.0, "max": 1.0, "step": 0.01}),
|
||||
"google_search": ("BOOLEAN", {"default": False}),
|
||||
"seed": ("INT", {"default": 69, "min": -1, "max": 2147483646, "step": 1}),
|
||||
},
|
||||
"optional": {
|
||||
@@ -46,71 +46,85 @@ class NanoBananaVertexNode:
|
||||
"image_5": ("IMAGE",),
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
RETURN_NAMES = ("image",)
|
||||
FUNCTION = "generate"
|
||||
CATEGORY = "image/generation"
|
||||
|
||||
|
||||
def setup_client(self, service_account_json, project_id, location):
|
||||
if not service_account_json.strip():
|
||||
raise ValueError("Service account JSON content is required.")
|
||||
if not project_id.strip():
|
||||
raise ValueError("Project ID is required.")
|
||||
|
||||
|
||||
try:
|
||||
json.loads(service_account_json)
|
||||
sa_info = json.loads(service_account_json)
|
||||
except json.JSONDecodeError as e:
|
||||
raise ValueError(f"Invalid JSON content: {str(e)}")
|
||||
|
||||
temp_file = tempfile.NamedTemporaryFile(mode='w', suffix='.json', delete=False)
|
||||
temp_file.write(service_account_json.strip())
|
||||
temp_file.close()
|
||||
|
||||
os.environ['GOOGLE_APPLICATION_CREDENTIALS'] = temp_file.name
|
||||
|
||||
return genai.Client(vertexai=True, project=project_id.strip(), location=location.strip())
|
||||
|
||||
credentials = service_account.Credentials.from_service_account_info(
|
||||
sa_info,
|
||||
scopes=["https://www.googleapis.com/auth/cloud-platform"]
|
||||
)
|
||||
|
||||
return genai.Client(
|
||||
vertexai=True,
|
||||
project=project_id.strip(),
|
||||
location=location.strip(),
|
||||
credentials=credentials,
|
||||
http_options=types.HttpOptions(
|
||||
retry_options=types.HttpRetryOptions(attempts=10, jitter=10)
|
||||
)
|
||||
)
|
||||
|
||||
def _convert_tensor_to_bytes(self, tensor):
|
||||
if tensor.dim() == 4:
|
||||
tensor = tensor[0]
|
||||
|
||||
arr = (tensor.cpu().numpy() * 255).astype(np.uint8)
|
||||
buf = io.BytesIO()
|
||||
Image.fromarray(arr).save(buf, format='PNG')
|
||||
return buf.getvalue()
|
||||
|
||||
def generate(self, project_id, location, service_account, model, aspect_ratio, resolution, temperature, top_p, seed,
|
||||
def generate(self, project_id, location, service_account, model, aspect_ratio, resolution,
|
||||
temperature, top_p, google_search, seed,
|
||||
prompt="", system_instruction="", **kwargs):
|
||||
|
||||
|
||||
client = self.setup_client(service_account, project_id, location)
|
||||
|
||||
|
||||
parts = []
|
||||
input_images = [kwargs.get(f"image_{i}") for i in range(1, 6)]
|
||||
for img in input_images:
|
||||
for i in range(1, 6):
|
||||
img = kwargs.get(f"image_{i}")
|
||||
if img is not None:
|
||||
img_bytes = self._convert_tensor_to_bytes(img)
|
||||
parts.append(types.Part.from_bytes(mime_type="image/png", data=img_bytes))
|
||||
|
||||
parts.append(types.Part.from_bytes(mime_type="image/png", data=self._convert_tensor_to_bytes(img)))
|
||||
|
||||
if prompt.strip():
|
||||
parts.append(types.Part.from_text(text=prompt))
|
||||
|
||||
|
||||
if not parts:
|
||||
raise ValueError("At least one image or prompt must be provided.")
|
||||
|
||||
tools = None
|
||||
if google_search:
|
||||
if "gemini-2.5" in model:
|
||||
print(f"Ignoring google_search: {model} does not support it.")
|
||||
else:
|
||||
tools = [types.Tool(googleSearch=types.GoogleSearch())]
|
||||
|
||||
img_config_params = {"aspect_ratio": aspect_ratio}
|
||||
if "gemini-3-pro" in model:
|
||||
img_config_params["image_size"] = resolution
|
||||
|
||||
|
||||
config = types.GenerateContentConfig(
|
||||
temperature=temperature,
|
||||
seed=seed,
|
||||
top_p=top_p,
|
||||
response_modalities=["IMAGE"],
|
||||
image_config=types.ImageConfig(**img_config_params),
|
||||
system_instruction=system_instruction.strip() if system_instruction.strip() else None
|
||||
system_instruction=system_instruction.strip() if system_instruction.strip() else None,
|
||||
tools=tools
|
||||
)
|
||||
|
||||
|
||||
try:
|
||||
response = client.models.generate_content(
|
||||
model=model,
|
||||
@@ -119,14 +133,11 @@ class NanoBananaVertexNode:
|
||||
)
|
||||
except Exception as e:
|
||||
raise RuntimeError(f"Gemini API Error: {str(e)}")
|
||||
|
||||
|
||||
try:
|
||||
img_data = response.candidates[0].content.parts[0].inline_data.data
|
||||
result_pil = Image.open(io.BytesIO(img_data)).convert("RGB")
|
||||
|
||||
result_tensor = torch.from_numpy(np.array(result_pil).astype(np.float32) / 255.0).unsqueeze(0)
|
||||
return (result_tensor,)
|
||||
|
||||
return (torch.from_numpy(np.array(result_pil).astype(np.float32) / 255.0).unsqueeze(0),)
|
||||
except (AttributeError, IndexError, TypeError):
|
||||
raise ValueError("API returned a response, but no valid image data was found.")
|
||||
|
||||
@@ -134,5 +145,6 @@ class NanoBananaVertexNode:
|
||||
def IS_CHANGED(cls, seed, **kwargs):
|
||||
return seed
|
||||
|
||||
|
||||
NODE_CLASS_MAPPINGS = {"NanoBananaVertexNode": NanoBananaVertexNode}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {"NanoBananaVertexNode": "Nano Banana (Vertex AI)"}
|
||||
+1
-1
@@ -1,7 +1,7 @@
|
||||
[project]
|
||||
name = "vertexapi"
|
||||
description = "A collection of powerful custom nodes for ComfyUI that connect your local workflows to closed-source AI models via Vertex AI."
|
||||
version = "1.0.3"
|
||||
version = "1.1.0"
|
||||
license = {file = "LICENSE"}
|
||||
# classifiers = [
|
||||
# # For OS-independent nodes (works on all operating systems)
|
||||
|
||||
+90
-94
@@ -1,16 +1,15 @@
|
||||
import time
|
||||
import os
|
||||
import io
|
||||
import json
|
||||
import tempfile
|
||||
import torch
|
||||
import numpy as np
|
||||
import av
|
||||
from PIL import Image
|
||||
from google import genai
|
||||
from google.genai import types
|
||||
from google.oauth2 import service_account
|
||||
|
||||
class GoogleVeoVertexVideoGenerator:
|
||||
class GoogleVeoVertexVideoGenerator:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
@@ -18,16 +17,16 @@ class GoogleVeoVertexVideoGenerator:
|
||||
"prompt": ("STRING", {"multiline": True, "default": "a cat reading a book"}),
|
||||
"project_id": ("STRING", {"multiline": False, "default": ""}),
|
||||
"location": ([
|
||||
"global", "us-central1", "us-east1", "us-east4", "us-east5", "us-south1",
|
||||
"us-west1", "us-west2", "us-west3", "us-west4",
|
||||
"northamerica-northeast1", "northamerica-northeast2",
|
||||
"southamerica-east1", "southamerica-west1", "africa-south1",
|
||||
"europe-west1", "europe-north1", "europe-west2", "europe-west3",
|
||||
"europe-west4", "europe-west6", "europe-west8", "europe-west9",
|
||||
"europe-west12", "europe-southwest1", "europe-central2",
|
||||
"asia-east1", "asia-east2", "asia-northeast1", "asia-northeast2",
|
||||
"asia-northeast3", "asia-south1", "asia-south2", "asia-southeast1",
|
||||
"asia-southeast2", "australia-southeast1", "australia-southeast2",
|
||||
"global", "us-central1", "us-east1", "us-east4", "us-east5", "us-south1",
|
||||
"us-west1", "us-west2", "us-west3", "us-west4",
|
||||
"northamerica-northeast1", "northamerica-northeast2",
|
||||
"southamerica-east1", "southamerica-west1", "africa-south1",
|
||||
"europe-west1", "europe-north1", "europe-west2", "europe-west3",
|
||||
"europe-west4", "europe-west6", "europe-west8", "europe-west9",
|
||||
"europe-west12", "europe-southwest1", "europe-central2",
|
||||
"asia-east1", "asia-east2", "asia-northeast1", "asia-northeast2",
|
||||
"asia-northeast3", "asia-south1", "asia-south2", "asia-southeast1",
|
||||
"asia-southeast2", "australia-southeast1", "australia-southeast2",
|
||||
"me-central1", "me-central2", "me-west1"
|
||||
], {"default": "us-central1"}),
|
||||
"service_account": ("STRING", {"multiline": True, "default": ""}),
|
||||
@@ -38,7 +37,8 @@ class GoogleVeoVertexVideoGenerator:
|
||||
"veo-3.0-generate-001",
|
||||
"veo-3.0-fast-generate-001",
|
||||
"veo-3.1-generate-001",
|
||||
"veo-3.1-fast-generate-001"
|
||||
"veo-3.1-fast-generate-001",
|
||||
"veo-3.1-lite-generate-preview"
|
||||
], {"default": "veo-3.0-generate-001"}),
|
||||
"resolution": (["720p", "1080p"], {"default": "720p"}),
|
||||
"aspect_ratio": (["16:9", "9:16"], {"default": "16:9"}),
|
||||
@@ -53,109 +53,105 @@ class GoogleVeoVertexVideoGenerator:
|
||||
"last_frame": ("IMAGE",),
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
RETURN_TYPES = ("IMAGE", "AUDIO")
|
||||
RETURN_NAMES = ("frames", "audio")
|
||||
FUNCTION = "generate_video"
|
||||
CATEGORY = "video/generation"
|
||||
OUTPUT_IS_LIST = (True, False)
|
||||
|
||||
def generate_video(self, prompt, project_id, location, service_account, model, resolution, aspect_ratio,
|
||||
duration_seconds, seed, generate_audio, fps, negative_prompt=None,
|
||||
first_frame=None, last_frame=None):
|
||||
|
||||
# Validate Service Account JSON
|
||||
if not service_account.strip():
|
||||
|
||||
def setup_client(self, service_account_json, project_id, location):
|
||||
if not service_account_json.strip():
|
||||
raise ValueError("Service account JSON content is required.")
|
||||
if not project_id.strip():
|
||||
raise ValueError("Project ID is required.")
|
||||
|
||||
|
||||
try:
|
||||
json.loads(service_account)
|
||||
sa_info = json.loads(service_account_json)
|
||||
except json.JSONDecodeError as e:
|
||||
raise ValueError(f"Invalid JSON content: {str(e)}")
|
||||
|
||||
creds_file = tempfile.NamedTemporaryFile(mode='w', suffix='.json', delete=False)
|
||||
creds_file.write(service_account.strip())
|
||||
creds_file.close()
|
||||
os.environ['GOOGLE_APPLICATION_CREDENTIALS'] = creds_file.name
|
||||
credentials = service_account.Credentials.from_service_account_info(
|
||||
sa_info,
|
||||
scopes=["https://www.googleapis.com/auth/cloud-platform"]
|
||||
)
|
||||
|
||||
try:
|
||||
client = genai.Client(vertexai=True, project=project_id, location=location)
|
||||
|
||||
config = types.GenerateVideosConfig(
|
||||
resolution=resolution,
|
||||
aspect_ratio=aspect_ratio,
|
||||
duration_seconds=duration_seconds,
|
||||
generate_audio=generate_audio,
|
||||
fps=int(fps),
|
||||
seed=seed if seed != -1 else None,
|
||||
negative_prompt=negative_prompt.strip() if negative_prompt and negative_prompt.strip() else None
|
||||
return genai.Client(
|
||||
vertexai=True,
|
||||
project=project_id.strip(),
|
||||
location=location.strip(),
|
||||
credentials=credentials,
|
||||
http_options=types.HttpOptions(
|
||||
retry_options=types.HttpRetryOptions(attempts=10, jitter=10)
|
||||
)
|
||||
)
|
||||
|
||||
# Helper: Tensor -> Bytes
|
||||
def tensor_to_bytes(t):
|
||||
# Handle batch dimension if present
|
||||
if t.dim() == 4:
|
||||
t = t[0]
|
||||
arr = (t.cpu().numpy() * 255).astype(np.uint8)
|
||||
b = io.BytesIO()
|
||||
Image.fromarray(arr).save(b, format="PNG")
|
||||
return b.getvalue()
|
||||
def generate_video(self, prompt, project_id, location, service_account, model, resolution, aspect_ratio,
|
||||
duration_seconds, seed, generate_audio, fps, negative_prompt=None,
|
||||
first_frame=None, last_frame=None):
|
||||
|
||||
gen_kwargs = {"model": model, "prompt": prompt, "config": config}
|
||||
|
||||
if first_frame is not None:
|
||||
gen_kwargs["image"] = types.Image(image_bytes=tensor_to_bytes(first_frame), mime_type="image/png")
|
||||
if last_frame is not None:
|
||||
setattr(config, 'last_frame', types.Image(image_bytes=tensor_to_bytes(last_frame), mime_type="image/png"))
|
||||
client = self.setup_client(service_account, project_id, location)
|
||||
|
||||
op = client.models.generate_videos(**gen_kwargs)
|
||||
print(f"Veo Operation: {op.name}")
|
||||
|
||||
while not op.done:
|
||||
time.sleep(5)
|
||||
op = client.operations.get(op)
|
||||
|
||||
if op.error: raise Exception(f"Veo Error: {op.error}")
|
||||
if not op.result.generated_videos: raise Exception("No videos generated")
|
||||
def tensor_to_bytes(t):
|
||||
if t.dim() == 4: t = t[0]
|
||||
arr = (t.cpu().numpy() * 255).astype(np.uint8)
|
||||
b = io.BytesIO()
|
||||
Image.fromarray(arr).save(b, format="PNG")
|
||||
return b.getvalue()
|
||||
|
||||
video_bytes = io.BytesIO(op.result.generated_videos[0].video.video_bytes)
|
||||
|
||||
# Decode Video
|
||||
config = types.GenerateVideosConfig(
|
||||
resolution=resolution,
|
||||
aspect_ratio=aspect_ratio,
|
||||
duration_seconds=duration_seconds,
|
||||
generate_audio=generate_audio,
|
||||
fps=int(fps),
|
||||
seed=seed if seed != -1 else None,
|
||||
negative_prompt=negative_prompt.strip() if negative_prompt and negative_prompt.strip() else None
|
||||
)
|
||||
|
||||
gen_kwargs = {"model": model, "prompt": prompt, "config": config}
|
||||
|
||||
if first_frame is not None:
|
||||
gen_kwargs["image"] = types.Image(image_bytes=tensor_to_bytes(first_frame), mime_type="image/png")
|
||||
if last_frame is not None:
|
||||
setattr(config, 'last_frame', types.Image(image_bytes=tensor_to_bytes(last_frame), mime_type="image/png"))
|
||||
|
||||
op = client.models.generate_videos(**gen_kwargs)
|
||||
print(f"Veo Operation: {op.name}")
|
||||
|
||||
while not op.done:
|
||||
time.sleep(5)
|
||||
op = client.operations.get(op)
|
||||
|
||||
if op.error: raise Exception(f"Veo Error: {op.error}")
|
||||
if not op.result.generated_videos: raise Exception("No videos generated")
|
||||
|
||||
video_bytes = io.BytesIO(op.result.generated_videos[0].video.video_bytes)
|
||||
|
||||
container = av.open(video_bytes)
|
||||
frames = []
|
||||
for frame in container.decode(video=0):
|
||||
img = frame.to_rgb().to_ndarray().astype(np.float32) / 255.0
|
||||
frames.append(torch.from_numpy(img).unsqueeze(0))
|
||||
container.close()
|
||||
|
||||
audio = None
|
||||
if generate_audio:
|
||||
video_bytes.seek(0)
|
||||
container = av.open(video_bytes)
|
||||
frames = []
|
||||
for frame in container.decode(video=0):
|
||||
img = frame.to_rgb().to_ndarray().astype(np.float32) / 255.0
|
||||
frames.append(torch.from_numpy(img).unsqueeze(0))
|
||||
if container.streams.audio:
|
||||
audio_data = [f.to_ndarray() for f in container.decode(audio=0)]
|
||||
if audio_data:
|
||||
waveform = torch.from_numpy(np.concatenate(audio_data, axis=1)).float()
|
||||
if audio_data[0].dtype == np.int16: waveform /= 32768.0
|
||||
elif audio_data[0].dtype == np.int32: waveform /= 2147483648.0
|
||||
audio = {"waveform": waveform.unsqueeze(0), "sample_rate": container.streams.audio[0].rate}
|
||||
container.close()
|
||||
|
||||
# Decode Audio
|
||||
audio = None
|
||||
if generate_audio:
|
||||
video_bytes.seek(0)
|
||||
container = av.open(video_bytes)
|
||||
if container.streams.audio:
|
||||
audio_data = [f.to_ndarray() for f in container.decode(audio=0)]
|
||||
if audio_data:
|
||||
waveform = torch.from_numpy(np.concatenate(audio_data, axis=1)).float()
|
||||
# Normalize 16/32-bit audio
|
||||
if audio_data[0].dtype == np.int16: waveform /= 32768.0
|
||||
elif audio_data[0].dtype == np.int32: waveform /= 2147483648.0
|
||||
|
||||
audio = {
|
||||
"waveform": waveform.unsqueeze(0),
|
||||
"sample_rate": container.streams.audio[0].rate
|
||||
}
|
||||
container.close()
|
||||
if not frames: raise Exception("Failed to decode video frames")
|
||||
|
||||
if not frames: raise Exception("Failed to decode video frames")
|
||||
|
||||
return ([torch.cat(frames, dim=0)], audio)
|
||||
return ([torch.cat(frames, dim=0)], audio)
|
||||
|
||||
finally:
|
||||
if os.path.exists(creds_file.name):
|
||||
os.remove(creds_file.name)
|
||||
|
||||
NODE_CLASS_MAPPINGS = {"GoogleVeoVertexVideoGenerator": GoogleVeoVertexVideoGenerator}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {"GoogleVeoVertexVideoGenerator": "Google Veo (Vertex AI)"}
|
||||
Reference in New Issue
Block a user