mirror of
https://git.datalinker.icu/comfyanonymous/ComfyUI
synced 2026-08-09 11:23:30 +08:00
converted Google Veo nodes
This commit is contained in:
parent
6dadfa2cb4
commit
b7916fdf9b
111
comfy_api_nodes/apis/veo_api.py
Normal file
111
comfy_api_nodes/apis/veo_api.py
Normal file
@ -0,0 +1,111 @@
|
|||||||
|
from typing import Optional, Union
|
||||||
|
from enum import Enum
|
||||||
|
|
||||||
|
from pydantic import BaseModel, Field
|
||||||
|
|
||||||
|
|
||||||
|
class Image2(BaseModel):
|
||||||
|
bytesBase64Encoded: str
|
||||||
|
gcsUri: Optional[str] = None
|
||||||
|
mimeType: Optional[str] = None
|
||||||
|
|
||||||
|
|
||||||
|
class Image3(BaseModel):
|
||||||
|
bytesBase64Encoded: Optional[str] = None
|
||||||
|
gcsUri: str
|
||||||
|
mimeType: Optional[str] = None
|
||||||
|
|
||||||
|
|
||||||
|
class Instance1(BaseModel):
|
||||||
|
image: Optional[Union[Image2, Image3]] = Field(
|
||||||
|
None, description='Optional image to guide video generation'
|
||||||
|
)
|
||||||
|
prompt: str = Field(..., description='Text description of the video')
|
||||||
|
|
||||||
|
|
||||||
|
class PersonGeneration1(str, Enum):
|
||||||
|
ALLOW = 'ALLOW'
|
||||||
|
BLOCK = 'BLOCK'
|
||||||
|
|
||||||
|
|
||||||
|
class Parameters1(BaseModel):
|
||||||
|
aspectRatio: Optional[str] = Field(None, examples=['16:9'])
|
||||||
|
durationSeconds: Optional[int] = None
|
||||||
|
enhancePrompt: Optional[bool] = None
|
||||||
|
generateAudio: Optional[bool] = Field(
|
||||||
|
None,
|
||||||
|
description='Generate audio for the video. Only supported by veo 3 models.',
|
||||||
|
)
|
||||||
|
negativePrompt: Optional[str] = None
|
||||||
|
personGeneration: Optional[PersonGeneration1] = None
|
||||||
|
sampleCount: Optional[int] = None
|
||||||
|
seed: Optional[int] = None
|
||||||
|
storageUri: Optional[str] = Field(
|
||||||
|
None, description='Optional Cloud Storage URI to upload the video'
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class VeoGenVidRequest(BaseModel):
|
||||||
|
instances: Optional[list[Instance1]] = None
|
||||||
|
parameters: Optional[Parameters1] = None
|
||||||
|
|
||||||
|
|
||||||
|
class VeoGenVidResponse(BaseModel):
|
||||||
|
name: str = Field(
|
||||||
|
...,
|
||||||
|
description='Operation resource name',
|
||||||
|
examples=[
|
||||||
|
'projects/PROJECT_ID/locations/us-central1/publishers/google/models/MODEL_ID/operations/a1b07c8e-7b5a-4aba-bb34-3e1ccb8afcc8'
|
||||||
|
],
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class VeoGenVidPollRequest(BaseModel):
|
||||||
|
operationName: str = Field(
|
||||||
|
...,
|
||||||
|
description='Full operation name (from predict response)',
|
||||||
|
examples=[
|
||||||
|
'projects/PROJECT_ID/locations/us-central1/publishers/google/models/MODEL_ID/operations/OPERATION_ID'
|
||||||
|
],
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class Video(BaseModel):
|
||||||
|
bytesBase64Encoded: Optional[str] = Field(
|
||||||
|
None, description='Base64-encoded video content'
|
||||||
|
)
|
||||||
|
gcsUri: Optional[str] = Field(None, description='Cloud Storage URI of the video')
|
||||||
|
mimeType: Optional[str] = Field(None, description='Video MIME type')
|
||||||
|
|
||||||
|
|
||||||
|
class Error1(BaseModel):
|
||||||
|
code: Optional[int] = Field(None, description='Error code')
|
||||||
|
message: Optional[str] = Field(None, description='Error message')
|
||||||
|
|
||||||
|
|
||||||
|
class Response1(BaseModel):
|
||||||
|
field_type: Optional[str] = Field(
|
||||||
|
None,
|
||||||
|
alias='@type',
|
||||||
|
examples=[
|
||||||
|
'type.googleapis.com/cloud.ai.large_models.vision.GenerateVideoResponse'
|
||||||
|
],
|
||||||
|
)
|
||||||
|
raiMediaFilteredCount: Optional[int] = Field(
|
||||||
|
None, description='Count of media filtered by responsible AI policies'
|
||||||
|
)
|
||||||
|
raiMediaFilteredReasons: Optional[list[str]] = Field(
|
||||||
|
None, description='Reasons why media was filtered by responsible AI policies'
|
||||||
|
)
|
||||||
|
videos: Optional[list[Video]] = None
|
||||||
|
|
||||||
|
|
||||||
|
class VeoGenVidPollResponse(BaseModel):
|
||||||
|
done: Optional[bool] = None
|
||||||
|
error: Optional[Error1] = Field(
|
||||||
|
None, description='Error details if operation failed'
|
||||||
|
)
|
||||||
|
name: Optional[str] = None
|
||||||
|
response: Optional[Response1] = Field(
|
||||||
|
None, description='The actual prediction response if done is true'
|
||||||
|
)
|
||||||
@ -1,28 +1,24 @@
|
|||||||
import logging
|
|
||||||
import base64
|
import base64
|
||||||
import aiohttp
|
|
||||||
import torch
|
|
||||||
from io import BytesIO
|
from io import BytesIO
|
||||||
from typing import Optional
|
|
||||||
from typing_extensions import override
|
from typing_extensions import override
|
||||||
|
|
||||||
from comfy_api.latest import ComfyExtension, IO
|
|
||||||
from comfy_api.input_impl.video_types import VideoFromFile
|
from comfy_api.input_impl.video_types import VideoFromFile
|
||||||
from comfy_api_nodes.apis import (
|
from comfy_api.latest import IO, ComfyExtension
|
||||||
VeoGenVidRequest,
|
from comfy_api_nodes.apis.veo_api import (
|
||||||
VeoGenVidResponse,
|
|
||||||
VeoGenVidPollRequest,
|
VeoGenVidPollRequest,
|
||||||
VeoGenVidPollResponse,
|
VeoGenVidPollResponse,
|
||||||
|
VeoGenVidRequest,
|
||||||
|
VeoGenVidResponse,
|
||||||
)
|
)
|
||||||
from comfy_api_nodes.apis.client import (
|
from comfy_api_nodes.util import (
|
||||||
ApiEndpoint,
|
ApiEndpoint,
|
||||||
HttpMethod,
|
download_url_to_video_output,
|
||||||
SynchronousOperation,
|
poll_op,
|
||||||
PollingOperation,
|
sync_op,
|
||||||
|
tensor_to_base64_string,
|
||||||
)
|
)
|
||||||
|
|
||||||
from comfy_api_nodes.util import downscale_image_tensor, tensor_to_base64_string
|
|
||||||
|
|
||||||
AVERAGE_DURATION_VIDEO_GEN = 32
|
AVERAGE_DURATION_VIDEO_GEN = 32
|
||||||
MODELS_MAP = {
|
MODELS_MAP = {
|
||||||
"veo-2.0-generate-001": "veo-2.0-generate-001",
|
"veo-2.0-generate-001": "veo-2.0-generate-001",
|
||||||
@ -32,28 +28,6 @@ MODELS_MAP = {
|
|||||||
"veo-3.0-fast-generate-001": "veo-3.0-fast-generate-001",
|
"veo-3.0-fast-generate-001": "veo-3.0-fast-generate-001",
|
||||||
}
|
}
|
||||||
|
|
||||||
def convert_image_to_base64(image: torch.Tensor):
|
|
||||||
if image is None:
|
|
||||||
return None
|
|
||||||
|
|
||||||
scaled_image = downscale_image_tensor(image, total_pixels=2048*2048)
|
|
||||||
return tensor_to_base64_string(scaled_image)
|
|
||||||
|
|
||||||
|
|
||||||
def get_video_url_from_response(poll_response: VeoGenVidPollResponse) -> Optional[str]:
|
|
||||||
if (
|
|
||||||
poll_response.response
|
|
||||||
and hasattr(poll_response.response, "videos")
|
|
||||||
and poll_response.response.videos
|
|
||||||
and len(poll_response.response.videos) > 0
|
|
||||||
):
|
|
||||||
video = poll_response.response.videos[0]
|
|
||||||
else:
|
|
||||||
return None
|
|
||||||
if hasattr(video, "gcsUri") and video.gcsUri:
|
|
||||||
return str(video.gcsUri)
|
|
||||||
return None
|
|
||||||
|
|
||||||
|
|
||||||
class VeoVideoGenerationNode(IO.ComfyNode):
|
class VeoVideoGenerationNode(IO.ComfyNode):
|
||||||
"""
|
"""
|
||||||
@ -166,18 +140,13 @@ class VeoVideoGenerationNode(IO.ComfyNode):
|
|||||||
# Prepare the instances for the request
|
# Prepare the instances for the request
|
||||||
instances = []
|
instances = []
|
||||||
|
|
||||||
instance = {
|
instance = {"prompt": prompt}
|
||||||
"prompt": prompt
|
|
||||||
}
|
|
||||||
|
|
||||||
# Add image if provided
|
# Add image if provided
|
||||||
if image is not None:
|
if image is not None:
|
||||||
image_base64 = convert_image_to_base64(image)
|
image_base64 = tensor_to_base64_string(image)
|
||||||
if image_base64:
|
if image_base64:
|
||||||
instance["image"] = {
|
instance["image"] = {"bytesBase64Encoded": image_base64, "mimeType": "image/png"}
|
||||||
"bytesBase64Encoded": image_base64,
|
|
||||||
"mimeType": "image/png"
|
|
||||||
}
|
|
||||||
|
|
||||||
instances.append(instance)
|
instances.append(instance)
|
||||||
|
|
||||||
@ -198,116 +167,74 @@ class VeoVideoGenerationNode(IO.ComfyNode):
|
|||||||
if "veo-3.0" in model:
|
if "veo-3.0" in model:
|
||||||
parameters["generateAudio"] = generate_audio
|
parameters["generateAudio"] = generate_audio
|
||||||
|
|
||||||
auth = {
|
initial_response = await sync_op(
|
||||||
"auth_token": cls.hidden.auth_token_comfy_org,
|
cls,
|
||||||
"comfy_api_key": cls.hidden.api_key_comfy_org,
|
ApiEndpoint(path=f"/proxy/veo/{model}/generate", method="POST"),
|
||||||
}
|
response_model=VeoGenVidResponse,
|
||||||
# Initial request to start video generation
|
data=VeoGenVidRequest(
|
||||||
initial_operation = SynchronousOperation(
|
|
||||||
endpoint=ApiEndpoint(
|
|
||||||
path=f"/proxy/veo/{model}/generate",
|
|
||||||
method=HttpMethod.POST,
|
|
||||||
request_model=VeoGenVidRequest,
|
|
||||||
response_model=VeoGenVidResponse
|
|
||||||
),
|
|
||||||
request=VeoGenVidRequest(
|
|
||||||
instances=instances,
|
instances=instances,
|
||||||
parameters=parameters
|
parameters=parameters,
|
||||||
),
|
),
|
||||||
auth_kwargs=auth,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
initial_response = await initial_operation.execute()
|
|
||||||
operation_name = initial_response.name
|
|
||||||
|
|
||||||
logging.info("Veo generation started with operation name: %s", operation_name)
|
|
||||||
|
|
||||||
# Define status extractor function
|
|
||||||
def status_extractor(response):
|
def status_extractor(response):
|
||||||
# Only return "completed" if the operation is done, regardless of success or failure
|
# Only return "completed" if the operation is done, regardless of success or failure
|
||||||
# We'll check for errors after polling completes
|
# We'll check for errors after polling completes
|
||||||
return "completed" if response.done else "pending"
|
return "completed" if response.done else "pending"
|
||||||
|
|
||||||
# Define progress extractor function
|
poll_response = await poll_op(
|
||||||
def progress_extractor(response):
|
cls,
|
||||||
# Could be enhanced if the API provides progress information
|
ApiEndpoint(path=f"/proxy/veo/{model}/poll", method="POST"),
|
||||||
return None
|
response_model=VeoGenVidPollResponse,
|
||||||
|
|
||||||
# Define the polling operation
|
|
||||||
poll_operation = PollingOperation(
|
|
||||||
poll_endpoint=ApiEndpoint(
|
|
||||||
path=f"/proxy/veo/{model}/poll",
|
|
||||||
method=HttpMethod.POST,
|
|
||||||
request_model=VeoGenVidPollRequest,
|
|
||||||
response_model=VeoGenVidPollResponse
|
|
||||||
),
|
|
||||||
completed_statuses=["completed"],
|
|
||||||
failed_statuses=[], # No failed statuses, we'll handle errors after polling
|
|
||||||
status_extractor=status_extractor,
|
status_extractor=status_extractor,
|
||||||
progress_extractor=progress_extractor,
|
data=VeoGenVidPollRequest(
|
||||||
request=VeoGenVidPollRequest(
|
operationName=initial_response.name,
|
||||||
operationName=operation_name
|
|
||||||
),
|
),
|
||||||
auth_kwargs=auth,
|
|
||||||
poll_interval=5.0,
|
poll_interval=5.0,
|
||||||
result_url_extractor=get_video_url_from_response,
|
|
||||||
node_id=cls.hidden.unique_id,
|
|
||||||
estimated_duration=AVERAGE_DURATION_VIDEO_GEN,
|
estimated_duration=AVERAGE_DURATION_VIDEO_GEN,
|
||||||
)
|
)
|
||||||
|
|
||||||
# Execute the polling operation
|
|
||||||
poll_response = await poll_operation.execute()
|
|
||||||
|
|
||||||
# Now check for errors in the final response
|
# Now check for errors in the final response
|
||||||
# Check for error in poll response
|
# Check for error in poll response
|
||||||
if hasattr(poll_response, 'error') and poll_response.error:
|
if poll_response.error:
|
||||||
error_message = f"Veo API error: {poll_response.error.message} (code: {poll_response.error.code})"
|
raise Exception(f"Veo API error: {poll_response.error.message} (code: {poll_response.error.code})")
|
||||||
logging.error(error_message)
|
|
||||||
raise Exception(error_message)
|
|
||||||
|
|
||||||
# Check for RAI filtered content
|
# Check for RAI filtered content
|
||||||
if (hasattr(poll_response.response, 'raiMediaFilteredCount') and
|
if (
|
||||||
poll_response.response.raiMediaFilteredCount > 0):
|
hasattr(poll_response.response, "raiMediaFilteredCount")
|
||||||
|
and poll_response.response.raiMediaFilteredCount > 0
|
||||||
|
):
|
||||||
|
|
||||||
# Extract reason message if available
|
# Extract reason message if available
|
||||||
if (hasattr(poll_response.response, 'raiMediaFilteredReasons') and
|
if (
|
||||||
poll_response.response.raiMediaFilteredReasons):
|
hasattr(poll_response.response, "raiMediaFilteredReasons")
|
||||||
|
and poll_response.response.raiMediaFilteredReasons
|
||||||
|
):
|
||||||
reason = poll_response.response.raiMediaFilteredReasons[0]
|
reason = poll_response.response.raiMediaFilteredReasons[0]
|
||||||
error_message = f"Content filtered by Google's Responsible AI practices: {reason} ({poll_response.response.raiMediaFilteredCount} videos filtered.)"
|
error_message = f"Content filtered by Google's Responsible AI practices: {reason} ({poll_response.response.raiMediaFilteredCount} videos filtered.)"
|
||||||
else:
|
else:
|
||||||
error_message = f"Content filtered by Google's Responsible AI practices ({poll_response.response.raiMediaFilteredCount} videos filtered.)"
|
error_message = f"Content filtered by Google's Responsible AI practices ({poll_response.response.raiMediaFilteredCount} videos filtered.)"
|
||||||
|
|
||||||
logging.error(error_message)
|
|
||||||
raise Exception(error_message)
|
raise Exception(error_message)
|
||||||
|
|
||||||
# Extract video data
|
# Extract video data
|
||||||
if poll_response.response and hasattr(poll_response.response, 'videos') and poll_response.response.videos and len(poll_response.response.videos) > 0:
|
if (
|
||||||
|
poll_response.response
|
||||||
|
and hasattr(poll_response.response, "videos")
|
||||||
|
and poll_response.response.videos
|
||||||
|
and len(poll_response.response.videos) > 0
|
||||||
|
):
|
||||||
video = poll_response.response.videos[0]
|
video = poll_response.response.videos[0]
|
||||||
|
|
||||||
# Check if video is provided as base64 or URL
|
# Check if video is provided as base64 or URL
|
||||||
if hasattr(video, 'bytesBase64Encoded') and video.bytesBase64Encoded:
|
if hasattr(video, "bytesBase64Encoded") and video.bytesBase64Encoded:
|
||||||
# Decode base64 string to bytes
|
return IO.NodeOutput(VideoFromFile(BytesIO(base64.b64decode(video.bytesBase64Encoded))))
|
||||||
video_data = base64.b64decode(video.bytesBase64Encoded)
|
|
||||||
elif hasattr(video, 'gcsUri') and video.gcsUri:
|
|
||||||
# Download from URL
|
|
||||||
async with aiohttp.ClientSession() as session:
|
|
||||||
async with session.get(video.gcsUri) as video_response:
|
|
||||||
video_data = await video_response.content.read()
|
|
||||||
else:
|
|
||||||
raise Exception("Video returned but no data or URL was provided")
|
|
||||||
else:
|
|
||||||
raise Exception("Video generation completed but no video was returned")
|
|
||||||
|
|
||||||
if not video_data:
|
if hasattr(video, "gcsUri") and video.gcsUri:
|
||||||
raise Exception("No video data was returned")
|
return IO.NodeOutput(await download_url_to_video_output(video.gcsUri))
|
||||||
|
|
||||||
logging.info("Video generation completed successfully")
|
raise Exception("Video returned but no data or URL was provided")
|
||||||
|
raise Exception("Video generation completed but no video was returned")
|
||||||
# Convert video data to BytesIO object
|
|
||||||
video_io = BytesIO(video_data)
|
|
||||||
|
|
||||||
# Return VideoFromFile object
|
|
||||||
return IO.NodeOutput(VideoFromFile(video_io))
|
|
||||||
|
|
||||||
|
|
||||||
class Veo3VideoGenerationNode(VeoVideoGenerationNode):
|
class Veo3VideoGenerationNode(VeoVideoGenerationNode):
|
||||||
@ -391,7 +318,10 @@ class Veo3VideoGenerationNode(VeoVideoGenerationNode):
|
|||||||
IO.Combo.Input(
|
IO.Combo.Input(
|
||||||
"model",
|
"model",
|
||||||
options=[
|
options=[
|
||||||
"veo-3.1-generate", "veo-3.1-fast-generate", "veo-3.0-generate-001", "veo-3.0-fast-generate-001"
|
"veo-3.1-generate",
|
||||||
|
"veo-3.1-fast-generate",
|
||||||
|
"veo-3.0-generate-001",
|
||||||
|
"veo-3.0-fast-generate-001",
|
||||||
],
|
],
|
||||||
default="veo-3.0-generate-001",
|
default="veo-3.0-generate-001",
|
||||||
tooltip="Veo 3 model to use for video generation",
|
tooltip="Veo 3 model to use for video generation",
|
||||||
@ -424,5 +354,6 @@ class VeoExtension(ComfyExtension):
|
|||||||
Veo3VideoGenerationNode,
|
Veo3VideoGenerationNode,
|
||||||
]
|
]
|
||||||
|
|
||||||
|
|
||||||
async def comfy_entrypoint() -> VeoExtension:
|
async def comfy_entrypoint() -> VeoExtension:
|
||||||
return VeoExtension()
|
return VeoExtension()
|
||||||
|
|||||||
@ -136,6 +136,7 @@ async def poll_op(
|
|||||||
completed_statuses: Optional[list[Union[str, int]]] = None,
|
completed_statuses: Optional[list[Union[str, int]]] = None,
|
||||||
failed_statuses: Optional[list[Union[str, int]]] = None,
|
failed_statuses: Optional[list[Union[str, int]]] = None,
|
||||||
queued_statuses: Optional[list[Union[str, int]]] = None,
|
queued_statuses: Optional[list[Union[str, int]]] = None,
|
||||||
|
data: Optional[BaseModel] = None,
|
||||||
poll_interval: float = 5.0,
|
poll_interval: float = 5.0,
|
||||||
max_poll_attempts: int = 120,
|
max_poll_attempts: int = 120,
|
||||||
timeout_per_poll: float = 120.0,
|
timeout_per_poll: float = 120.0,
|
||||||
@ -155,6 +156,7 @@ async def poll_op(
|
|||||||
completed_statuses=completed_statuses,
|
completed_statuses=completed_statuses,
|
||||||
failed_statuses=failed_statuses,
|
failed_statuses=failed_statuses,
|
||||||
queued_statuses=queued_statuses,
|
queued_statuses=queued_statuses,
|
||||||
|
data=data,
|
||||||
poll_interval=poll_interval,
|
poll_interval=poll_interval,
|
||||||
max_poll_attempts=max_poll_attempts,
|
max_poll_attempts=max_poll_attempts,
|
||||||
timeout_per_poll=timeout_per_poll,
|
timeout_per_poll=timeout_per_poll,
|
||||||
@ -229,6 +231,7 @@ async def poll_op_raw(
|
|||||||
completed_statuses: Optional[list[Union[str, int]]] = None,
|
completed_statuses: Optional[list[Union[str, int]]] = None,
|
||||||
failed_statuses: Optional[list[Union[str, int]]] = None,
|
failed_statuses: Optional[list[Union[str, int]]] = None,
|
||||||
queued_statuses: Optional[list[Union[str, int]]] = None,
|
queued_statuses: Optional[list[Union[str, int]]] = None,
|
||||||
|
data: Optional[Union[dict[str, Any], BaseModel]] = None,
|
||||||
poll_interval: float = 5.0,
|
poll_interval: float = 5.0,
|
||||||
max_poll_attempts: int = 120,
|
max_poll_attempts: int = 120,
|
||||||
timeout_per_poll: float = 120.0,
|
timeout_per_poll: float = 120.0,
|
||||||
@ -289,6 +292,7 @@ async def poll_op_raw(
|
|||||||
resp_json = await sync_op_raw(
|
resp_json = await sync_op_raw(
|
||||||
cls,
|
cls,
|
||||||
poll_endpoint,
|
poll_endpoint,
|
||||||
|
data=data,
|
||||||
timeout=timeout_per_poll,
|
timeout=timeout_per_poll,
|
||||||
max_retries=max_retries_per_poll,
|
max_retries=max_retries_per_poll,
|
||||||
retry_delay=retry_delay_per_poll,
|
retry_delay=retry_delay_per_poll,
|
||||||
|
|||||||
Loading…
x
Reference in New Issue
Block a user