From 41218e4873220205f633b3d3ab0a4207141be205 Mon Sep 17 00:00:00 2001 From: Alexander Piskun <13381981+bigcat88@users.noreply.github.com> Date: Mon, 24 Aug 2026 20:30:15 +0400 Subject: [PATCH] [Partner Nodes] feat(Wan): add Wan 3.0 video generation nodes (#15843) Signed-off-by: Alexander Piskun --- comfy_api_nodes/apis/wan.py | 26 +++ comfy_api_nodes/nodes_wan.py | 421 +++++++++++++++++++++++++++++++++++ 2 files changed, 447 insertions(+) diff --git a/comfy_api_nodes/apis/wan.py b/comfy_api_nodes/apis/wan.py index c64acae97..523ac68c1 100644 --- a/comfy_api_nodes/apis/wan.py +++ b/comfy_api_nodes/apis/wan.py @@ -184,6 +184,32 @@ class Wan27Text2VideoTaskCreationRequest(BaseModel): parameters: Wan27Text2VideoParametersField = Field(...) +class Wan3MediaItem(BaseModel): + type: str = Field(...) + url: str = Field(...) + + +class Wan3InputField(BaseModel): + prompt: str | None = Field(None) + media: list[Wan3MediaItem] | None = Field(None) + + +class Wan3ParametersField(BaseModel): + resolution: str = Field(...) + ratio: str = Field(...) + duration: int = Field(..., ge=-1, le=30) + seed: int = Field(..., ge=0, le=2147483647) + audio: bool = Field(True) + prompt_extend: bool = Field(True) + watermark: bool = Field(False) + + +class Wan3TaskCreationRequest(BaseModel): + model: str = Field(...) + input: Wan3InputField = Field(...) + parameters: Wan3ParametersField = Field(...) + + class TaskCreationOutputField(BaseModel): task_id: str = Field(...) task_status: str = Field(...) diff --git a/comfy_api_nodes/nodes_wan.py b/comfy_api_nodes/nodes_wan.py index 1782739fd..b11c528bc 100644 --- a/comfy_api_nodes/nodes_wan.py +++ b/comfy_api_nodes/nodes_wan.py @@ -34,6 +34,10 @@ from comfy_api_nodes.apis.wan import ( Wan27VideoEditInputField, Wan27VideoEditParametersField, Wan27VideoEditTaskCreationRequest, + Wan3InputField, + Wan3MediaItem, + Wan3ParametersField, + Wan3TaskCreationRequest, ) from comfy_api_nodes.util import ( ApiEndpoint, @@ -53,10 +57,41 @@ from comfy_api_nodes.util import ( validate_string, validate_video_duration, ) +from comfy_api_nodes.util.client import FAILED_STATUSES, QUEUED_STATUSES RES_IN_PARENS = re.compile(r"\((\d+)\s*[x×]\s*(\d+)\)") +WAN3_QUEUED_STATUSES = [*QUEUED_STATUSES, "pending"] +WAN3_FAILED_STATUSES = [*FAILED_STATUSES, "unknown"] + +_WAN3_REF_TAG_RE = re.compile(r"@(image|video|audio)(?P\d*)(?!\w)", re.IGNORECASE | re.ASCII) + + +def _wan3_rewrite_reference_prompt(prompt: str, counts: dict[str, int]) -> str: + parts = [] + pos = 0 + prev_end = -1 + for match in _WAN3_REF_TAG_RE.finditer(prompt): + start = match.start() + before = prompt[start - 1] if start > 0 else "" + if (before.isascii() and (before.isalnum() or before == "_")) and start != prev_end: + continue + kind = match.group(1).lower() + idx = int(match.group("idx") or 1) + total = counts[kind] + if not 1 <= idx <= total: + raise ValueError( + f"The prompt references @{kind.capitalize()}{idx}, " + f"but only {total} reference {kind} inputs are connected." + ) + parts.append(prompt[pos:start]) + parts.append(f"{kind.capitalize()} {idx}") + pos = match.end() + prev_end = match.end() + parts.append(prompt[pos:]) + return "".join(parts) + class WanTextToImageApi(IO.ComfyNode): @classmethod @@ -1648,6 +1683,390 @@ class Wan2ReferenceVideoApi(IO.ComfyNode): return IO.NodeOutput(await download_url_to_video_output(response.output.video_url)) +class Wan3ReferenceToVideoApi(IO.ComfyNode): + @classmethod + def define_schema(cls): + return IO.Schema( + node_id="Wan3ReferenceToVideoApi", + display_name="Wan 3.0 Reference to Video", + category="partner/video/Wan", + description="Generates a video from a text prompt and optional reference images, videos, and audio " + "using the Wan 3.0 model. Reference media can be combined freely and mentioned in the prompt " + "as @Image1, @Video1, @Audio1.", + inputs=[ + IO.DynamicCombo.Input( + "model", + options=[ + IO.DynamicCombo.Option( + "wan3.0-video", + [ + IO.String.Input( + "prompt", + multiline=True, + default="", + tooltip="Prompt describing the elements and visual features. " + "Supports English and Chinese. Refer to connected reference media " + "as @Image1, @Video1, @Audio1, numbered per type in input order.", + ), + IO.Combo.Input( + "resolution", + options=["1080P", "720P", "480P"], + ), + IO.Combo.Input( + "ratio", + options=["adaptive", "16:9", "9:16", "1:1", "4:3", "3:4"], + tooltip="Aspect ratio of the output video. With 'adaptive', the output " + "dimensions are derived from the input media.", + ), + IO.Combo.Input( + "duration", + options=["auto", *(str(i) for i in range(2, 31))], + default="5", + tooltip="Output duration in seconds. With 'auto', the model chooses " + "a duration that fits the prompt and reference media. The combined " + "duration of reference videos and output must not exceed 30 seconds.", + ), + IO.Boolean.Input( + "audio", + default=True, + tooltip="Whether the output video contains an audio track.", + ), + IO.Boolean.Input( + "prompt_extend", + default=True, + tooltip="Whether to enhance the prompt with AI assistance.", + advanced=True, + ), + IO.Autogrow.Input( + "reference_images", + template=IO.Autogrow.TemplateNames( + IO.Image.Input("reference_image"), + names=[f"image{i}" for i in range(1, 11)], + min=0, + ), + ), + IO.Autogrow.Input( + "reference_videos", + template=IO.Autogrow.TemplateNames( + IO.Video.Input("reference_video"), + names=[f"video{i}" for i in range(1, 6)], + min=0, + ), + ), + IO.Autogrow.Input( + "reference_audios", + template=IO.Autogrow.TemplateNames( + IO.Audio.Input("reference_audio"), + names=[f"audio{i}" for i in range(1, 6)], + min=0, + ), + ), + ], + ), + ], + ), + IO.Int.Input( + "seed", + default=42, + min=0, + max=2147483647, + step=1, + display_mode=IO.NumberDisplay.number, + control_after_generate=True, + tooltip="Seed to use for generation.", + ), + IO.Boolean.Input( + "watermark", + default=False, + tooltip="Whether to add an AI-generated watermark to the result.", + advanced=True, + ), + ], + outputs=[ + IO.Video.Output(), + ], + hidden=[ + IO.Hidden.auth_token_comfy_org, + IO.Hidden.api_key_comfy_org, + IO.Hidden.unique_id, + ], + is_api_node=True, + price_badge=IO.PriceBadge( + depends_on=IO.PriceBadgeDepends(widgets=["model", "model.resolution", "model.duration"]), + expr=""" + ( + $ppsTable := { "480p": 0.0715, "720p": 0.143, "1080p": 0.286 }; + $pps := $lookup($ppsTable, $lookup(widgets, "model.resolution")); + $dur := $lookup(widgets, "model.duration"); + $dur = "auto" + ? { "type": "usd", "usd": $pps, "format": {"suffix": "/second"} } + : ( + $minUsd := $round($number($dur) * $pps, 2); + $maxUsd := $round($min([$number($dur) + 15, 30]) * $pps, 2); + $minUsd = $maxUsd + ? { "type": "usd", "usd": $minUsd } + : { "type": "range_usd", "min_usd": $minUsd, "max_usd": $maxUsd } + ) + ) + """, + ), + ) + + @classmethod + async def execute( + cls, + model: dict, + seed: int, + watermark: bool, + ): + reference_images = model.get("reference_images", {}) + reference_videos = model.get("reference_videos", {}) + reference_audios = model.get("reference_audios", {}) + for key in reference_images: + if get_number_of_images(reference_images[key]) != 1: + raise ValueError(f"Reference image input '{key}' must contain exactly one image, not a batch.") + total_video_seconds = 0.0 + for key in reference_videos: + validate_video_duration(reference_videos[key], max_duration=15) + try: + total_video_seconds += reference_videos[key].get_duration() + except Exception: + pass + if total_video_seconds > 15.0001: + raise ValueError( + f"The total duration of reference videos ({total_video_seconds:.2f}s) exceeds the 15s limit." + ) + total_audio_seconds = 0.0 + for key in reference_audios: + validate_audio_duration(reference_audios[key], max_duration=15) + total_audio_seconds += reference_audios[key]["waveform"].shape[-1] / int( + reference_audios[key]["sample_rate"] + ) + if total_audio_seconds > 15.0001: + raise ValueError( + f"The total duration of reference audios ({total_audio_seconds:.2f}s) exceeds the 15s limit." + ) + duration = -1 if model["duration"] == "auto" else int(model["duration"]) + if duration != -1 and total_video_seconds + duration > 30.0001: + raise ValueError( + f"Reference video duration ({total_video_seconds:.2f}s) plus output duration ({duration}s) " + "exceeds the 30s combined limit." + ) + prompt = _wan3_rewrite_reference_prompt( + model["prompt"], + {"image": len(reference_images), "video": len(reference_videos), "audio": len(reference_audios)}, + ) + validate_string(prompt, strip_whitespace=False, max_length=20000) + if not prompt.strip() and not (reference_images or reference_videos or reference_audios): + raise ValueError("Provide a prompt or at least one reference input.") + media = [] + for key in reference_images: + media.append( + Wan3MediaItem(type="reference_image", url=await upload_image_to_comfyapi(cls, reference_images[key])) + ) + for key in reference_videos: + media.append( + Wan3MediaItem(type="reference_video", url=await upload_video_to_comfyapi(cls, reference_videos[key])) + ) + for key in reference_audios: + media.append( + Wan3MediaItem( + type="reference_audio", + url=await upload_audio_to_comfyapi( + cls, + reference_audios[key], + container_format="mp3", + codec_name="libmp3lame", + mime_type="audio/mpeg", + ), + ) + ) + initial_response = await sync_op( + cls, + ApiEndpoint(path="/proxy/wan/api/v1/services/aigc/video-generation/video-synthesis", method="POST"), + response_model=TaskCreationResponse, + data=Wan3TaskCreationRequest( + model=model["model"], + input=Wan3InputField(prompt=prompt or None, media=media or None), + parameters=Wan3ParametersField( + resolution=model["resolution"], + ratio=model["ratio"], + duration=duration, + seed=seed, + audio=model["audio"], + prompt_extend=model["prompt_extend"], + watermark=watermark, + ), + ), + ) + if not initial_response.output: + raise Exception(f"An unknown error occurred: {initial_response.code} - {initial_response.message}") + response = await poll_op( + cls, + ApiEndpoint(path=f"/proxy/wan/api/v1/tasks/{initial_response.output.task_id}"), + response_model=VideoTaskStatusResponse, + status_extractor=lambda x: x.output.task_status, + queued_statuses=WAN3_QUEUED_STATUSES, + failed_statuses=WAN3_FAILED_STATUSES, + poll_interval=10, + ) + return IO.NodeOutput(await download_url_to_video_output(response.output.video_url)) + + +class Wan3ImageToVideoApi(IO.ComfyNode): + @classmethod + def define_schema(cls): + return IO.Schema( + node_id="Wan3ImageToVideoApi", + display_name="Wan 3.0 Image to Video", + category="partner/video/Wan", + description="Generates a video from a first-frame image, with optional last-frame control, " + "using the Wan 3.0 model.", + inputs=[ + IO.DynamicCombo.Input( + "model", + options=[ + IO.DynamicCombo.Option( + "wan3.0-video", + [ + IO.String.Input( + "prompt", + multiline=True, + default="", + tooltip="Prompt describing the elements and visual features. " + "Supports English and Chinese.", + ), + IO.Combo.Input( + "resolution", + options=["1080P", "720P", "480P"], + ), + IO.Combo.Input( + "ratio", + options=["adaptive", "16:9", "9:16", "1:1", "4:3", "3:4"], + tooltip="Aspect ratio of the output video. With 'adaptive', the output " + "dimensions are derived from the first frame.", + ), + IO.Combo.Input( + "duration", + options=["auto", *(str(i) for i in range(2, 31))], + default="5", + tooltip="Output duration in seconds. With 'auto', the model chooses " + "a duration that fits the prompt.", + ), + IO.Boolean.Input( + "audio", + default=True, + tooltip="Whether the output video contains an audio track.", + ), + IO.Boolean.Input( + "prompt_extend", + default=True, + tooltip="Whether to enhance the prompt with AI assistance.", + advanced=True, + ), + ], + ), + ], + ), + IO.Image.Input( + "first_frame", + tooltip="First frame image.", + ), + IO.Image.Input( + "last_frame", + optional=True, + tooltip="Last frame image. The model generates a video transitioning from first to last frame.", + ), + IO.Int.Input( + "seed", + default=42, + min=0, + max=2147483647, + step=1, + display_mode=IO.NumberDisplay.number, + control_after_generate=True, + tooltip="Seed to use for generation.", + ), + IO.Boolean.Input( + "watermark", + default=False, + tooltip="Whether to add an AI-generated watermark to the result.", + advanced=True, + ), + ], + outputs=[ + IO.Video.Output(), + ], + hidden=[ + IO.Hidden.auth_token_comfy_org, + IO.Hidden.api_key_comfy_org, + IO.Hidden.unique_id, + ], + is_api_node=True, + price_badge=IO.PriceBadge( + depends_on=IO.PriceBadgeDepends(widgets=["model", "model.resolution", "model.duration"]), + expr=""" + ( + $ppsTable := { "480p": 0.0715, "720p": 0.143, "1080p": 0.286 }; + $pps := $lookup($ppsTable, $lookup(widgets, "model.resolution")); + $dur := $lookup(widgets, "model.duration"); + $dur = "auto" + ? { "type": "usd", "usd": $pps, "format": {"suffix": "/second"} } + : { "type": "usd", "usd": $round($number($dur) * $pps, 2) } + ) + """, + ), + ) + + @classmethod + async def execute( + cls, + model: dict, + first_frame: Input.Image, + seed: int, + watermark: bool, + last_frame: Input.Image | None = None, + ): + if get_number_of_images(first_frame) != 1: + raise ValueError("Exactly one first_frame image is required.") + if last_frame is not None and get_number_of_images(last_frame) != 1: + raise ValueError("Exactly one last_frame image is required.") + validate_string(model["prompt"], strip_whitespace=False, max_length=20000) + media = [Wan3MediaItem(type="first_frame", url=await upload_image_to_comfyapi(cls, first_frame))] + if last_frame is not None: + media.append(Wan3MediaItem(type="last_frame", url=await upload_image_to_comfyapi(cls, last_frame))) + initial_response = await sync_op( + cls, + ApiEndpoint(path="/proxy/wan/api/v1/services/aigc/video-generation/video-synthesis", method="POST"), + response_model=TaskCreationResponse, + data=Wan3TaskCreationRequest( + model=model["model"], + input=Wan3InputField(prompt=model["prompt"] or None, media=media), + parameters=Wan3ParametersField( + resolution=model["resolution"], + ratio=model["ratio"], + duration=-1 if model["duration"] == "auto" else int(model["duration"]), + seed=seed, + audio=model["audio"], + prompt_extend=model["prompt_extend"], + watermark=watermark, + ), + ), + ) + if not initial_response.output: + raise Exception(f"An unknown error occurred: {initial_response.code} - {initial_response.message}") + response = await poll_op( + cls, + ApiEndpoint(path=f"/proxy/wan/api/v1/tasks/{initial_response.output.task_id}"), + response_model=VideoTaskStatusResponse, + status_extractor=lambda x: x.output.task_status, + queued_statuses=WAN3_QUEUED_STATUSES, + failed_statuses=WAN3_FAILED_STATUSES, + poll_interval=10, + ) + return IO.NodeOutput(await download_url_to_video_output(response.output.video_url)) + + class HappyHorseTextToVideoApi(IO.ComfyNode): @classmethod def define_schema(cls): @@ -2342,6 +2761,8 @@ class WanApiExtension(ComfyExtension): Wan2VideoContinuationApi, Wan2VideoEditApi, Wan2ReferenceVideoApi, + Wan3ReferenceToVideoApi, + Wan3ImageToVideoApi, HappyHorseTextToVideoApi, HappyHorseImageToVideoApi, HappyHorseVideoEditApi,