[Partner Nodes] feat(Wan): add Wan 3.0 video generation nodes (#15843)

Signed-off-by: Alexander Piskun <bigcat88@icloud.com>
This commit is contained in:
Alexander Piskun
2026-08-24 20:30:15 +04:00
committed by GitHub
parent b78cec879b
commit 41218e4873
2 changed files with 447 additions and 0 deletions

View File

@@ -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(...)

View File

@@ -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<idx>\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,