diff --git a/README.md b/README.md index 650fb70..2784c0e 100644 --- a/README.md +++ b/README.md @@ -3,14 +3,19 @@ [![License: MIT](https://img.shields.io/badge/License-MIT-yellow.svg)](LICENSE) -A custom node for ComfyUI that provides **seamless integration** with the **Wan models** (text-to-image and image-to-video) from **Alibaba Cloud Model Studio**. This solution delivers cutting-edge image and video generation capabilities directly within ComfyUI. +A custom node for ComfyUI that provides **seamless integration** with the **Wan models** from **Alibaba Cloud Model Studio**. This solution delivers cutting-edge image and video generation capabilities directly within ComfyUI, supporting both international and Mainland China regions. ### Why Choose Alibaba Cloud Model Studio? This is a direct integration with Alibaba Cloud's Model Studio service, not a third-party wrapper or local model implementation. Benefits include: - **Enterprise-Grade Infrastructure**: Leverages Alibaba Cloud's battle-tested AI platform serving millions of requests daily -- **State-of-the-Art Models**: Access to the latest Wan models (wan2.2-t2i-flash, wan2.2-t2i-plus, wan2.2-i2v-flash, wan2.2-i2v-plus) with continuous updates +- **State-of-the-Art Models**: Access to the latest Wan models with continuous updates: + - **Text-to-Image**: wan2.2-t2i-flash (Speed Edition), wan2.2-t2i-plus (Professional Edition) + - **Image-to-Video**: wan2.2-i2v-flash (Speed Edition), wan2.2-i2v-plus (Professional Edition) + - **Text-to-Video**: wan2.2-t2v-plus (Professional Edition) + - **Image-to-Video (First/Last Frames)**: wan2.1-kf2v-plus (Professional Edition) + - **Universal Video Editing (VACE)**: wan2.1-vace-plus (Professional Edition) - Split into 5 specialized nodes for better usability - **Commercial Licensing**: Properly licensed for commercial use through Alibaba Cloud's terms of service - **Scalable Architecture**: Handles high-volume workloads with Alibaba Cloud's reliable infrastructure - **Security Compliance**: Follows Alibaba Cloud's security best practices with secure API key management @@ -21,12 +26,53 @@ This is a direct integration with Alibaba Cloud's Model Studio service, not a th **Model Authorization Required**: If you're using a non-default workspace or project in Alibaba Cloud, you may need to explicitly authorize access to the Wan models in your DashScope console. +## Regional Support + +This node supports both international and Mainland China Alibaba Cloud regions. By default, it uses the international region endpoints, but you can easily switch to Mainland China endpoints by modifying the variables in `wan_base.py`: + +- **International Region** (default): + - Video POST: `https://dashscope-intl.aliyuncs.com/api/v1/services/aigc/video-generation/video-synthesis` + - II2V POST: `https://dashscope-intl.aliyuncs.com/api/v1/services/aigc/image2video/video-synthesis` + - T2I POST: `https://dashscope-intl.aliyuncs.com/api/v1/services/aigc/text2image/image-synthesis` + - GET: `https://dashscope-intl.aliyuncs.com/api/v1/tasks/{task_id}` + +- **Mainland China Region**: + - Video POST: `https://dashscope.aliyuncs.com/api/v1/services/aigc/video-generation/video-synthesis` + - II2V POST: `https://dashscope.aliyuncs.com/api/v1/services/aigc/image2video/video-synthesis` + - T2I POST: `https://dashscope.aliyuncs.com/api/v1/services/aigc/text2image/image-synthesis` + - GET: `https://dashscope.aliyuncs.com/api/v1/tasks/{task_id}` + +To switch regions, simply modify the `API_ENDPOINT_POST_VIDEO`, `API_ENDPOINT_POST_II2V`, `API_ENDPOINT_POST_T2I`, and `API_ENDPOINT_GET` variables in `wan_base.py` to the corresponding Mainland China endpoints listed above. + +## Centralized Endpoint Management + +All API endpoints are centrally managed in the `wan_base.py` file, making it easy to maintain and switch between regions. This approach ensures consistency across all nodes and simplifies future updates. The centralized management includes: + +- `API_ENDPOINT_POST_VIDEO`: For general video generation nodes (I2V, T2V, VACE) +- `API_ENDPOINT_POST_II2V`: For image-to-video with first/last frames (II2V) +- `API_ENDPOINT_POST_T2I`: For text-to-image generation (T2I) +- `API_ENDPOINT_GET`: For task result polling (shared across all nodes) + +## Available Nodes + +| Node Name | Function | Model | Description | +|-----------|----------|-------|-------------| +| Wan Text-to-Image Generator | T2I | wan2.2-t2i-flash, wan2.2-t2i-plus | Generate images from text prompts with multiple resolution options | +| Wan Image-to-Video Generator | I2V | wan2.2-i2v-flash, wan2.2-i2v-plus | Create 5-second videos from a single image and text prompt | +| Wan Text-to-Video Generator | T2V | wan2.2-t2v-plus | Generate 5-second videos directly from text prompts | +| Wan Image-to-Video (First/Last Frame) Generator | II2V | wan2.1-kf2v-plus | Create 5-second videos using both first and last frame images | +| Wan VACE - Multi-Image Reference | VACE | wan2.1-vace-plus | Generate videos from multiple reference images | +| Wan VACE - Video Repainting | VACE | wan2.1-vace-plus | Repaint videos while preserving motion | +| Wan VACE - Local Video Editing | VACE | wan2.1-vace-plus | Locally edit specific areas of videos | +| Wan VACE - Video Extension | VACE | wan2.1-vace-plus | Extend videos with additional content | +| Wan VACE - Video Outpainting | VACE | wan2.1-vace-plus | Scale videos in different directions | + ## Features -- Generate images from text T2I using Wan models with selectable model types -- Generate 5-second videos from images and text prompts (I2V) using Wan models -- Configurable parameters: seed, resolution, prompt extension, watermark, negative prompts -- Powered by Alibaba Cloud's advanced Wan models +- **Regional Support**: Works with both international and Mainland China Alibaba Cloud regions +- **Configurable Parameters**: Seed, resolution, prompt extension, watermark, negative prompts, and more +- **Specialized VACE Nodes**: The Universal Video Editing (VACE) model has been split into 5 specialized nodes for better usability and focused functionality +- **Powered by Alibaba Cloud's Advanced Wan Models**: Access to state-of-the-art models with continuous updates ## Installation @@ -100,6 +146,18 @@ You can use services like Imgur, cloud storage providers, or your own web server - **seed**: Random seed for generation (0 for random) - **watermark**: Add Wan watermark to output +### Text-to-Video Generator +- **model**: Select the Wan model to use (wan2.2-t2v-plus) +- **prompt** (required): The text prompt for video generation +- **resolution**: Output video resolution (480P, 1080P) +- **negative_prompt**: Text describing content to avoid in the video +- **prompt_extend**: Enable intelligent prompt rewriting for better results +- **seed**: Random seed for generation (0 for random) +- **watermark**: Add Wan watermark to output +- **output_dir**: Directory where the generated video will be saved. Can be browsed and selected in ComfyUI. + +**Note**: To preview the generated video in ComfyUI, connect the output of this node to a "Load Video (Path)" node from ComfyUI-VideoHelperSuite. + ### Image-to-Video Generator - **model**: Select the Wan model to use (wan2.2-i2v-flash or wan2.2-i2v-plus) - **image_url**: Publicly accessible URL to the image for the first frame of the video @@ -113,12 +171,137 @@ You can use services like Imgur, cloud storage providers, or your own web server **Note**: To preview the generated video in ComfyUI, connect the output of this node to a "Load Video (Path)" node from ComfyUI-VideoHelperSuite. +### Image-to-Video (First/Last Frame) Generator +- **model**: Select the Wan model to use (wan2.1-kf2v-plus) +- **first_frame_url**: Publicly accessible URL to the first frame image +- **last_frame_url**: Publicly accessible URL to the last frame image +- **prompt** (required): The text prompt describing the video content and transition +- **resolution**: Output video resolution (720P) +- **negative_prompt**: Text describing content to avoid in the video +- **prompt_extend**: Enable intelligent prompt rewriting for better results +- **seed**: Random seed for generation (0 for random) +- **watermark**: Add Wan watermark to output +- **output_dir**: Directory where the generated video will be saved. Can be browsed and selected in ComfyUI. + +**Note**: To preview the generated video in ComfyUI, connect the output of this node to a "Load Video (Path)" node from ComfyUI-VideoHelperSuite. + +### Wan VACE - Multi-Image Reference +This node generates videos from multiple reference images using the Wan VACE model. + +**Parameters:** +- **model**: Select the Wan model to use (wan2.1-vace-plus) +- **prompt** (required): The text prompt describing the desired video content +- **ref_images_url** (required): Newline-separated URLs for reference images +- **obj_or_bg** (optional): Newline-separated values (obj/bg) corresponding to ref_images_url. + - If not provided, the node automatically assigns "obj" to all images except the last one, which is assigned "bg" + - Example: For 3 images, it automatically becomes ["obj", "obj", "bg"] +- **size**: Output video resolution (1280*720, 720*1280, 960*960, 832*1088, 1088*832) +- **seed**: Random seed for generation (0 for random) +- **prompt_extend**: Enable intelligent prompt rewriting for better results +- **watermark**: Add Wan watermark to output +- **output_dir**: Directory where the generated video will be saved. Can be browsed and selected in ComfyUI. + +**Automatic obj_or_bg Handling:** +The node intelligently handles the obj_or_bg parameter: +- For a single reference image: Automatically set to ["obj"] +- For multiple reference images: Automatically set to ["obj", "obj", ..., "bg"] where the last image is treated as background + +**Note**: To preview the generated video in ComfyUI, connect the output of this node to a "Load Video (Path)" node from ComfyUI-VideoHelperSuite. + +### Wan VACE - Video Repainting +This node repaints videos while preserving motion using the Wan VACE model. + +**Parameters:** +- **model**: Select the Wan model to use (wan2.1-vace-plus) +- **prompt** (required): The text prompt describing the desired video content +- **video_url** (required): URL of the input video to repaint +- **ref_images_url** (optional): Newline-separated URLs for reference images (only 1 image supported) +- **control_condition**: Method for video feature extraction (posebodyface, posebody, depth, scribble) +- **strength**: Control strength of the video feature extraction method (0.0-1.0, default: 1.0) +- **seed**: Random seed for generation (0 for random) +- **prompt_extend**: Enable intelligent prompt rewriting for better results +- **watermark**: Add Wan watermark to output +- **output_dir**: Directory where the generated video will be saved. Can be browsed and selected in ComfyUI. + +**Note**: To preview the generated video in ComfyUI, connect the output of this node to a "Load Video (Path)" node from ComfyUI-VideoHelperSuite. + +### Wan VACE - Local Video Editing +This node locally edits specific areas of videos using the Wan VACE model. + +**Parameters:** +- **model**: Select the Wan model to use (wan2.1-vace-plus) +- **prompt** (required): The text prompt describing the desired video content +- **video_url** (required): URL of the input video to edit +- **ref_images_url** (optional): Newline-separated URLs for reference images (only 1 image supported) +- **mask_image_url** (optional): URL of the mask image +- **mask_frame_id** (optional): Frame ID where the masked object appears (default: 1) +- **mask_video_url** (optional): URL of the mask video +- **control_condition** (optional): Method for video feature extraction (posebodyface, posebody, depth, scribble) +- **mask_type**: Behavior of the editing area (tracking, fixed) +- **expand_ratio**: Ratio for expanding the mask area outward (0.0-1.0, default: 0.05) +- **expand_mode**: Shape of the mask area (hull, bbox, original) +- **size**: Output video resolution (1280*720, 720*1280, 960*960, 832*1088, 1088*832) +- **seed**: Random seed for generation (0 for random) +- **prompt_extend**: Enable intelligent prompt rewriting for better results +- **watermark**: Add Wan watermark to output +- **output_dir**: Directory where the generated video will be saved. Can be browsed and selected in ComfyUI. + +**Note**: To preview the generated video in ComfyUI, connect the output of this node to a "Load Video (Path)" node from ComfyUI-VideoHelperSuite. + +### Wan VACE - Video Extension +This node extends videos with additional content using the Wan VACE model. + +**Parameters:** +- **model**: Select the Wan model to use (wan2.1-vace-plus) +- **prompt** (required): The text prompt describing the desired video content +- **first_frame_url** (optional): URL of the first frame image +- **last_frame_url** (optional): URL of the last frame image +- **first_clip_url** (optional): URL of the first video segment +- **last_clip_url** (optional): URL of the last video segment +- **video_url** (optional): URL of the reference video for motion features +- **control_condition** (optional): Method for video feature extraction (posebodyface, posebody, depth, scribble) +- **seed**: Random seed for generation (0 for random) +- **prompt_extend**: Enable intelligent prompt rewriting for better results +- **watermark**: Add Wan watermark to output +- **output_dir**: Directory where the generated video will be saved. Can be browsed and selected in ComfyUI. + +**Note**: To preview the generated video in ComfyUI, connect the output of this node to a "Load Video (Path)" node from ComfyUI-VideoHelperSuite. + +### Wan VACE - Video Outpainting +This node scales videos in different directions using the Wan VACE model. + +**Parameters:** +- **model**: Select the Wan model to use (wan2.1-vace-plus) +- **prompt** (required): The text prompt describing the desired video content +- **video_url** (required): URL of the input video to outpaint +- **top_scale**: Scale upward proportionally (1.0-2.0, default: 1.0) +- **bottom_scale**: Scale downward proportionally (1.0-2.0, default: 1.0) +- **left_scale**: Scale to the left proportionally (1.0-2.0, default: 1.0) +- **right_scale**: Scale to the right proportionally (1.0-2.0, default: 1.0) +- **seed**: Random seed for generation (0 for random) +- **prompt_extend**: Enable intelligent prompt rewriting for better results +- **watermark**: Add Wan watermark to output +- **output_dir**: Directory where the generated video will be saved. Can be browsed and selected in ComfyUI. + +**Note**: To preview the generated video in ComfyUI, connect the output of this node to a "Load Video (Path)" node from ComfyUI-VideoHelperSuite. + ## Examples ### Text-to-Image Generation Prompt: "Generate an image of a cat swimming under the water" -![Text-to-Image Example](media/ComfyUI_wan-t2i.png) +![Text-to-Image Example](media/ComfyUI_Wan-t2i.png) + +### Text-to-Video Generation +1. Add the "Wan Text-to-Video Generator" node to your workflow +2. Select the desired model (wan2.2-t2v-plus) +3. Connect a text input with your prompt (e.g., "A kitten running in the moonlight") +4. Optionally configure the output directory where the video will be saved (can be browsed in ComfyUI) +5. Execute the node +6. The node will return a path to the downloaded video file +7. To preview the video, connect the output to a "Load Video (Path)" node from ComfyUI-VideoHelperSuite + +![Text-to-Video Example](media/ComfyUI_Wan-t2v.png) ### Image-to-Video Generation 1. First frame: Provide a URL to an image (e.g., "https://example.com/your_image.png") @@ -126,7 +309,19 @@ Prompt: "Generate an image of a cat swimming under the water" 3. Output directory: "./videos" (default) or any custom path 4. To preview: Connect the output to a "Load Video (Path)" node from ComfyUI-VideoHelperSuite -![Image-to-Video Example](media/ComfyUI_wan-i2v.png) +![Image-to-Video Example](media/ComfyUI_Wan-i2v.png) + +### Image-to-Video (First/Last Frame) Generation +1. Add the "Wan Image-to-Video (First/Last Frame) Generator" node to your workflow +2. Select the desired model (wan2.1-kf2v-plus) +3. Provide publicly accessible URLs to the first and last frame images +4. Connect a text input with your prompt describing the video content and transition +5. Optionally configure the output directory where the video will be saved (can be browsed in ComfyUI) +6. Execute the node +7. The node will return a path to the downloaded video file +8. To preview the video, connect the output to a "Load Video (Path)" node from ComfyUI-VideoHelperSuite + +![Image-first-last-frame-to-Video Example](media/ComfyUI_Wan-ii2v.png) ## Security diff --git a/__init__.py b/__init__.py index 13b464a..13a2868 100644 --- a/__init__.py +++ b/__init__.py @@ -1,18 +1,41 @@ """ ComfyUI_Wan - A custom node for ComfyUI that integrates Wan models -for text-to-image and image-to-video generation. +for text-to-image, image-to-video, and text-to-video generation. """ -from .wan_nodes import WanT2IGenerator, WanI2VGenerator +from .wan_t2i import WanT2IGenerator +from .wan_i2v import WanI2VGenerator +from .wan_t2v import WanT2VGenerator +from .wan_ii2v import WanII2VGenerator +from .wan_vace_image_reference import WanVACEImageReference +from .wan_vace_video_repainting import WanVACEVideoRepainting +from .wan_vace_video_edit import WanVACEVideoEdit +from .wan_vace_video_extension import WanVACEVideoExtension +from .wan_vace_video_outpainting import WanVACEVideoOutpainting NODE_CLASS_MAPPINGS = { "WanT2IGenerator": WanT2IGenerator, "WanI2VGenerator": WanI2VGenerator, + "WanT2VGenerator": WanT2VGenerator, + "WanII2VGenerator": WanII2VGenerator, + "WanVACEImageReference": WanVACEImageReference, + "WanVACEVideoRepainting": WanVACEVideoRepainting, + "WanVACEVideoEdit": WanVACEVideoEdit, + "WanVACEVideoExtension": WanVACEVideoExtension, + "WanVACEVideoOutpainting": WanVACEVideoOutpainting, } NODE_DISPLAY_NAME_MAPPINGS = { "WanT2IGenerator": "Wan Text-to-Image Generator", "WanI2VGenerator": "Wan Image-to-Video Generator", + "WanT2VGenerator": "Wan Text-to-Video Generator", + "WanII2VGenerator": "Wan Image-to-Video (First/Last Frame) Generator", + "WanVACEImageReference": "Wan VACE - Multi-Image Reference", + "WanVACEVideoRepainting": "Wan VACE - Video Repainting", + "WanVACEVideoEdit": "Wan VACE - Local Video Editing", + "WanVACEVideoExtension": "Wan VACE - Video Extension", + "WanVACEVideoOutpainting": "Wan VACE - Video Outpainting", } -__all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS'] \ No newline at end of file +__all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS'] + diff --git a/media/ComfyUI_Wan-VACE-multi-image-reference.png b/media/ComfyUI_Wan-VACE-multi-image-reference.png new file mode 100644 index 0000000..ec7848f Binary files /dev/null and b/media/ComfyUI_Wan-VACE-multi-image-reference.png differ diff --git a/media/ComfyUI_Wan-ii2v.png b/media/ComfyUI_Wan-ii2v.png new file mode 100644 index 0000000..c2454dc Binary files /dev/null and b/media/ComfyUI_Wan-ii2v.png differ diff --git a/media/ComfyUI_Wan-t2v.png b/media/ComfyUI_Wan-t2v.png new file mode 100644 index 0000000..5013086 Binary files /dev/null and b/media/ComfyUI_Wan-t2v.png differ diff --git a/media/ComfyUI_wan-i2v.png b/media/ComfyUI_wan-i2v.png index bb8dcbb..2aabe7f 100644 Binary files a/media/ComfyUI_wan-i2v.png and b/media/ComfyUI_wan-i2v.png differ diff --git a/wan_base.py b/wan_base.py new file mode 100644 index 0000000..5716d5c --- /dev/null +++ b/wan_base.py @@ -0,0 +1,83 @@ +import os +import json +import requests +from PIL import Image +import numpy as np +import torch +import io +import base64 +from dotenv import load_dotenv +import sys +import pathlib + +# Import ComfyUI's folder_paths for directory browsing +try: + import folder_paths + COMFYUI_AVAILABLE = True +except ImportError: + COMFYUI_AVAILABLE = False + print("folder_paths not available, using default directory handling") + +# Load environment variables from .env file +# Try to load .env file from the current directory first +env_path = pathlib.Path(__file__).parent / '.env' +if env_path.exists(): + load_dotenv(dotenv_path=env_path) +else: + # Fallback to default behavior + load_dotenv() + +class WanAPIBase: + """Base class for Wan API interactions""" + + # API endpoints - International region (default) + # To use Mainland China region, change these URLs: + # Video POST: https://dashscope.aliyuncs.com/api/v1/services/aigc/video-generation/video-synthesis + # II2V POST: https://dashscope.aliyuncs.com/api/v1/services/aigc/image2video/video-synthesis + # T2I POST: https://dashscope.aliyuncs.com/api/v1/services/aigc/text2image/image-synthesis + # GET: https://dashscope.aliyuncs.com/api/v1/tasks/{task_id} + API_ENDPOINT_POST_VIDEO = "https://dashscope-intl.aliyuncs.com/api/v1/services/aigc/video-generation/video-synthesis" + API_ENDPOINT_POST_II2V = "https://dashscope-intl.aliyuncs.com/api/v1/services/aigc/image2video/video-synthesis" + API_ENDPOINT_POST_T2I = "https://dashscope-intl.aliyuncs.com/api/v1/services/aigc/text2image/image-synthesis" + API_ENDPOINT_GET = "https://dashscope-intl.aliyuncs.com/api/v1/tasks/{task_id}" + + def __init__(self): + self.api_key = os.getenv('DASHSCOPE_API_KEY') + # Strip any extra quotes or whitespace + if self.api_key: + self.api_key = self.api_key.strip().strip('"\'') + print(f"Initialized WanAPIBase with API key: {self.api_key[:8] if self.api_key else 'None'}...{self.api_key[-4:] if self.api_key else ''}") + + def check_api_key(self): + """Check if API key is set in environment variables""" + if not self.api_key: + raise ValueError("DASHSCOPE_API_KEY environment variable not set. " + "Please set it before using this node.") + return self.api_key + + def prepare_images(self, images): + """Convert images to base64 strings for API submission""" + image_data = [] + for i, image in enumerate(images, 1): + if image is not None: + # Convert tensor to PIL Image + if isinstance(image, torch.Tensor): + # Convert tensor to numpy array + image_np = image.cpu().numpy() + # If the tensor is in [0, 1] range, convert to [0, 255] + if image_np.max() <= 1.0: + image_np = (image_np * 255).astype(np.uint8) + # If tensor has shape [H, W, C], convert to PIL + pil_image = Image.fromarray(image_np.squeeze()) + else: + pil_image = image + + # Convert PIL image to base64 + buffer = io.BytesIO() + pil_image.save(buffer, format="PNG") + img_str = base64.b64encode(buffer.getvalue()).decode() + image_data.append({ + "id": str(i), + "data": img_str + }) + return image_data diff --git a/wan_i2v.py b/wan_i2v.py new file mode 100644 index 0000000..12707b2 --- /dev/null +++ b/wan_i2v.py @@ -0,0 +1,273 @@ +import os +import json +import requests +from PIL import Image +import numpy as np +import torch +import io +import base64 +from dotenv import load_dotenv +import sys +import pathlib +from datetime import datetime + +# Import the base class and COMFYUI_AVAILABLE flag +from .wan_base import WanAPIBase, COMFYUI_AVAILABLE + +# Try to import folder_paths if available +try: + import folder_paths +except ImportError: + pass + +class WanI2VGenerator(WanAPIBase): + """Node for image-to-video generation using Wan model""" + + # Define available Wan i2v models + MODEL_OPTIONS = [ + "wan2.2-i2v-flash", # Speed Edition + "wan2.2-i2v-plus" # Professional Edition + ] + + # Define allowed resolutions for Wan i2v models (using uppercase P as required by API) + RESOLUTION_OPTIONS = [ + "480P", + "720P", + "1080P" + ] + + def __init__(self): + super().__init__() + # Use the centralized API endpoint from the base class + # To use Mainland China region, modify API_ENDPOINT_POST_VIDEO in wan_base.py + self.api_url = self.API_ENDPOINT_POST_VIDEO + + @classmethod + def INPUT_TYPES(cls): + # Define output directory options + if COMFYUI_AVAILABLE: + # Use ComfyUI's output directory with browseable option + output_dir_options = { + "default": "./videos", + "tooltip": "Directory where the generated video will be saved. Browse to select a custom directory." + } + else: + # Fallback to string input + output_dir_options = { + "default": "./videos", + "multiline": False + } + + return { + "required": { + "model": (cls.MODEL_OPTIONS, { + "default": "wan2.2-i2v-flash" + }), + "image_url": ("STRING", { + "default": "https://example.com/your_image.png" + }), + "prompt": ("STRING", { + "multiline": True, + "default": "A cat running on the grass" + }) + }, + "optional": { + "negative_prompt": ("STRING", { + "multiline": True, + "default": "" + }), + "resolution": (cls.RESOLUTION_OPTIONS, { + "default": "720P" + }), + "prompt_extend": ("BOOLEAN", { + "default": True + }), + "watermark": ("BOOLEAN", { + "default": False + }), + "seed": ("INT", { + "default": 0, + "min": 0, + "max": 2147483647 + }), + "output_dir": ("STRING", output_dir_options) + } + } + + RETURN_TYPES = ("STRING",) # Returns path to downloaded video file + FUNCTION = "generate" + CATEGORY = "Ru4ls/Wan" + + def generate(self, model, image_url, prompt, negative_prompt="", resolution="720P", + prompt_extend=True, watermark=False, seed=0, output_dir="./videos"): + # Check API key + self.check_api_key() + + # Prepare API payload for image-to-video generation + payload = { + "model": model, + "input": { + "prompt": prompt, + "img_url": image_url + }, + "parameters": { + "resolution": resolution, + "prompt_extend": prompt_extend, + "watermark": watermark + } + } + + # Add optional parameters if they have non-default values + if negative_prompt: + payload["input"]["negative_prompt"] = negative_prompt + if seed > 0: + payload["parameters"]["seed"] = seed + + # Set headers according to DashScope documentation + headers = { + "Authorization": f"Bearer {self.api_key}", + "Content-Type": "application/json", + "X-DashScope-Async": "enable" # Wan requires async processing + } + + try: + # Make API request + print(f"Making API request to {self.api_url}") + response = requests.post(self.api_url, headers=headers, json=payload) + print(f"Response status code: {response.status_code}") + if hasattr(response, 'text'): + print(f"Response text: {response.text[:500]}...") # Print first 500 chars + response.raise_for_status() + + # Parse response to get task_id + result = response.json() + print(f"API response received: {json.dumps(result, indent=2)[:200]}...") # Print first 200 chars + + # Check if this is a task creation response + if "output" in result and "task_id" in result["output"]: + task_id = result["output"]["task_id"] + task_status = result["output"]["task_status"] + print(f"Task created with ID: {task_id}, status: {task_status}") + + # Now we need to poll for the result + task_result = self.poll_task_result(task_id, output_dir) + return (task_result,) # Return path to downloaded video file + else: + raise ValueError(f"Unexpected API response format: {result}") + + except requests.exceptions.RequestException as e: + # More detailed error handling + if hasattr(e, 'response') and e.response is not None: + status_code = e.response.status_code + response_text = e.response.text + print(f"API request failed with status {status_code}: {response_text}") + if status_code == 401: + raise RuntimeError(f"API request failed: 401 Unauthorized. " + f"This usually means your API key is invalid or not properly configured. " + f"Error details: {response_text}") + elif status_code == 403: + raise RuntimeError(f"API request failed: 403 Forbidden. " + f"This usually means your API key is valid but you don't have access to this model. " + f"Error details: {response_text}") + elif status_code == 400: + raise RuntimeError(f"API request failed: 400 Bad Request. " + f"This usually means there's an issue with the request format. " + f"Error details: {response_text}") + else: + raise RuntimeError(f"API request failed: {status_code} {e.response.reason}. Response: {response_text}") + else: + raise RuntimeError(f"API request failed: {str(e)}") + except Exception as e: + raise RuntimeError(f"Failed to process API response: {str(e)}") + + def poll_task_result(self, task_id, output_dir="./videos"): + """Poll for task result until completion and download video""" + import time + + # URL for querying task results + # To use Mainland China region, modify API_ENDPOINT_GET in wan_base.py + query_url = self.API_ENDPOINT_GET.format(task_id=task_id) + + headers = { + "Authorization": f"Bearer {self.api_key}", + "Content-Type": "application/json" + } + + max_attempts = 60 # Maximum polling attempts (may take longer for video) + attempt = 0 + + while attempt < max_attempts: + try: + print(f"Polling task {task_id}, attempt {attempt + 1}/{max_attempts}") + response = requests.get(query_url, headers=headers) + response.raise_for_status() + + result = response.json() + task_status = result["output"]["task_status"] + print(f"Task status: {task_status}") + + if task_status == "SUCCEEDED": + # Task completed successfully + if "video_url" in result["output"]: + video_url = result["output"]["video_url"] + + # Download the video + video_response = requests.get(video_url) + video_response.raise_for_status() + + # Create a unique filename for the video + timestamp = datetime.now().strftime("%Y%m%d_%H%M%S") + video_filename = f"wan_i2v_{timestamp}.mp4" + + # Handle output directory based on ComfyUI availability + if COMFYUI_AVAILABLE and not output_dir.startswith(("./", "/")): + # Use ComfyUI's output directory structure + if output_dir.endswith("/"): + output_dir = output_dir[:-1] + full_output_folder = folder_paths.get_output_directory() + output_path = os.path.join(full_output_folder, output_dir) + else: + # Resolve output directory path (existing logic) + if output_dir.startswith("./"): + # Relative to the node directory + output_path = os.path.join(os.path.dirname(__file__), output_dir[2:]) + else: + output_path = output_dir + + # Create output directory if it doesn't exist + os.makedirs(output_path, exist_ok=True) + + # Save video to file + video_path = os.path.join(output_path, video_filename) + with open(video_path, "wb") as f: + f.write(video_response.content) + + print(f"Video downloaded and saved to: {video_path}") + # Return path relative to ComfyUI output directory if using ComfyUI + if COMFYUI_AVAILABLE and not output_dir.startswith(("./", "/")): + return os.path.join(output_dir, video_filename) if output_dir != "videos/" else video_filename + else: + return video_path # Return full path + else: + raise ValueError(f"Unexpected API response format: {result}") + + elif task_status == "FAILED": + # Task failed + error_code = result["output"].get("code", "Unknown") + error_message = result["output"].get("message", "Unknown error") + raise RuntimeError(f"Task failed with code: {error_code}, message: {error_message}") + + elif task_status in ["PENDING", "RUNNING"]: + # Task still in progress, wait and retry + time.sleep(10) # Wait 10 seconds before retrying (video generation may take longer) + attempt += 1 + continue + + else: + raise ValueError(f"Unexpected task status: {task_status}") + + except requests.exceptions.RequestException as e: + raise RuntimeError(f"Failed to query task status: {str(e)}") + + # If we've reached here, we've exceeded max attempts + raise RuntimeError(f"Task did not complete within the expected time ({max_attempts} attempts)") \ No newline at end of file diff --git a/wan_ii2v.py b/wan_ii2v.py new file mode 100644 index 0000000..46d3b31 --- /dev/null +++ b/wan_ii2v.py @@ -0,0 +1,274 @@ +import os +import json +import requests +from PIL import Image +import numpy as np +import torch +import io +import base64 +from dotenv import load_dotenv +import sys +import pathlib +from datetime import datetime + +# Import the base class and COMFYUI_AVAILABLE flag +from .wan_base import WanAPIBase, COMFYUI_AVAILABLE + +# Try to import folder_paths if available +try: + import folder_paths +except ImportError: + pass + +class WanII2VGenerator(WanAPIBase): + """Node for image-to-video generation using first and last frames with Wan model""" + + # Define available Wan ii2v models + MODEL_OPTIONS = [ + "wan2.1-kf2v-plus" # Professional Edition + ] + + # Define allowed resolutions for Wan ii2v models + RESOLUTION_OPTIONS = [ + "720P" + ] + + def __init__(self): + super().__init__() + # Use the centralized API endpoint from the base class + # To use Mainland China region, modify API_ENDPOINT_POST_II2V in wan_base.py + self.api_url = self.API_ENDPOINT_POST_II2V + + @classmethod + def INPUT_TYPES(cls): + # Define output directory options + if COMFYUI_AVAILABLE: + # Use ComfyUI's output directory with browseable option + output_dir_options = { + "default": "./videos", + "tooltip": "Directory where the generated video will be saved. Browse to select a custom directory." + } + else: + # Fallback to string input + output_dir_options = { + "default": "./videos", + "multiline": False + } + + return { + "required": { + "model": (cls.MODEL_OPTIONS, { + "default": "wan2.1-kf2v-plus" + }), + "first_frame_url": ("STRING", { + "default": "https://example.com/first_frame.png" + }), + "last_frame_url": ("STRING", { + "default": "https://example.com/last_frame.png" + }), + "prompt": ("STRING", { + "multiline": True, + "default": "A black kitten looks up at the sky curiously, the camera gradually rises from eye level, and finally shoots from a top-down angle to capture the kitten's curious eyes." + }) + }, + "optional": { + "negative_prompt": ("STRING", { + "multiline": True, + "default": "" + }), + "resolution": (cls.RESOLUTION_OPTIONS, { + "default": "720P" + }), + "prompt_extend": ("BOOLEAN", { + "default": True + }), + "watermark": ("BOOLEAN", { + "default": False + }), + "seed": ("INT", { + "default": 0, + "min": 0, + "max": 2147483647 + }), + "output_dir": ("STRING", output_dir_options) + } + } + + RETURN_TYPES = ("STRING",) # Returns path to downloaded video file + FUNCTION = "generate" + CATEGORY = "Ru4ls/Wan" + + def generate(self, model, first_frame_url, last_frame_url, prompt, negative_prompt="", + resolution="720P", prompt_extend=True, watermark=False, seed=0, output_dir="./videos"): + # Check API key + self.check_api_key() + + # Prepare API payload for image-to-video generation with first and last frames + payload = { + "model": model, + "input": { + "first_frame_url": first_frame_url, + "last_frame_url": last_frame_url, + "prompt": prompt + }, + "parameters": { + "resolution": resolution, + "prompt_extend": prompt_extend, + "watermark": watermark + } + } + + # Add optional parameters if they have non-default values + if negative_prompt: + payload["input"]["negative_prompt"] = negative_prompt + if seed > 0: + payload["parameters"]["seed"] = seed + + # Set headers according to DashScope documentation + headers = { + "Authorization": f"Bearer {self.api_key}", + "Content-Type": "application/json", + "X-DashScope-Async": "enable" # Wan requires async processing + } + + try: + # Make API request + print(f"Making API request to {self.api_url}") + response = requests.post(self.api_url, headers=headers, json=payload) + print(f"Response status code: {response.status_code}") + if hasattr(response, 'text'): + print(f"Response text: {response.text[:500]}...") # Print first 500 chars + response.raise_for_status() + + # Parse response to get task_id + result = response.json() + print(f"API response received: {json.dumps(result, indent=2)[:200]}...") # Print first 200 chars + + # Check if this is a task creation response + if "output" in result and "task_id" in result["output"]: + task_id = result["output"]["task_id"] + task_status = result["output"]["task_status"] + print(f"Task created with ID: {task_id}, status: {task_status}") + + # Now we need to poll for the result + task_result = self.poll_task_result(task_id, output_dir) + return (task_result,) # Return path to downloaded video file + else: + raise ValueError(f"Unexpected API response format: {result}") + + except requests.exceptions.RequestException as e: + # More detailed error handling + if hasattr(e, 'response') and e.response is not None: + status_code = e.response.status_code + response_text = e.response.text + print(f"API request failed with status {status_code}: {response_text}") + if status_code == 401: + raise RuntimeError(f"API request failed: 401 Unauthorized. " + f"This usually means your API key is invalid or not properly configured. " + f"Error details: {response_text}") + elif status_code == 403: + raise RuntimeError(f"API request failed: 403 Forbidden. " + f"This usually means your API key is valid but you don't have access to this model. " + f"Error details: {response_text}") + elif status_code == 400: + raise RuntimeError(f"API request failed: 400 Bad Request. " + f"This usually means there's an issue with the request format. " + f"Error details: {response_text}") + else: + raise RuntimeError(f"API request failed: {status_code} {e.response.reason}. Response: {response_text}") + else: + raise RuntimeError(f"API request failed: {str(e)}") + except Exception as e: + raise RuntimeError(f"Failed to process API response: {str(e)}") + + def poll_task_result(self, task_id, output_dir="./videos"): + """Poll for task result until completion and download video""" + import time + + # URL for querying task results + # To use Mainland China region, modify API_ENDPOINT_GET in wan_base.py + query_url = self.API_ENDPOINT_GET.format(task_id=task_id) + + headers = { + "Authorization": f"Bearer {self.api_key}", + "Content-Type": "application/json" + } + + max_attempts = 60 # Maximum polling attempts (may take longer for video) + attempt = 0 + + while attempt < max_attempts: + try: + print(f"Polling task {task_id}, attempt {attempt + 1}/{max_attempts}") + response = requests.get(query_url, headers=headers) + response.raise_for_status() + + result = response.json() + task_status = result["output"]["task_status"] + print(f"Task status: {task_status}") + + if task_status == "SUCCEEDED": + # Task completed successfully + if "video_url" in result["output"]: + video_url = result["output"]["video_url"] + + # Download the video + video_response = requests.get(video_url) + video_response.raise_for_status() + + # Create a unique filename for the video + timestamp = datetime.now().strftime("%Y%m%d_%H%M%S") + video_filename = f"wan_ii2v_{timestamp}.mp4" + + # Handle output directory based on ComfyUI availability + if COMFYUI_AVAILABLE and not output_dir.startswith(("./", "/")): + # Use ComfyUI's output directory structure + if output_dir.endswith("/"): + output_dir = output_dir[:-1] + full_output_folder = folder_paths.get_output_directory() + output_path = os.path.join(full_output_folder, output_dir) + else: + # Resolve output directory path (existing logic) + if output_dir.startswith("./"): + # Relative to the node directory + output_path = os.path.join(os.path.dirname(__file__), output_dir[2:]) + else: + output_path = output_dir + + # Create output directory if it doesn't exist + os.makedirs(output_path, exist_ok=True) + + # Save video to file + video_path = os.path.join(output_path, video_filename) + with open(video_path, "wb") as f: + f.write(video_response.content) + + print(f"Video downloaded and saved to: {video_path}") + # Return path relative to ComfyUI output directory if using ComfyUI + if COMFYUI_AVAILABLE and not output_dir.startswith(("./", "/")): + return os.path.join(output_dir, video_filename) if output_dir != "videos/" else video_filename + else: + return video_path # Return full path + else: + raise ValueError(f"Unexpected API response format: {result}") + + elif task_status == "FAILED": + # Task failed + error_code = result["output"].get("code", "Unknown") + error_message = result["output"].get("message", "Unknown error") + raise RuntimeError(f"Task failed with code: {error_code}, message: {error_message}") + + elif task_status in ["PENDING", "RUNNING"]: + # Task still in progress, wait and retry + time.sleep(10) # Wait 10 seconds before retrying (video generation may take longer) + attempt += 1 + continue + + else: + raise ValueError(f"Unexpected task status: {task_status}") + + except requests.exceptions.RequestException as e: + raise RuntimeError(f"Failed to query task status: {str(e)}") + + # If we've reached here, we've exceeded max attempts + raise RuntimeError(f"Task did not complete within the expected time ({max_attempts} attempts)") \ No newline at end of file diff --git a/wan_nodes.py b/wan_nodes.py index 5f7f103..4dee274 100644 --- a/wan_nodes.py +++ b/wan_nodes.py @@ -1,563 +1,17 @@ -import os -import json -import requests -from PIL import Image -import numpy as np -import torch -import io -import base64 -from dotenv import load_dotenv -import sys -import pathlib +""" +Backward compatibility module - imports all Wan nodes from their separate modules. +""" -# Import ComfyUI's folder_paths for directory browsing -try: - import folder_paths - COMFYUI_AVAILABLE = True -except ImportError: - COMFYUI_AVAILABLE = False - print("folder_paths not available, using default directory handling") +from .wan_t2i import WanT2IGenerator +from .wan_i2v import WanI2VGenerator +from .wan_t2v import WanT2VGenerator +from .wan_ii2v import WanII2VGenerator +from .wan_vace_image_reference import WanVACEImageReference +from .wan_vace_video_repainting import WanVACEVideoRepainting +from .wan_vace_video_edit import WanVACEVideoEdit +from .wan_vace_video_extension import WanVACEVideoExtension +from .wan_vace_video_outpainting import WanVACEVideoOutpainting -# Load environment variables from .env file -# Try to load .env file from the current directory first -env_path = pathlib.Path(__file__).parent / '.env' -if env_path.exists(): - load_dotenv(dotenv_path=env_path) -else: - # Fallback to default behavior - load_dotenv() - -# Debug: Print environment variable status -api_key = os.getenv('DASHSCOPE_API_KEY') -if api_key: - # Strip any extra quotes or whitespace - api_key = api_key.strip().strip('"\'') - print(f"API Key loaded: {api_key[:8]}...{api_key[-4:]}") # Print partial key for security -else: - print("API Key not found in environment variables") - - -class WanAPIBase: - """Base class for Wan API interactions""" - - def __init__(self): - self.api_key = os.getenv('DASHSCOPE_API_KEY') - # Strip any extra quotes or whitespace - if self.api_key: - self.api_key = self.api_key.strip().strip('"\'') - print(f"Initialized WanAPIBase with API key: {self.api_key[:8] if self.api_key else 'None'}...{self.api_key[-4:] if self.api_key else ''}") - - def check_api_key(self): - """Check if API key is set in environment variables""" - if not self.api_key: - raise ValueError("DASHSCOPE_API_KEY environment variable not set. " - "Please set it before using this node.") - return self.api_key - - def prepare_images(self, images): - """Convert images to base64 strings for API submission""" - image_data = [] - for i, image in enumerate(images, 1): - if image is not None: - # Convert tensor to PIL Image - if isinstance(image, torch.Tensor): - # Convert tensor to numpy array - image_np = image.cpu().numpy() - # If the tensor is in [0, 1] range, convert to [0, 255] - if image_np.max() <= 1.0: - image_np = (image_np * 255).astype(np.uint8) - # If tensor has shape [H, W, C], convert to PIL - pil_image = Image.fromarray(image_np.squeeze()) - else: - pil_image = image - - # Convert PIL image to base64 - buffer = io.BytesIO() - pil_image.save(buffer, format="PNG") - img_str = base64.b64encode(buffer.getvalue()).decode() - image_data.append({ - "id": str(i), - "data": img_str - }) - return image_data - - -class WanT2IGenerator(WanAPIBase): - """Node for text-to-image generation using Wan model""" - - # Define available Wan models - MODEL_OPTIONS = [ - "wan2.2-t2i-flash", # Speed Edition - "wan2.2-t2i-plus" # Professional Edition - ] - - # Define allowed sizes for Wan models with descriptive names - # Based on the documentation, Wan supports sizes from 512 to 1440 pixels - SIZE_OPTIONS = [ - "1024*1024", # 1:1 square (default) - "1152*896", # 9:7 landscape - "896*1152", # 7:9 portrait - "1280*720", # 16:9 landscape - "720*1280", # 9:16 portrait - "1440*512", # Wide landscape - "512*1440" # Tall portrait - ] - - def __init__(self): - super().__init__() - self.api_url = "https://dashscope-intl.aliyuncs.com/api/v1/services/aigc/text2image/image-synthesis" - self.model = "wan2.2-t2i-flash" # Using Wan Speed Edition as default - - @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "model": (cls.MODEL_OPTIONS, { - "default": "wan2.2-t2i-flash" - }), - "prompt": ("STRING", { - "multiline": True, - "default": "Generate an image of a cat" - }), - "size": (cls.SIZE_OPTIONS, { - "default": "1024*1024" - }) - }, - "optional": { - "negative_prompt": ("STRING", { - "multiline": True, - "default": "" - }), - "prompt_extend": ("BOOLEAN", { - "default": True - }), - "watermark": ("BOOLEAN", { - "default": False - }), - "seed": ("INT", { - "default": 0, - "min": 0, - "max": 2147483647 - }) - } - } - - RETURN_TYPES = ("IMAGE",) - FUNCTION = "generate" - CATEGORY = "Ru4ls/Wan" - - def generate(self, model, prompt, size, negative_prompt="", prompt_extend=True, watermark=False, seed=0): - # Check API key - self.check_api_key() - - # Set the selected model - self.model = model - - # Debug: Print API key status - print(f"Using API key: {self.api_key[:8]}...{self.api_key[-4:] if self.api_key else 'None'}") - print(f"Selected model: {self.model}") - print(f"Using API endpoint: {self.api_url}") - - # Prepare API payload for text-to-image generation - using the Wan format - payload = { - "model": self.model, - "input": { - "prompt": prompt - }, - "parameters": { - "size": size, - "prompt_extend": prompt_extend, - "watermark": watermark, - "n": 1 # Generate only one image - } - } - - # Add optional parameters if they have non-default values - if negative_prompt: - payload["input"]["negative_prompt"] = negative_prompt - if seed > 0: - payload["parameters"]["seed"] = seed - - # Set headers according to DashScope documentation - headers = { - "Authorization": f"Bearer {self.api_key}", - "Content-Type": "application/json", - "X-DashScope-Async": "enable" # Wan requires async processing - } - - # Debug: Print request details - print(f"Request headers: {{'Authorization': 'Bearer {self.api_key[:8]}...', 'Content-Type': 'application/json', 'X-DashScope-Async': 'enable'}}") - print(f"Request payload model: {payload['model']}") - print(f"Request payload prompt: {payload['input']['prompt'][:100]}...") - print(f"Request payload size: {payload['parameters']['size']}") - print(f"Request payload prompt_extend: {payload['parameters']['prompt_extend']}") - print(f"Request payload watermark: {payload['parameters']['watermark']}") - - try: - # Make API request - print(f"Making API request to {self.api_url}") - response = requests.post(self.api_url, headers=headers, json=payload) - print(f"Response status code: {response.status_code}") - if hasattr(response, 'text'): - print(f"Response text: {response.text[:500]}...") # Print first 500 chars - response.raise_for_status() - - # Parse response to get task_id - result = response.json() - print(f"API response received: {json.dumps(result, indent=2)[:200]}...") # Print first 200 chars - - # Check if this is a task creation response - if "output" in result and "task_id" in result["output"]: - task_id = result["output"]["task_id"] - task_status = result["output"]["task_status"] - print(f"Task created with ID: {task_id}, status: {task_status}") - - # Now we need to poll for the result - task_result = self.poll_task_result(task_id) - return task_result - else: - raise ValueError(f"Unexpected API response format: {result}") - - except requests.exceptions.RequestException as e: - # More detailed error handling - if hasattr(e, 'response') and e.response is not None: - status_code = e.response.status_code - response_text = e.response.text - print(f"API request failed with status {status_code}: {response_text}") - if status_code == 401: - raise RuntimeError(f"API request failed: 401 Unauthorized. " - f"This usually means your API key is invalid or not properly configured. " - f"Error details: {response_text}") - elif status_code == 403: - raise RuntimeError(f"API request failed: 403 Forbidden. " - f"This usually means your API key is valid but you don't have access to this model. " - f"Error details: {response_text}") - elif status_code == 400: - raise RuntimeError(f"API request failed: 400 Bad Request. " - f"This usually means there's an issue with the request format. " - f"Error details: {response_text}") - else: - raise RuntimeError(f"API request failed: {status_code} {e.response.reason}. Response: {response_text}") - else: - raise RuntimeError(f"API request failed: {str(e)}") - except Exception as e: - raise RuntimeError(f"Failed to process API response: {str(e)}") - - def poll_task_result(self, task_id): - """Poll for task result until completion""" - import time - - # URL for querying task results - query_url = f"https://dashscope-intl.aliyuncs.com/api/v1/tasks/{task_id}" - - headers = { - "Authorization": f"Bearer {self.api_key}", - "Content-Type": "application/json" - } - - max_attempts = 30 # Maximum polling attempts - attempt = 0 - - while attempt < max_attempts: - try: - print(f"Polling task {task_id}, attempt {attempt + 1}/{max_attempts}") - response = requests.get(query_url, headers=headers) - response.raise_for_status() - - result = response.json() - task_status = result["output"]["task_status"] - print(f"Task status: {task_status}") - - if task_status == "SUCCEEDED": - # Task completed successfully - results = result["output"]["results"] - if len(results) > 0 and "url" in results[0]: - image_url = results[0]["url"] - # Download the generated image - image_response = requests.get(image_url) - image_response.raise_for_status() - - # Convert to tensor - image = Image.open(io.BytesIO(image_response.content)) - image_tensor = torch.from_numpy(np.array(image).astype(np.float32) / 255.0) - image_tensor = image_tensor.unsqueeze(0) # Add batch dimension - - return (image_tensor,) - else: - raise ValueError(f"Unexpected API response format: {result}") - - elif task_status == "FAILED": - # Task failed - error_code = result["output"].get("code", "Unknown") - error_message = result["output"].get("message", "Unknown error") - raise RuntimeError(f"Task failed with code: {error_code}, message: {error_message}") - - elif task_status in ["PENDING", "RUNNING"]: - # Task still in progress, wait and retry - time.sleep(5) # Wait 5 seconds before retrying - attempt += 1 - continue - - else: - raise ValueError(f"Unexpected task status: {task_status}") - - except requests.exceptions.RequestException as e: - raise RuntimeError(f"Failed to query task status: {str(e)}") - - # If we've reached here, we've exceeded max attempts - raise RuntimeError(f"Task did not complete within the expected time ({max_attempts} attempts)") - - -class WanI2VGenerator(WanAPIBase): - """Node for image-to-video generation using Wan model""" - - # Define available Wan i2v models - MODEL_OPTIONS = [ - "wan2.2-i2v-flash", # Speed Edition - "wan2.2-i2v-plus" # Professional Edition - ] - - # Define allowed resolutions for Wan i2v models (using uppercase P as required by API) - RESOLUTION_OPTIONS = [ - "480P", - "720P", - "1080P" - ] - - def __init__(self): - super().__init__() - self.api_url = "https://dashscope-intl.aliyuncs.com/api/v1/services/aigc/video-generation/video-synthesis" - - @classmethod - def INPUT_TYPES(cls): - # Define output directory options - if COMFYUI_AVAILABLE: - # Use ComfyUI's output directory with browseable option - output_dir_options = { - "default": "videos/", - "tooltip": "Directory where the generated video will be saved. Browse to select a custom directory." - } - else: - # Fallback to string input - output_dir_options = { - "default": "./videos", - "multiline": False - } - - return { - "required": { - "model": (cls.MODEL_OPTIONS, { - "default": "wan2.2-i2v-flash" - }), - "image_url": ("STRING", { - "default": "https://example.com/your_image.png" - }), - "prompt": ("STRING", { - "multiline": True, - "default": "A cat running on the grass" - }) - }, - "optional": { - "negative_prompt": ("STRING", { - "multiline": True, - "default": "" - }), - "resolution": (cls.RESOLUTION_OPTIONS, { - "default": "720P" - }), - "prompt_extend": ("BOOLEAN", { - "default": True - }), - "watermark": ("BOOLEAN", { - "default": False - }), - "seed": ("INT", { - "default": 0, - "min": 0, - "max": 2147483647 - }), - "output_dir": ("STRING", output_dir_options) - } - } - - RETURN_TYPES = ("STRING",) # Returns path to downloaded video file - FUNCTION = "generate" - CATEGORY = "Ru4ls/Wan" - - def generate(self, model, image_url, prompt, negative_prompt="", resolution="720p", - prompt_extend=True, watermark=False, seed=0, output_dir="videos/"): - # Check API key - self.check_api_key() - - # Prepare API payload for image-to-video generation - payload = { - "model": model, - "input": { - "prompt": prompt, - "img_url": image_url - }, - "parameters": { - "resolution": resolution, - "prompt_extend": prompt_extend, - "watermark": watermark - } - } - - # Add optional parameters if they have non-default values - if negative_prompt: - payload["input"]["negative_prompt"] = negative_prompt - if seed > 0: - payload["parameters"]["seed"] = seed - - # Set headers according to DashScope documentation - headers = { - "Authorization": f"Bearer {self.api_key}", - "Content-Type": "application/json", - "X-DashScope-Async": "enable" # Wan requires async processing - } - - try: - # Make API request - print(f"Making API request to {self.api_url}") - response = requests.post(self.api_url, headers=headers, json=payload) - print(f"Response status code: {response.status_code}") - if hasattr(response, 'text'): - print(f"Response text: {response.text[:500]}...") # Print first 500 chars - response.raise_for_status() - - # Parse response to get task_id - result = response.json() - print(f"API response received: {json.dumps(result, indent=2)[:200]}...") # Print first 200 chars - - # Check if this is a task creation response - if "output" in result and "task_id" in result["output"]: - task_id = result["output"]["task_id"] - task_status = result["output"]["task_status"] - print(f"Task created with ID: {task_id}, status: {task_status}") - - # Now we need to poll for the result - task_result = self.poll_task_result(task_id, output_dir) - return (task_result,) # Return path to downloaded video file - else: - raise ValueError(f"Unexpected API response format: {result}") - - except requests.exceptions.RequestException as e: - # More detailed error handling - if hasattr(e, 'response') and e.response is not None: - status_code = e.response.status_code - response_text = e.response.text - print(f"API request failed with status {status_code}: {response_text}") - if status_code == 401: - raise RuntimeError(f"API request failed: 401 Unauthorized. " - f"This usually means your API key is invalid or not properly configured. " - f"Error details: {response_text}") - elif status_code == 403: - raise RuntimeError(f"API request failed: 403 Forbidden. " - f"This usually means your API key is valid but you don't have access to this model. " - f"Error details: {response_text}") - elif status_code == 400: - raise RuntimeError(f"API request failed: 400 Bad Request. " - f"This usually means there's an issue with the request format. " - f"Error details: {response_text}") - else: - raise RuntimeError(f"API request failed: {status_code} {e.response.reason}. Response: {response_text}") - else: - raise RuntimeError(f"API request failed: {str(e)}") - except Exception as e: - raise RuntimeError(f"Failed to process API response: {str(e)}") - - def poll_task_result(self, task_id, output_dir="videos/"): - """Poll for task result until completion and download video""" - import time - import os - from datetime import datetime - - # URL for querying task results - query_url = f"https://dashscope-intl.aliyuncs.com/api/v1/tasks/{task_id}" - - headers = { - "Authorization": f"Bearer {self.api_key}", - "Content-Type": "application/json" - } - - max_attempts = 60 # Maximum polling attempts (may take longer for video) - attempt = 0 - - while attempt < max_attempts: - try: - print(f"Polling task {task_id}, attempt {attempt + 1}/{max_attempts}") - response = requests.get(query_url, headers=headers) - response.raise_for_status() - - result = response.json() - task_status = result["output"]["task_status"] - print(f"Task status: {task_status}") - - if task_status == "SUCCEEDED": - # Task completed successfully - if "video_url" in result["output"]: - video_url = result["output"]["video_url"] - - # Download the video - video_response = requests.get(video_url) - video_response.raise_for_status() - - # Create a unique filename for the video - timestamp = datetime.now().strftime("%Y%m%d_%H%M%S") - video_filename = f"wan_i2v_{timestamp}.mp4" - - # Handle output directory based on ComfyUI availability - if COMFYUI_AVAILABLE and not output_dir.startswith(("./", "/")): - # Use ComfyUI's output directory structure - if output_dir.endswith("/"): - output_dir = output_dir[:-1] - full_output_folder = folder_paths.get_output_directory() - output_path = os.path.join(full_output_folder, output_dir) - else: - # Resolve output directory path (existing logic) - if output_dir.startswith("./"): - # Relative to the node directory - output_path = os.path.join(os.path.dirname(__file__), output_dir[2:]) - else: - output_path = output_dir - - # Create output directory if it doesn't exist - os.makedirs(output_path, exist_ok=True) - - # Save video to file - video_path = os.path.join(output_path, video_filename) - with open(video_path, "wb") as f: - f.write(video_response.content) - - print(f"Video downloaded and saved to: {video_path}") - # Return path relative to ComfyUI output directory if using ComfyUI - if COMFYUI_AVAILABLE and not output_dir.startswith(("./", "/")): - return os.path.join(output_dir, video_filename) if output_dir != "videos/" else video_filename - else: - return video_path # Return full path - else: - raise ValueError(f"Unexpected API response format: {result}") - - elif task_status == "FAILED": - # Task failed - error_code = result["output"].get("code", "Unknown") - error_message = result["output"].get("message", "Unknown error") - raise RuntimeError(f"Task failed with code: {error_code}, message: {error_message}") - - elif task_status in ["PENDING", "RUNNING"]: - # Task still in progress, wait and retry - time.sleep(10) # Wait 10 seconds before retrying (video generation may take longer) - attempt += 1 - continue - - else: - raise ValueError(f"Unexpected task status: {task_status}") - - except requests.exceptions.RequestException as e: - raise RuntimeError(f"Failed to query task status: {str(e)}") - - # If we've reached here, we've exceeded max attempts - raise RuntimeError(f"Task did not complete within the expected time ({max_attempts} attempts)") - - -# Node class mappings are in __init__.py \ No newline at end of file +__all__ = ['WanT2IGenerator', 'WanI2VGenerator', 'WanT2VGenerator', 'WanII2VGenerator', + 'WanVACEImageReference', 'WanVACEVideoRepainting', 'WanVACEVideoEdit', + 'WanVACEVideoExtension', 'WanVACEVideoOutpainting'] \ No newline at end of file diff --git a/wan_t2i.py b/wan_t2i.py new file mode 100644 index 0000000..67295bb --- /dev/null +++ b/wan_t2i.py @@ -0,0 +1,248 @@ +import os +import json +import requests +from PIL import Image +import numpy as np +import torch +import io +import base64 +from dotenv import load_dotenv +import sys +import pathlib + +# Import the base class +from .wan_base import WanAPIBase, COMFYUI_AVAILABLE + +# Try to import folder_paths if available +try: + import folder_paths +except ImportError: + pass + +class WanT2IGenerator(WanAPIBase): + """Node for text-to-image generation using Wan model""" + + # Define available Wan models + MODEL_OPTIONS = [ + "wan2.2-t2i-flash", # Speed Edition + "wan2.2-t2i-plus" # Professional Edition + ] + + # Define allowed sizes for Wan models with descriptive names + # Based on the documentation, Wan supports sizes from 512 to 1440 pixels + SIZE_OPTIONS = [ + "1024*1024", # 1:1 square (default) + "1152*896", # 9:7 landscape + "896*1152", # 7:9 portrait + "1280*720", # 16:9 landscape + "720*1280", # 9:16 portrait + "1440*512", # Wide landscape + "512*1440" # Tall portrait + ] + + def __init__(self): + super().__init__() + # Use the centralized API endpoint from the base class + # To use Mainland China region, modify API_ENDPOINT_POST_T2I in wan_base.py + self.api_url = self.API_ENDPOINT_POST_T2I + self.model = "wan2.2-t2i-flash" # Using Wan Speed Edition as default + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "model": (cls.MODEL_OPTIONS, { + "default": "wan2.2-t2i-flash" + }), + "prompt": ("STRING", { + "multiline": True, + "default": "Generate an image of a cat" + }), + "size": (cls.SIZE_OPTIONS, { + "default": "1024*1024" + }) + }, + "optional": { + "negative_prompt": ("STRING", { + "multiline": True, + "default": "" + }), + "prompt_extend": ("BOOLEAN", { + "default": True + }), + "watermark": ("BOOLEAN", { + "default": False + }), + "seed": ("INT", { + "default": 0, + "min": 0, + "max": 2147483647 + }) + } + } + + RETURN_TYPES = ("IMAGE",) + FUNCTION = "generate" + CATEGORY = "Ru4ls/Wan" + + def generate(self, model, prompt, size, negative_prompt="", prompt_extend=True, watermark=False, seed=0): + # Check API key + self.check_api_key() + + # Set the selected model + self.model = model + + # Debug: Print API key status + print(f"Using API key: {self.api_key[:8]}...{self.api_key[-4:] if self.api_key else 'None'}") + print(f"Selected model: {self.model}") + print(f"Using API endpoint: {self.api_url}") + + # Prepare API payload for text-to-image generation - using the Wan format + payload = { + "model": self.model, + "input": { + "prompt": prompt + }, + "parameters": { + "size": size, + "prompt_extend": prompt_extend, + "watermark": watermark, + "n": 1 # Generate only one image + } + } + + # Add optional parameters if they have non-default values + if negative_prompt: + payload["input"]["negative_prompt"] = negative_prompt + if seed > 0: + payload["parameters"]["seed"] = seed + + # Set headers according to DashScope documentation + headers = { + "Authorization": f"Bearer {self.api_key}", + "Content-Type": "application/json", + "X-DashScope-Async": "enable" # Wan requires async processing + } + + # Debug: Print request details + print(f"Request headers: {{'Authorization': 'Bearer {self.api_key[:8]}...', 'Content-Type': 'application/json', 'X-DashScope-Async': 'enable'}}") + print(f"Request payload model: {payload['model']}") + print(f"Request payload prompt: {payload['input']['prompt'][:100]}...") + print(f"Request payload size: {payload['parameters']['size']}") + print(f"Request payload prompt_extend: {payload['parameters']['prompt_extend']}") + print(f"Request payload watermark: {payload['parameters']['watermark']}") + + try: + # Make API request + print(f"Making API request to {self.api_url}") + response = requests.post(self.api_url, headers=headers, json=payload) + print(f"Response status code: {response.status_code}") + if hasattr(response, 'text'): + print(f"Response text: {response.text[:500]}...") # Print first 500 chars + response.raise_for_status() + + # Parse response to get task_id + result = response.json() + print(f"API response received: {json.dumps(result, indent=2)[:200]}...") # Print first 200 chars + + # Check if this is a task creation response + if "output" in result and "task_id" in result["output"]: + task_id = result["output"]["task_id"] + task_status = result["output"]["task_status"] + print(f"Task created with ID: {task_id}, status: {task_status}") + + # Now we need to poll for the result + task_result = self.poll_task_result(task_id) + return task_result + else: + raise ValueError(f"Unexpected API response format: {result}") + + except requests.exceptions.RequestException as e: + # More detailed error handling + if hasattr(e, 'response') and e.response is not None: + status_code = e.response.status_code + response_text = e.response.text + print(f"API request failed with status {status_code}: {response_text}") + if status_code == 401: + raise RuntimeError(f"API request failed: 401 Unauthorized. " + f"This usually means your API key is invalid or not properly configured. " + f"Error details: {response_text}") + elif status_code == 403: + raise RuntimeError(f"API request failed: 403 Forbidden. " + f"This usually means your API key is valid but you don't have access to this model. " + f"Error details: {response_text}") + elif status_code == 400: + raise RuntimeError(f"API request failed: 400 Bad Request. " + f"This usually means there's an issue with the request format. " + f"Error details: {response_text}") + else: + raise RuntimeError(f"API request failed: {status_code} {e.response.reason}. Response: {response_text}") + else: + raise RuntimeError(f"API request failed: {str(e)}") + except Exception as e: + raise RuntimeError(f"Failed to process API response: {str(e)}") + + def poll_task_result(self, task_id): + """Poll for task result until completion""" + import time + + # URL for querying task results + # To use Mainland China region, modify API_ENDPOINT_GET in wan_base.py + query_url = self.API_ENDPOINT_GET.format(task_id=task_id) + + headers = { + "Authorization": f"Bearer {self.api_key}", + "Content-Type": "application/json" + } + + max_attempts = 30 # Maximum polling attempts + attempt = 0 + + while attempt < max_attempts: + try: + print(f"Polling task {task_id}, attempt {attempt + 1}/{max_attempts}") + response = requests.get(query_url, headers=headers) + response.raise_for_status() + + result = response.json() + task_status = result["output"]["task_status"] + print(f"Task status: {task_status}") + + if task_status == "SUCCEEDED": + # Task completed successfully + results = result["output"]["results"] + if len(results) > 0 and "url" in results[0]: + image_url = results[0]["url"] + # Download the generated image + image_response = requests.get(image_url) + image_response.raise_for_status() + + # Convert to tensor + image = Image.open(io.BytesIO(image_response.content)) + image_tensor = torch.from_numpy(np.array(image).astype(np.float32) / 255.0) + image_tensor = image_tensor.unsqueeze(0) # Add batch dimension + + return (image_tensor,) + else: + raise ValueError(f"Unexpected API response format: {result}") + + elif task_status == "FAILED": + # Task failed + error_code = result["output"].get("code", "Unknown") + error_message = result["output"].get("message", "Unknown error") + raise RuntimeError(f"Task failed with code: {error_code}, message: {error_message}") + + elif task_status in ["PENDING", "RUNNING"]: + # Task still in progress, wait and retry + time.sleep(5) # Wait 5 seconds before retrying + attempt += 1 + continue + + else: + raise ValueError(f"Unexpected task status: {task_status}") + + except requests.exceptions.RequestException as e: + raise RuntimeError(f"Failed to query task status: {str(e)}") + + # If we've reached here, we've exceeded max attempts + raise RuntimeError(f"Task did not complete within the expected time ({max_attempts} attempts)") \ No newline at end of file diff --git a/wan_t2v.py b/wan_t2v.py new file mode 100644 index 0000000..2866cd0 --- /dev/null +++ b/wan_t2v.py @@ -0,0 +1,287 @@ +import os +import json +import requests +from PIL import Image +import numpy as np +import torch +import io +import base64 +from dotenv import load_dotenv +import sys +import pathlib +from datetime import datetime + +# Import the base class and COMFYUI_AVAILABLE flag +from .wan_base import WanAPIBase, COMFYUI_AVAILABLE + +# Try to import folder_paths if available +try: + import folder_paths +except ImportError: + pass + +class WanT2VGenerator(WanAPIBase): + """Node for text-to-video generation using Wan model""" + + # Define available Wan t2v models + MODEL_OPTIONS = [ + "wan2.2-t2v-plus" # Professional Edition + ] + + # Define allowed resolutions for Wan t2v models + RESOLUTION_OPTIONS = [ + "480P", + "1080P" + ] + + def __init__(self): + super().__init__() + # Use the centralized API endpoint from the base class + # To use Mainland China region, modify API_ENDPOINT_POST_VIDEO in wan_base.py + self.api_url = self.API_ENDPOINT_POST_VIDEO + + @classmethod + def INPUT_TYPES(cls): + # Define output directory options + if COMFYUI_AVAILABLE: + # Use ComfyUI's output directory with browseable option + output_dir_options = { + "default": "./videos", + "tooltip": "Directory where the generated video will be saved. Browse to select a custom directory." + } + else: + # Fallback to string input + output_dir_options = { + "default": "./videos", + "multiline": False + } + + return { + "required": { + "model": (cls.MODEL_OPTIONS, { + "default": "wan2.2-t2v-plus" + }), + "prompt": ("STRING", { + "multiline": True, + "default": "A kitten running in the moonlight" + }) + }, + "optional": { + "negative_prompt": ("STRING", { + "multiline": True, + "default": "" + }), + "resolution": (cls.RESOLUTION_OPTIONS, { + "default": "1080P" + }), + "prompt_extend": ("BOOLEAN", { + "default": True + }), + "watermark": ("BOOLEAN", { + "default": False + }), + "seed": ("INT", { + "default": 0, + "min": 0, + "max": 2147483647 + }), + "output_dir": ("STRING", output_dir_options) + } + } + + RETURN_TYPES = ("STRING",) # Returns path to downloaded video file + FUNCTION = "generate" + CATEGORY = "Ru4ls/Wan" + + def generate(self, model, prompt, negative_prompt="", resolution="1080P", + prompt_extend=True, watermark=False, seed=0, output_dir="./videos"): + # Check API key + self.check_api_key() + + # Prepare API payload for text-to-video generation + payload = { + "model": model, + "input": { + "prompt": prompt + }, + "parameters": { + "prompt_extend": prompt_extend, + "watermark": watermark + } + } + + # Add resolution parameter based on selection + # Convert resolution tier to specific size values + resolution_sizes = { + "480P": { + "16:9": "832*480", + "9:16": "480*832", + "1:1": "624*624" + }, + "1080P": { + "16:9": "1920*1080", + "9:16": "1080*1920", + "1:1": "1440*1440", + "4:3": "1632*1248", + "3:4": "1248*1632" + } + } + + # Default to 16:9 aspect ratio + if resolution in resolution_sizes: + payload["parameters"]["size"] = resolution_sizes[resolution]["16:9"] + + # Add optional parameters if they have non-default values + if negative_prompt: + payload["input"]["negative_prompt"] = negative_prompt + if seed > 0: + payload["parameters"]["seed"] = seed + + # Set headers according to DashScope documentation + headers = { + "Authorization": f"Bearer {self.api_key}", + "Content-Type": "application/json", + "X-DashScope-Async": "enable" # Wan requires async processing + } + + try: + # Make API request + print(f"Making API request to {self.api_url}") + response = requests.post(self.api_url, headers=headers, json=payload) + print(f"Response status code: {response.status_code}") + if hasattr(response, 'text'): + print(f"Response text: {response.text[:500]}...") # Print first 500 chars + response.raise_for_status() + + # Parse response to get task_id + result = response.json() + print(f"API response received: {json.dumps(result, indent=2)[:200]}...") # Print first 200 chars + + # Check if this is a task creation response + if "output" in result and "task_id" in result["output"]: + task_id = result["output"]["task_id"] + task_status = result["output"]["task_status"] + print(f"Task created with ID: {task_id}, status: {task_status}") + + # Now we need to poll for the result + task_result = self.poll_task_result(task_id, output_dir) + return (task_result,) # Return path to downloaded video file + else: + raise ValueError(f"Unexpected API response format: {result}") + + except requests.exceptions.RequestException as e: + # More detailed error handling + if hasattr(e, 'response') and e.response is not None: + status_code = e.response.status_code + response_text = e.response.text + print(f"API request failed with status {status_code}: {response_text}") + if status_code == 401: + raise RuntimeError(f"API request failed: 401 Unauthorized. " + f"This usually means your API key is invalid or not properly configured. " + f"Error details: {response_text}") + elif status_code == 403: + raise RuntimeError(f"API request failed: 403 Forbidden. " + f"This usually means your API key is valid but you don't have access to this model. " + f"Error details: {response_text}") + elif status_code == 400: + raise RuntimeError(f"API request failed: 400 Bad Request. " + f"This usually means there's an issue with the request format. " + f"Error details: {response_text}") + else: + raise RuntimeError(f"API request failed: {status_code} {e.response.reason}. Response: {response_text}") + else: + raise RuntimeError(f"API request failed: {str(e)}") + except Exception as e: + raise RuntimeError(f"Failed to process API response: {str(e)}") + + def poll_task_result(self, task_id, output_dir="./videos"): + """Poll for task result until completion and download video""" + import time + + # URL for querying task results + # To use Mainland China region, modify API_ENDPOINT_GET in wan_base.py + query_url = self.API_ENDPOINT_GET.format(task_id=task_id) + + headers = { + "Authorization": f"Bearer {self.api_key}", + "Content-Type": "application/json" + } + + max_attempts = 60 # Maximum polling attempts (may take longer for video) + attempt = 0 + + while attempt < max_attempts: + try: + print(f"Polling task {task_id}, attempt {attempt + 1}/{max_attempts}") + response = requests.get(query_url, headers=headers) + response.raise_for_status() + + result = response.json() + task_status = result["output"]["task_status"] + print(f"Task status: {task_status}") + + if task_status == "SUCCEEDED": + # Task completed successfully + if "video_url" in result["output"]: + video_url = result["output"]["video_url"] + + # Download the video + video_response = requests.get(video_url) + video_response.raise_for_status() + + # Create a unique filename for the video + timestamp = datetime.now().strftime("%Y%m%d_%H%M%S") + video_filename = f"wan_t2v_{timestamp}.mp4" + + # Handle output directory based on ComfyUI availability + if COMFYUI_AVAILABLE and not output_dir.startswith(("./", "/")): + # Use ComfyUI's output directory structure + if output_dir.endswith("/"): + output_dir = output_dir[:-1] + full_output_folder = folder_paths.get_output_directory() + output_path = os.path.join(full_output_folder, output_dir) + else: + # Resolve output directory path (existing logic) + if output_dir.startswith("./"): + # Relative to the node directory + output_path = os.path.join(os.path.dirname(__file__), output_dir[2:]) + else: + output_path = output_dir + + # Create output directory if it doesn't exist + os.makedirs(output_path, exist_ok=True) + + # Save video to file + video_path = os.path.join(output_path, video_filename) + with open(video_path, "wb") as f: + f.write(video_response.content) + + print(f"Video downloaded and saved to: {video_path}") + # Return path relative to ComfyUI output directory if using ComfyUI + if COMFYUI_AVAILABLE and not output_dir.startswith(("./", "/")): + return os.path.join(output_dir, video_filename) if output_dir != "videos/" else video_filename + else: + return video_path # Return full path + else: + raise ValueError(f"Unexpected API response format: {result}") + + elif task_status == "FAILED": + # Task failed + error_code = result["output"].get("code", "Unknown") + error_message = result["output"].get("message", "Unknown error") + raise RuntimeError(f"Task failed with code: {error_code}, message: {error_message}") + + elif task_status in ["PENDING", "RUNNING"]: + # Task still in progress, wait and retry + time.sleep(10) # Wait 10 seconds before retrying (video generation may take longer) + attempt += 1 + continue + + else: + raise ValueError(f"Unexpected task status: {task_status}") + + except requests.exceptions.RequestException as e: + raise RuntimeError(f"Failed to query task status: {str(e)}") + + # If we've reached here, we've exceeded max attempts + raise RuntimeError(f"Task did not complete within the expected time ({max_attempts} attempts)") \ No newline at end of file diff --git a/wan_vace_image_reference.py b/wan_vace_image_reference.py new file mode 100644 index 0000000..4dcc576 --- /dev/null +++ b/wan_vace_image_reference.py @@ -0,0 +1,310 @@ +""" +Wan VACE Multi-Image Reference Node for ComfyUI +""" + +import os +import json +import requests +from PIL import Image +import numpy as np +import torch +import io +import base64 +from dotenv import load_dotenv +import sys +import pathlib +from datetime import datetime + +# Import the base class and COMFYUI_AVAILABLE flag +from .wan_base import WanAPIBase, COMFYUI_AVAILABLE + +# Try to import folder_paths if available +try: + import folder_paths +except ImportError: + pass + +class WanVACEImageReference(WanAPIBase): + """Node for multi-image reference using Wan VACE model""" + + # Define available Wan VACE models + MODEL_OPTIONS = [ + "wan2.1-vace-plus" # Professional Edition + ] + + # Define video resolutions + RESOLUTION_OPTIONS = [ + "1280*720", # 16:9 aspect ratio (default) + "720*1280", # 9:16 aspect ratio + "960*960", # 1:1 aspect ratio + "832*1088", # 3:4 aspect ratio + "1088*832" # 4:3 aspect ratio + ] + + def __init__(self): + super().__init__() + # Use the centralized API endpoint from the base class + # To use Mainland China region, modify API_ENDPOINT_POST_VIDEO in wan_base.py + self.api_url = self.API_ENDPOINT_POST_VIDEO + + @classmethod + def INPUT_TYPES(cls): + # Define output directory options + if COMFYUI_AVAILABLE: + # Use ComfyUI's output directory with browseable option + output_dir_options = { + "default": "./videos", + "tooltip": "Directory where the generated video will be saved. Browse to select a custom directory." + } + else: + # Fallback to string input + output_dir_options = { + "default": "./videos", + "multiline": False + } + + return { + "required": { + "model": (cls.MODEL_OPTIONS, { + "default": "wan2.1-vace-plus" + }), + "prompt": ("STRING", { + "multiline": True, + "default": "Generate a video based on the provided reference images" + }), + "ref_images_url": ("STRING", { + "multiline": True, + "default": "", + "tooltip": "Newline-separated URLs for reference images" + }) + }, + "optional": { + "obj_or_bg": ("STRING", { + "multiline": True, + "default": "", + "tooltip": "Newline-separated values (obj/bg) corresponding to ref_images_url" + }), + "size": (cls.RESOLUTION_OPTIONS, { + "default": "1280*720" + }), + "seed": ("INT", { + "default": 0, + "min": 0, + "max": 2147483647 + }), + "prompt_extend": ("BOOLEAN", { + "default": False + }), + "watermark": ("BOOLEAN", { + "default": False + }), + "output_dir": ("STRING", output_dir_options) + } + } + + RETURN_TYPES = ("STRING",) # Returns path to downloaded video file + FUNCTION = "generate" + CATEGORY = "Ru4ls/Wan/VACE" + + def generate(self, model, prompt, ref_images_url, obj_or_bg="", size="1280*720", + seed=0, prompt_extend=False, watermark=False, output_dir="./videos"): + + # Check API key + self.check_api_key() + + # Validate required inputs + if not ref_images_url or not ref_images_url.strip(): + raise ValueError("Reference images URLs are required for image reference function") + + # Handle ref_images_url as a list + ref_images_list = [url.strip() for url in ref_images_url.split('\n') if url.strip()] + if not ref_images_list: + raise ValueError("At least one valid reference image URL is required") + + # Prepare API payload + payload = { + "model": model, + "input": { + "function": "image_reference", + "prompt": prompt, + "ref_images_url": ref_images_list + }, + "parameters": { + "size": size, + "prompt_extend": prompt_extend, + "watermark": watermark + } + } + + # Add seed if provided + if seed > 0: + payload["parameters"]["seed"] = seed + + # Handle obj_or_bg as a list + if obj_or_bg: + obj_or_bg_list = [item.strip() for item in obj_or_bg.split('\n') if item.strip()] + if obj_or_bg_list: + # Validate that obj_or_bg_list has the same length as ref_images_list + if len(obj_or_bg_list) != len(ref_images_list): + raise ValueError("obj_or_bg list must have the same length as ref_images_url list") + payload["parameters"]["obj_or_bg"] = obj_or_bg_list + else: + # If obj_or_bg is not provided, automatically generate it: + # - For single image: default to ["obj"] + # - For multiple images: assume last image is background, rest are entities + if len(ref_images_list) == 1: + payload["parameters"]["obj_or_bg"] = ["obj"] + else: + # For multiple images, assume last is bg, rest are obj + obj_or_bg_auto = ["obj"] * (len(ref_images_list) - 1) + ["bg"] + payload["parameters"]["obj_or_bg"] = obj_or_bg_auto + + # Set headers according to DashScope documentation + headers = { + "Authorization": f"Bearer {self.api_key}", + "Content-Type": "application/json", + "X-DashScope-Async": "enable" # Wan requires async processing + } + + try: + # Make API request + print(f"Making API request to {self.api_url}") + print(f"Payload: {json.dumps(payload, indent=2)}") + response = requests.post(self.api_url, headers=headers, json=payload) + print(f"Response status code: {response.status_code}") + if hasattr(response, 'text'): + print(f"Response text: {response.text[:500]}...") # Print first 500 chars + response.raise_for_status() + + # Parse response to get task_id + result = response.json() + print(f"API response received: {json.dumps(result, indent=2)[:200]}...") # Print first 200 chars + + # Check if this is a task creation response + if "output" in result and "task_id" in result["output"]: + task_id = result["output"]["task_id"] + task_status = result["output"]["task_status"] + print(f"Task created with ID: {task_id}, status: {task_status}") + + # Now we need to poll for the result + task_result = self.poll_task_result(task_id, output_dir) + return (task_result,) # Return path to downloaded video file + else: + raise ValueError(f"Unexpected API response format: {result}") + + except requests.exceptions.RequestException as e: + # More detailed error handling + if hasattr(e, 'response') and e.response is not None: + status_code = e.response.status_code + response_text = e.response.text + print(f"API request failed with status {status_code}: {response_text}") + if status_code == 401: + raise RuntimeError(f"API request failed: 401 Unauthorized. " + f"This usually means your API key is invalid or not properly configured. " + f"Error details: {response_text}") + elif status_code == 403: + raise RuntimeError(f"API request failed: 403 Forbidden. " + f"This usually means your API key is valid but you don't have access to this model. " + f"Error details: {response_text}") + elif status_code == 400: + raise RuntimeError(f"API request failed: 400 Bad Request. " + f"This usually means there's an issue with the request format. " + f"Error details: {response_text}") + else: + raise RuntimeError(f"API request failed: {status_code} {e.response.reason}. Response: {response_text}") + else: + raise RuntimeError(f"API request failed: {str(e)}") + except Exception as e: + raise RuntimeError(f"Failed to process API response: {str(e)}") + + def poll_task_result(self, task_id, output_dir="./videos"): + """Poll for task result until completion and download video""" + import time + + # URL for querying task results + # To use Mainland China region, modify API_ENDPOINT_GET in wan_base.py + query_url = self.API_ENDPOINT_GET.format(task_id=task_id) + + headers = { + "Authorization": f"Bearer {self.api_key}", + "Content-Type": "application/json" + } + + max_attempts = 60 # Maximum polling attempts (may take longer for video) + attempt = 0 + + while attempt < max_attempts: + try: + print(f"Polling task {task_id}, attempt {attempt + 1}/{max_attempts}") + response = requests.get(query_url, headers=headers) + response.raise_for_status() + + result = response.json() + task_status = result["output"]["task_status"] + print(f"Task status: {task_status}") + + if task_status == "SUCCEEDED": + # Task completed successfully + if "video_url" in result["output"]: + video_url = result["output"]["video_url"] + + # Download the video + video_response = requests.get(video_url) + video_response.raise_for_status() + + # Create a unique filename for the video + timestamp = datetime.now().strftime("%Y%m%d_%H%M%S") + video_filename = f"wan_vace_image_reference_{timestamp}.mp4" + + # Handle output directory based on ComfyUI availability + if COMFYUI_AVAILABLE and not output_dir.startswith(("./", "/")): + # Use ComfyUI's output directory structure + if output_dir.endswith("/"): + output_dir = output_dir[:-1] + full_output_folder = folder_paths.get_output_directory() + output_path = os.path.join(full_output_folder, output_dir) + else: + # Resolve output directory path (existing logic) + if output_dir.startswith("./"): + # Relative to the node directory + output_path = os.path.join(os.path.dirname(__file__), output_dir[2:]) + else: + output_path = output_dir + + # Create output directory if it doesn't exist + os.makedirs(output_path, exist_ok=True) + + # Save video to file + video_path = os.path.join(output_path, video_filename) + with open(video_path, "wb") as f: + f.write(video_response.content) + + print(f"Video downloaded and saved to: {video_path}") + # Return path relative to ComfyUI output directory if using ComfyUI + if COMFYUI_AVAILABLE and not output_dir.startswith(("./", "/")): + return os.path.join(output_dir, video_filename) if output_dir != "./videos" else video_filename + else: + return video_path # Return full path + else: + raise ValueError(f"Unexpected API response format: {result}") + + elif task_status == "FAILED": + # Task failed + error_code = result["output"].get("code", "Unknown") + error_message = result["output"].get("message", "Unknown error") + raise RuntimeError(f"Task failed with code: {error_code}, message: {error_message}") + + elif task_status in ["PENDING", "RUNNING"]: + # Task still in progress, wait and retry + time.sleep(10) # Wait 10 seconds before retrying (video generation may take longer) + attempt += 1 + continue + + else: + raise ValueError(f"Unexpected task status: {task_status}") + + except requests.exceptions.RequestException as e: + raise RuntimeError(f"Failed to query task status: {str(e)}") + + # If we've reached here, we've exceeded max attempts + raise RuntimeError(f"Task did not complete within the expected time ({max_attempts} attempts)") diff --git a/wan_vace_video_edit.py b/wan_vace_video_edit.py new file mode 100644 index 0000000..aa558d6 --- /dev/null +++ b/wan_vace_video_edit.py @@ -0,0 +1,361 @@ +""" +Wan VACE Local Video Editing Node for ComfyUI +""" + +import os +import json +import requests +from PIL import Image +import numpy as np +import torch +import io +import base64 +from dotenv import load_dotenv +import sys +import pathlib +from datetime import datetime + +# Import the base class and COMFYUI_AVAILABLE flag +from .wan_base import WanAPIBase, COMFYUI_AVAILABLE + +# Try to import folder_paths if available +try: + import folder_paths +except ImportError: + pass + +class WanVACEVideoEdit(WanAPIBase): + """Node for local video editing using Wan VACE model""" + + # Define available Wan VACE models + MODEL_OPTIONS = [ + "wan2.1-vace-plus" # Professional Edition + ] + + # Define control conditions for local editing + CONTROL_CONDITION_OPTIONS = [ + "", # No control condition + "posebodyface", # Extract facial expressions and body movements + "posebody", # Extract body movements only + "depth", # Extract composition and motion contours + "scribble" # Extract line art structure + ] + + # Define mask types for local editing + MASK_TYPE_OPTIONS = [ + "tracking", # Dynamic tracking of the target object + "fixed" # Fixed mask area + ] + + # Define expand modes for local editing + EXPAND_MODE_OPTIONS = [ + "hull", # Polygon mode + "bbox", # Bounding box mode + "original" # Raw mode + ] + + # Define video resolutions + RESOLUTION_OPTIONS = [ + "1280*720", # 16:9 aspect ratio (default) + "720*1280", # 9:16 aspect ratio + "960*960", # 1:1 aspect ratio + "832*1088", # 3:4 aspect ratio + "1088*832" # 4:3 aspect ratio + ] + + def __init__(self): + super().__init__() + # Use the centralized API endpoint from the base class + # To use Mainland China region, modify API_ENDPOINT_POST_VIDEO in wan_base.py + self.api_url = self.API_ENDPOINT_POST_VIDEO + + @classmethod + def INPUT_TYPES(cls): + # Define output directory options + if COMFYUI_AVAILABLE: + # Use ComfyUI's output directory with browseable option + output_dir_options = { + "default": "./videos", + "tooltip": "Directory where the generated video will be saved. Browse to select a custom directory." + } + else: + # Fallback to string input + output_dir_options = { + "default": "./videos", + "multiline": False + } + + return { + "required": { + "model": (cls.MODEL_OPTIONS, { + "default": "wan2.1-vace-plus" + }), + "prompt": ("STRING", { + "multiline": True, + "default": "Edit the video with the following description" + }), + "video_url": ("STRING", { + "default": "", + "tooltip": "URL of the input video" + }) + }, + "optional": { + "ref_images_url": ("STRING", { + "multiline": True, + "default": "", + "tooltip": "Newline-separated URLs for reference images (only 1 image supported)" + }), + "mask_image_url": ("STRING", { + "default": "", + "tooltip": "URL of the mask image" + }), + "mask_frame_id": ("INT", { + "default": 1, + "min": 1, + "max": 1000 + }), + "mask_video_url": ("STRING", { + "default": "", + "tooltip": "URL of the mask video" + }), + "control_condition": (cls.CONTROL_CONDITION_OPTIONS, { + "default": "" + }), + "mask_type": (cls.MASK_TYPE_OPTIONS, { + "default": "tracking" + }), + "expand_ratio": ("FLOAT", { + "default": 0.05, + "min": 0.0, + "max": 1.0, + "step": 0.01 + }), + "expand_mode": (cls.EXPAND_MODE_OPTIONS, { + "default": "hull" + }), + "size": (cls.RESOLUTION_OPTIONS, { + "default": "1280*720" + }), + "seed": ("INT", { + "default": 0, + "min": 0, + "max": 2147483647 + }), + "prompt_extend": ("BOOLEAN", { + "default": False + }), + "watermark": ("BOOLEAN", { + "default": False + }), + "output_dir": ("STRING", output_dir_options) + } + } + + RETURN_TYPES = ("STRING",) # Returns path to downloaded video file + FUNCTION = "generate" + CATEGORY = "Ru4ls/Wan/VACE" + + def generate(self, model, prompt, video_url, ref_images_url="", mask_image_url="", + mask_frame_id=1, mask_video_url="", control_condition="", mask_type="tracking", + expand_ratio=0.05, expand_mode="hull", size="1280*720", seed=0, + prompt_extend=False, watermark=False, output_dir="./videos"): + + # Check API key + self.check_api_key() + + # Prepare API payload + payload = { + "model": model, + "input": { + "function": "video_edit", + "prompt": prompt, + "video_url": video_url + }, + "parameters": { + "mask_type": mask_type, + "expand_mode": expand_mode, + "size": size, + "prompt_extend": prompt_extend, + "watermark": watermark + } + } + + # Add seed if provided + if seed > 0: + payload["parameters"]["seed"] = seed + + # Add control condition if provided + if control_condition: + payload["parameters"]["control_condition"] = control_condition + + # Add expand_ratio if not default + if expand_ratio != 0.05: + payload["parameters"]["expand_ratio"] = expand_ratio + + # Add mask_image_url if provided + if mask_image_url: + payload["input"]["mask_image_url"] = mask_image_url + + # Add mask_frame_id if not default + if mask_frame_id != 1: + payload["input"]["mask_frame_id"] = mask_frame_id + + # Add mask_video_url if provided + if mask_video_url: + payload["input"]["mask_video_url"] = mask_video_url + + # Handle ref_images_url as a list (only 1 image supported) + if ref_images_url: + ref_images_list = [url.strip() for url in ref_images_url.split('\n') if url.strip()] + if ref_images_list: + payload["input"]["ref_images_url"] = ref_images_list[:1] # Only take the first image + + # Set headers according to DashScope documentation + headers = { + "Authorization": f"Bearer {self.api_key}", + "Content-Type": "application/json", + "X-DashScope-Async": "enable" # Wan requires async processing + } + + try: + # Make API request + print(f"Making API request to {self.api_url}") + print(f"Payload: {json.dumps(payload, indent=2)}") + response = requests.post(self.api_url, headers=headers, json=payload) + print(f"Response status code: {response.status_code}") + if hasattr(response, 'text'): + print(f"Response text: {response.text[:500]}...") # Print first 500 chars + response.raise_for_status() + + # Parse response to get task_id + result = response.json() + print(f"API response received: {json.dumps(result, indent=2)[:200]}...") # Print first 200 chars + + # Check if this is a task creation response + if "output" in result and "task_id" in result["output"]: + task_id = result["output"]["task_id"] + task_status = result["output"]["task_status"] + print(f"Task created with ID: {task_id}, status: {task_status}") + + # Now we need to poll for the result + task_result = self.poll_task_result(task_id, output_dir) + return (task_result,) # Return path to downloaded video file + else: + raise ValueError(f"Unexpected API response format: {result}") + + except requests.exceptions.RequestException as e: + # More detailed error handling + if hasattr(e, 'response') and e.response is not None: + status_code = e.response.status_code + response_text = e.response.text + print(f"API request failed with status {status_code}: {response_text}") + if status_code == 401: + raise RuntimeError(f"API request failed: 401 Unauthorized. " + f"This usually means your API key is invalid or not properly configured. " + f"Error details: {response_text}") + elif status_code == 403: + raise RuntimeError(f"API request failed: 403 Forbidden. " + f"This usually means your API key is valid but you don't have access to this model. " + f"Error details: {response_text}") + elif status_code == 400: + raise RuntimeError(f"API request failed: 400 Bad Request. " + f"This usually means there's an issue with the request format. " + f"Error details: {response_text}") + else: + raise RuntimeError(f"API request failed: {status_code} {e.response.reason}. Response: {response_text}") + else: + raise RuntimeError(f"API request failed: {str(e)}") + except Exception as e: + raise RuntimeError(f"Failed to process API response: {str(e)}") + + def poll_task_result(self, task_id, output_dir="./videos"): + """"Poll for task result until completion and download video""" + import time + + # URL for querying task results + # To use Mainland China region, modify API_ENDPOINT_GET in wan_base.py + query_url = self.API_ENDPOINT_GET.format(task_id=task_id) + + headers = { + "Authorization": f"Bearer {self.api_key}", + "Content-Type": "application/json" + } + + max_attempts = 60 # Maximum polling attempts (may take longer for video) + attempt = 0 + + while attempt < max_attempts: + try: + print(f"Polling task {task_id}, attempt {attempt + 1}/{max_attempts}") + response = requests.get(query_url, headers=headers) + response.raise_for_status() + + result = response.json() + task_status = result["output"]["task_status"] + print(f"Task status: {task_status}") + + if task_status == "SUCCEEDED": + # Task completed successfully + if "video_url" in result["output"]: + video_url = result["output"]["video_url"] + + # Download the video + video_response = requests.get(video_url) + video_response.raise_for_status() + + # Create a unique filename for the video + timestamp = datetime.now().strftime("%Y%m%d_%H%M%S") + video_filename = f"wan_vace_video_edit_{timestamp}.mp4" + + # Handle output directory based on ComfyUI availability + if COMFYUI_AVAILABLE and not output_dir.startswith(("./", "/")): + # Use ComfyUI's output directory structure + if output_dir.endswith("/"): + output_dir = output_dir[:-1] + full_output_folder = folder_paths.get_output_directory() + output_path = os.path.join(full_output_folder, output_dir) + else: + # Resolve output directory path (existing logic) + if output_dir.startswith("./"): + # Relative to the node directory + output_path = os.path.join(os.path.dirname(__file__), output_dir[2:]) + else: + output_path = output_dir + + # Create output directory if it doesn't exist + os.makedirs(output_path, exist_ok=True) + + # Save video to file + video_path = os.path.join(output_path, video_filename) + with open(video_path, "wb") as f: + f.write(video_response.content) + + print(f"Video downloaded and saved to: {video_path}") + # Return path relative to ComfyUI output directory if using ComfyUI + if COMFYUI_AVAILABLE and not output_dir.startswith(("./", "/")): + return os.path.join(output_dir, video_filename) if output_dir != "./videos" else video_filename + else: + return video_path # Return full path + else: + raise ValueError(f"Unexpected API response format: {result}") + + elif task_status == "FAILED": + # Task failed + error_code = result["output"].get("code", "Unknown") + error_message = result["output"].get("message", "Unknown error") + raise RuntimeError(f"Task failed with code: {error_code}, message: {error_message}") + + elif task_status in ["PENDING", "RUNNING"]: + # Task still in progress, wait and retry + time.sleep(10) # Wait 10 seconds before retrying (video generation may take longer) + attempt += 1 + continue + + else: + raise ValueError(f"Unexpected task status: {task_status}") + + except requests.exceptions.RequestException as e: + raise RuntimeError(f"Failed to query task status: {str(e)}") + + # If we've reached here, we've exceeded max attempts + raise RuntimeError(f"Task did not complete within the expected time ({max_attempts} attempts)") diff --git a/wan_vace_video_extension.py b/wan_vace_video_extension.py new file mode 100644 index 0000000..089bd66 --- /dev/null +++ b/wan_vace_video_extension.py @@ -0,0 +1,315 @@ +""" +Wan VACE Video Extension Node for ComfyUI +""" + +import os +import json +import requests +from PIL import Image +import numpy as np +import torch +import io +import base64 +from dotenv import load_dotenv +import sys +import pathlib +from datetime import datetime + +# Import the base class and COMFYUI_AVAILABLE flag +from .wan_base import WanAPIBase, COMFYUI_AVAILABLE + +# Try to import folder_paths if available +try: + import folder_paths +except ImportError: + pass + +class WanVACEVideoExtension(WanAPIBase): + """Node for video extension using Wan VACE model""" + + # Define available Wan VACE models + MODEL_OPTIONS = [ + "wan2.1-vace-plus" # Professional Edition + ] + + # Define control conditions for video extension + CONTROL_CONDITION_OPTIONS = [ + "", # No control condition + "posebodyface", # Extract facial expressions and body movements + "posebody", # Extract body movements only + "depth", # Extract composition and motion contours + "scribble" # Extract line art structure + ] + + def __init__(self): + super().__init__() + # Use the centralized API endpoint from the base class + # To use Mainland China region, modify API_ENDPOINT_POST_VIDEO in wan_base.py + self.api_url = self.API_ENDPOINT_POST_VIDEO + + @classmethod + def INPUT_TYPES(cls): + # Define output directory options + if COMFYUI_AVAILABLE: + # Use ComfyUI's output directory with browseable option + output_dir_options = { + "default": "./videos", + "tooltip": "Directory where the generated video will be saved. Browse to select a custom directory." + } + else: + # Fallback to string input + output_dir_options = { + "default": "./videos", + "multiline": False + } + + return { + "required": { + "model": (cls.MODEL_OPTIONS, { + "default": "wan2.1-vace-plus" + }), + "prompt": ("STRING", { + "multiline": True, + "default": "Extend the video with the following description" + }) + }, + "optional": { + "first_frame_url": ("STRING", { + "default": "", + "tooltip": "URL of the first frame image" + }), + "last_frame_url": ("STRING", { + "default": "", + "tooltip": "URL of the last frame image" + }), + "first_clip_url": ("STRING", { + "default": "", + "tooltip": "URL of the first video segment" + }), + "last_clip_url": ("STRING", { + "default": "", + "tooltip": "URL of the last video segment" + }), + "video_url": ("STRING", { + "default": "", + "tooltip": "URL of the reference video for motion features" + }), + "control_condition": (cls.CONTROL_CONDITION_OPTIONS, { + "default": "" + }), + "seed": ("INT", { + "default": 0, + "min": 0, + "max": 2147483647 + }), + "prompt_extend": ("BOOLEAN", { + "default": False + }), + "watermark": ("BOOLEAN", { + "default": False + }), + "output_dir": ("STRING", output_dir_options) + } + } + + RETURN_TYPES = ("STRING",) # Returns path to downloaded video file + FUNCTION = "generate" + CATEGORY = "Ru4ls/Wan/VACE" + + def generate(self, model, prompt, first_frame_url="", last_frame_url="", + first_clip_url="", last_clip_url="", video_url="", control_condition="", + seed=0, prompt_extend=False, watermark=False, output_dir="./videos"): + + # Check API key + self.check_api_key() + + # Prepare API payload + payload = { + "model": model, + "input": { + "function": "video_extension", + "prompt": prompt + }, + "parameters": { + "prompt_extend": prompt_extend, + "watermark": watermark + } + } + + # Add seed if provided + if seed > 0: + payload["parameters"]["seed"] = seed + + # Add control condition if provided + if control_condition: + payload["parameters"]["control_condition"] = control_condition + + # Add first_frame_url if provided + if first_frame_url: + payload["input"]["first_frame_url"] = first_frame_url + + # Add last_frame_url if provided + if last_frame_url: + payload["input"]["last_frame_url"] = last_frame_url + + # Add first_clip_url if provided + if first_clip_url: + payload["input"]["first_clip_url"] = first_clip_url + + # Add last_clip_url if provided + if last_clip_url: + payload["input"]["last_clip_url"] = last_clip_url + + # Add video_url if provided + if video_url: + payload["input"]["video_url"] = video_url + + # Set headers according to DashScope documentation + headers = { + "Authorization": f"Bearer {self.api_key}", + "Content-Type": "application/json", + "X-DashScope-Async": "enable" # Wan requires async processing + } + + try: + # Make API request + print(f"Making API request to {self.api_url}") + print(f"Payload: {json.dumps(payload, indent=2)}") + response = requests.post(self.api_url, headers=headers, json=payload) + print(f"Response status code: {response.status_code}") + if hasattr(response, 'text'): + print(f"Response text: {response.text[:500]}...") # Print first 500 chars + response.raise_for_status() + + # Parse response to get task_id + result = response.json() + print(f"API response received: {json.dumps(result, indent=2)[:200]}...") # Print first 200 chars + + # Check if this is a task creation response + if "output" in result and "task_id" in result["output"]: + task_id = result["output"]["task_id"] + task_status = result["output"]["task_status"] + print(f"Task created with ID: {task_id}, status: {task_status}") + + # Now we need to poll for the result + task_result = self.poll_task_result(task_id, output_dir) + return (task_result,) # Return path to downloaded video file + else: + raise ValueError(f"Unexpected API response format: {result}") + + except requests.exceptions.RequestException as e: + # More detailed error handling + if hasattr(e, 'response') and e.response is not None: + status_code = e.response.status_code + response_text = e.response.text + print(f"API request failed with status {status_code}: {response_text}") + if status_code == 401: + raise RuntimeError(f"API request failed: 401 Unauthorized. " + f"This usually means your API key is invalid or not properly configured. " + f"Error details: {response_text}") + elif status_code == 403: + raise RuntimeError(f"API request failed: 403 Forbidden. " + f"This usually means your API key is valid but you don't have access to this model. " + f"Error details: {response_text}") + elif status_code == 400: + raise RuntimeError(f"API request failed: 400 Bad Request. " + f"This usually means there's an issue with the request format. " + f"Error details: {response_text}") + else: + raise RuntimeError(f"API request failed: {status_code} {e.response.reason}. Response: {response_text}") + else: + raise RuntimeError(f"API request failed: {str(e)}") + except Exception as e: + raise RuntimeError(f"Failed to process API response: {str(e)}") + + def poll_task_result(self, task_id, output_dir="./videos"): + """Poll for task result until completion and download video""" + import time + + # URL for querying task results + # To use Mainland China region, modify API_ENDPOINT_GET in wan_base.py + query_url = self.API_ENDPOINT_GET.format(task_id=task_id) + + headers = { + "Authorization": f"Bearer {self.api_key}", + "Content-Type": "application/json" + } + + max_attempts = 60 # Maximum polling attempts (may take longer for video) + attempt = 0 + + while attempt < max_attempts: + try: + print(f"Polling task {task_id}, attempt {attempt + 1}/{max_attempts}") + response = requests.get(query_url, headers=headers) + response.raise_for_status() + + result = response.json() + task_status = result["output"]["task_status"] + print(f"Task status: {task_status}") + + if task_status == "SUCCEEDED": + # Task completed successfully + if "video_url" in result["output"]: + video_url = result["output"]["video_url"] + + # Download the video + video_response = requests.get(video_url) + video_response.raise_for_status() + + # Create a unique filename for the video + timestamp = datetime.now().strftime("%Y%m%d_%H%M%S") + video_filename = f"wan_vace_video_extension_{timestamp}.mp4" + + # Handle output directory based on ComfyUI availability + if COMFYUI_AVAILABLE and not output_dir.startswith(("./", "/")): + # Use ComfyUI's output directory structure + if output_dir.endswith("/"): + output_dir = output_dir[:-1] + full_output_folder = folder_paths.get_output_directory() + output_path = os.path.join(full_output_folder, output_dir) + else: + # Resolve output directory path (existing logic) + if output_dir.startswith("./"): + # Relative to the node directory + output_path = os.path.join(os.path.dirname(__file__), output_dir[2:]) + else: + output_path = output_dir + + # Create output directory if it doesn't exist + os.makedirs(output_path, exist_ok=True) + + # Save video to file + video_path = os.path.join(output_path, video_filename) + with open(video_path, "wb") as f: + f.write(video_response.content) + + print(f"Video downloaded and saved to: {video_path}") + # Return path relative to ComfyUI output directory if using ComfyUI + if COMFYUI_AVAILABLE and not output_dir.startswith(("./", "/")): + return os.path.join(output_dir, video_filename) if output_dir != "./videos" else video_filename + else: + return video_path # Return full path + else: + raise ValueError(f"Unexpected API response format: {result}") + + elif task_status == "FAILED": + # Task failed + error_code = result["output"].get("code", "Unknown") + error_message = result["output"].get("message", "Unknown error") + raise RuntimeError(f"Task failed with code: {error_code}, message: {error_message}") + + elif task_status in ["PENDING", "RUNNING"]: + # Task still in progress, wait and retry + time.sleep(10) # Wait 10 seconds before retrying (video generation may take longer) + attempt += 1 + continue + + else: + raise ValueError(f"Unexpected task status: {task_status}") + + except requests.exceptions.RequestException as e: + raise RuntimeError(f"Failed to query task status: {str(e)}") + + # If we've reached here, we've exceeded max attempts + raise RuntimeError(f"Task did not complete within the expected time ({max_attempts} attempts)") diff --git a/wan_vace_video_outpainting.py b/wan_vace_video_outpainting.py new file mode 100644 index 0000000..b451752 --- /dev/null +++ b/wan_vace_video_outpainting.py @@ -0,0 +1,298 @@ +""" +Wan VACE Video Outpainting Node for ComfyUI +""" + +import os +import json +import requests +from PIL import Image +import numpy as np +import torch +import io +import base64 +from dotenv import load_dotenv +import sys +import pathlib +from datetime import datetime + +# Import the base class and COMFYUI_AVAILABLE flag +from .wan_base import WanAPIBase, COMFYUI_AVAILABLE + +# Try to import folder_paths if available +try: + import folder_paths +except ImportError: + pass + +class WanVACEVideoOutpainting(WanAPIBase): + """Node for video outpainting using Wan VACE model""" + + # Define available Wan VACE models + MODEL_OPTIONS = [ + "wan2.1-vace-plus" # Professional Edition + ] + + def __init__(self): + super().__init__() + # Use the centralized API endpoint from the base class + # To use Mainland China region, modify API_ENDPOINT_POST_VIDEO in wan_base.py + self.api_url = self.API_ENDPOINT_POST_VIDEO + + @classmethod + def INPUT_TYPES(cls): + # Define output directory options + if COMFYUI_AVAILABLE: + # Use ComfyUI's output directory with browseable option + output_dir_options = { + "default": "./videos", + "tooltip": "Directory where the generated video will be saved. Browse to select a custom directory." + } + else: + # Fallback to string input + output_dir_options = { + "default": "./videos", + "multiline": False + } + + return { + "required": { + "model": (cls.MODEL_OPTIONS, { + "default": "wan2.1-vace-plus" + }), + "prompt": ("STRING", { + "multiline": True, + "default": "Outpaint the video with the following description" + }), + "video_url": ("STRING", { + "default": "", + "tooltip": "URL of the input video" + }) + }, + "optional": { + "top_scale": ("FLOAT", { + "default": 1.0, + "min": 1.0, + "max": 2.0, + "step": 0.1 + }), + "bottom_scale": ("FLOAT", { + "default": 1.0, + "min": 1.0, + "max": 2.0, + "step": 0.1 + }), + "left_scale": ("FLOAT", { + "default": 1.0, + "min": 1.0, + "max": 2.0, + "step": 0.1 + }), + "right_scale": ("FLOAT", { + "default": 1.0, + "min": 1.0, + "max": 2.0, + "step": 0.1 + }), + "seed": ("INT", { + "default": 0, + "min": 0, + "max": 2147483647 + }), + "prompt_extend": ("BOOLEAN", { + "default": False + }), + "watermark": ("BOOLEAN", { + "default": False + }), + "output_dir": ("STRING", output_dir_options) + } + } + + RETURN_TYPES = ("STRING",) # Returns path to downloaded video file + FUNCTION = "generate" + CATEGORY = "Ru4ls/Wan/VACE" + + def generate(self, model, prompt, video_url, top_scale=1.0, bottom_scale=1.0, + left_scale=1.0, right_scale=1.0, seed=0, prompt_extend=False, + watermark=False, output_dir="./videos"): + + # Check API key + self.check_api_key() + + # Prepare API payload + payload = { + "model": model, + "input": { + "function": "video_outpainting", + "prompt": prompt, + "video_url": video_url + }, + "parameters": { + "prompt_extend": prompt_extend, + "watermark": watermark + } + } + + # Add seed if provided + if seed > 0: + payload["parameters"]["seed"] = seed + + # Add scale parameters if not default + if top_scale != 1.0: + payload["parameters"]["top_scale"] = top_scale + if bottom_scale != 1.0: + payload["parameters"]["bottom_scale"] = bottom_scale + if left_scale != 1.0: + payload["parameters"]["left_scale"] = left_scale + if right_scale != 1.0: + payload["parameters"]["right_scale"] = right_scale + + # Set headers according to DashScope documentation + headers = { + "Authorization": f"Bearer {self.api_key}", + "Content-Type": "application/json", + "X-DashScope-Async": "enable" # Wan requires async processing + } + + try: + # Make API request + print(f"Making API request to {self.api_url}") + print(f"Payload: {json.dumps(payload, indent=2)}") + response = requests.post(self.api_url, headers=headers, json=payload) + print(f"Response status code: {response.status_code}") + if hasattr(response, 'text'): + print(f"Response text: {response.text[:500]}...") # Print first 500 chars + response.raise_for_status() + + # Parse response to get task_id + result = response.json() + print(f"API response received: {json.dumps(result, indent=2)[:200]}...") # Print first 200 chars + + # Check if this is a task creation response + if "output" in result and "task_id" in result["output"]: + task_id = result["output"]["task_id"] + task_status = result["output"]["task_status"] + print(f"Task created with ID: {task_id}, status: {task_status}") + + # Now we need to poll for the result + task_result = self.poll_task_result(task_id, output_dir) + return (task_result,) # Return path to downloaded video file + else: + raise ValueError(f"Unexpected API response format: {result}") + + except requests.exceptions.RequestException as e: + # More detailed error handling + if hasattr(e, 'response') and e.response is not None: + status_code = e.response.status_code + response_text = e.response.text + print(f"API request failed with status {status_code}: {response_text}") + if status_code == 401: + raise RuntimeError(f"API request failed: 401 Unauthorized. " + f"This usually means your API key is invalid or not properly configured. " + f"Error details: {response_text}") + elif status_code == 403: + raise RuntimeError(f"API request failed: 403 Forbidden. " + f"This usually means your API key is valid but you don't have access to this model. " + f"Error details: {response_text}") + elif status_code == 400: + raise RuntimeError(f"API request failed: 400 Bad Request. " + f"This usually means there's an issue with the request format. " + f"Error details: {response_text}") + else: + raise RuntimeError(f"API request failed: {status_code} {e.response.reason}. Response: {response_text}") + else: + raise RuntimeError(f"API request failed: {str(e)}") + except Exception as e: + raise RuntimeError(f"Failed to process API response: {str(e)}") + + def poll_task_result(self, task_id, output_dir="./videos"): + """Poll for task result until completion and download video""" + import time + + # URL for querying task results + # To use Mainland China region, modify API_ENDPOINT_GET in wan_base.py + query_url = self.API_ENDPOINT_GET.format(task_id=task_id) + + headers = { + "Authorization": f"Bearer {self.api_key}", + "Content-Type": "application/json" + } + + max_attempts = 60 # Maximum polling attempts (may take longer for video) + attempt = 0 + + while attempt < max_attempts: + try: + print(f"Polling task {task_id}, attempt {attempt + 1}/{max_attempts}") + response = requests.get(query_url, headers=headers) + response.raise_for_status() + + result = response.json() + task_status = result["output"]["task_status"] + print(f"Task status: {task_status}") + + if task_status == "SUCCEEDED": + # Task completed successfully + if "video_url" in result["output"]: + video_url = result["output"]["video_url"] + + # Download the video + video_response = requests.get(video_url) + video_response.raise_for_status() + + # Create a unique filename for the video + timestamp = datetime.now().strftime("%Y%m%d_%H%M%S") + video_filename = f"wan_vace_video_outpainting_{timestamp}.mp4" + + # Handle output directory based on ComfyUI availability + if COMFYUI_AVAILABLE and not output_dir.startswith(("./", "/")): + # Use ComfyUI's output directory structure + if output_dir.endswith("/"): + output_dir = output_dir[:-1] + full_output_folder = folder_paths.get_output_directory() + output_path = os.path.join(full_output_folder, output_dir) + else: + # Resolve output directory path (existing logic) + if output_dir.startswith("./"): + # Relative to the node directory + output_path = os.path.join(os.path.dirname(__file__), output_dir[2:]) + else: + output_path = output_dir + + # Create output directory if it doesn't exist + os.makedirs(output_path, exist_ok=True) + + # Save video to file + video_path = os.path.join(output_path, video_filename) + with open(video_path, "wb") as f: + f.write(video_response.content) + + print(f"Video downloaded and saved to: {video_path}") + # Return path relative to ComfyUI output directory if using ComfyUI + if COMFYUI_AVAILABLE and not output_dir.startswith(("./", "/")): + return os.path.join(output_dir, video_filename) if output_dir != "./videos" else video_filename + else: + return video_path # Return full path + else: + raise ValueError(f"Unexpected API response format: {result}") + + elif task_status == "FAILED": + # Task failed + error_code = result["output"].get("code", "Unknown") + error_message = result["output"].get("message", "Unknown error") + raise RuntimeError(f"Task failed with code: {error_code}, message: {error_message}") + + elif task_status in ["PENDING", "RUNNING"]: + # Task still in progress, wait and retry + time.sleep(10) # Wait 10 seconds before retrying (video generation may take longer) + attempt += 1 + continue + + else: + raise ValueError(f"Unexpected task status: {task_status}") + + except requests.exceptions.RequestException as e: + raise RuntimeError(f"Failed to query task status: {str(e)}") + + # If we've reached here, we've exceeded max attempts + raise RuntimeError(f"Task did not complete within the expected time ({max_attempts} attempts)") diff --git a/wan_vace_video_repainting.py b/wan_vace_video_repainting.py new file mode 100644 index 0000000..708a3bf --- /dev/null +++ b/wan_vace_video_repainting.py @@ -0,0 +1,296 @@ +""" +Wan VACE Video Repainting Node for ComfyUI +""" + +import os +import json +import requests +from PIL import Image +import numpy as np +import torch +import io +import base64 +from dotenv import load_dotenv +import sys +import pathlib +from datetime import datetime + +# Import the base class and COMFYUI_AVAILABLE flag +from .wan_base import WanAPIBase, COMFYUI_AVAILABLE + +# Try to import folder_paths if available +try: + import folder_paths +except ImportError: + pass + +class WanVACEVideoRepainting(WanAPIBase): + """Node for video repainting using Wan VACE model""" + + # Define available Wan VACE models + MODEL_OPTIONS = [ + "wan2.1-vace-plus" # Professional Edition + ] + + # Define control conditions for video repainting + CONTROL_CONDITION_OPTIONS = [ + "posebodyface", # Extract facial expressions and body movements + "posebody", # Extract body movements only + "depth", # Extract composition and motion contours + "scribble" # Extract line art structure + ] + + def __init__(self): + super().__init__() + # Use the centralized API endpoint from the base class + # To use Mainland China region, modify API_ENDPOINT_POST_VIDEO in wan_base.py + self.api_url = self.API_ENDPOINT_POST_VIDEO + + @classmethod + def INPUT_TYPES(cls): + # Define output directory options + if COMFYUI_AVAILABLE: + # Use ComfyUI's output directory with browseable option + output_dir_options = { + "default": "./videos", + "tooltip": "Directory where the generated video will be saved. Browse to select a custom directory." + } + else: + # Fallback to string input + output_dir_options = { + "default": "./videos", + "multiline": False + } + + return { + "required": { + "model": (cls.MODEL_OPTIONS, { + "default": "wan2.1-vace-plus" + }), + "prompt": ("STRING", { + "multiline": True, + "default": "Repaint the video with the following description" + }), + "video_url": ("STRING", { + "default": "", + "tooltip": "URL of the input video" + }) + }, + "optional": { + "ref_images_url": ("STRING", { + "multiline": True, + "default": "", + "tooltip": "Newline-separated URLs for reference images (only 1 image supported)" + }), + "control_condition": (cls.CONTROL_CONDITION_OPTIONS, { + "default": "depth" + }), + "strength": ("FLOAT", { + "default": 1.0, + "min": 0.0, + "max": 1.0, + "step": 0.1 + }), + "seed": ("INT", { + "default": 0, + "min": 0, + "max": 2147483647 + }), + "prompt_extend": ("BOOLEAN", { + "default": False + }), + "watermark": ("BOOLEAN", { + "default": False + }), + "output_dir": ("STRING", output_dir_options) + } + } + + RETURN_TYPES = ("STRING",) # Returns path to downloaded video file + FUNCTION = "generate" + CATEGORY = "Ru4ls/Wan/VACE" + + def generate(self, model, prompt, video_url, ref_images_url="", control_condition="depth", + strength=1.0, seed=0, prompt_extend=False, watermark=False, output_dir="./videos"): + + # Check API key + self.check_api_key() + + # Prepare API payload + payload = { + "model": model, + "input": { + "function": "video_repainting", + "prompt": prompt, + "video_url": video_url + }, + "parameters": { + "control_condition": control_condition, + "prompt_extend": prompt_extend, + "watermark": watermark + } + } + + # Add seed if provided + if seed > 0: + payload["parameters"]["seed"] = seed + + # Add strength if not default + if strength != 1.0: + payload["parameters"]["strength"] = strength + + # Handle ref_images_url as a list (only 1 image supported) + if ref_images_url: + ref_images_list = [url.strip() for url in ref_images_url.split('\n') if url.strip()] + if ref_images_list: + payload["input"]["ref_images_url"] = ref_images_list[:1] # Only take the first image + + # Set headers according to DashScope documentation + headers = { + "Authorization": f"Bearer {self.api_key}", + "Content-Type": "application/json", + "X-DashScope-Async": "enable" # Wan requires async processing + } + + try: + # Make API request + print(f"Making API request to {self.api_url}") + print(f"Payload: {json.dumps(payload, indent=2)}") + response = requests.post(self.api_url, headers=headers, json=payload) + print(f"Response status code: {response.status_code}") + if hasattr(response, 'text'): + print(f"Response text: {response.text[:500]}...") # Print first 500 chars + response.raise_for_status() + + # Parse response to get task_id + result = response.json() + print(f"API response received: {json.dumps(result, indent=2)[:200]}...") # Print first 200 chars + + # Check if this is a task creation response + if "output" in result and "task_id" in result["output"]: + task_id = result["output"]["task_id"] + task_status = result["output"]["task_status"] + print(f"Task created with ID: {task_id}, status: {task_status}") + + # Now we need to poll for the result + task_result = self.poll_task_result(task_id, output_dir) + return (task_result,) # Return path to downloaded video file + else: + raise ValueError(f"Unexpected API response format: {result}") + + except requests.exceptions.RequestException as e: + # More detailed error handling + if hasattr(e, 'response') and e.response is not None: + status_code = e.response.status_code + response_text = e.response.text + print(f"API request failed with status {status_code}: {response_text}") + if status_code == 401: + raise RuntimeError(f"API request failed: 401 Unauthorized. " + f"This usually means your API key is invalid or not properly configured. " + f"Error details: {response_text}") + elif status_code == 403: + raise RuntimeError(f"API request failed: 403 Forbidden. " + f"This usually means your API key is valid but you don't have access to this model. " + f"Error details: {response_text}") + elif status_code == 400: + raise RuntimeError(f"API request failed: 400 Bad Request. " + f"This usually means there's an issue with the request format. " + f"Error details: {response_text}") + else: + raise RuntimeError(f"API request failed: {status_code} {e.response.reason}. Response: {response_text}") + else: + raise RuntimeError(f"API request failed: {str(e)}") + except Exception as e: + raise RuntimeError(f"Failed to process API response: {str(e)}") + + def poll_task_result(self, task_id, output_dir="./videos"): + """Poll for task result until completion and download video""" + import time + + # URL for querying task results + # To use Mainland China region, modify API_ENDPOINT_GET in wan_base.py + query_url = self.API_ENDPOINT_GET.format(task_id=task_id) + + headers = { + "Authorization": f"Bearer {self.api_key}", + "Content-Type": "application/json" + } + + max_attempts = 60 # Maximum polling attempts (may take longer for video) + attempt = 0 + + while attempt < max_attempts: + try: + print(f"Polling task {task_id}, attempt {attempt + 1}/{max_attempts}") + response = requests.get(query_url, headers=headers) + response.raise_for_status() + + result = response.json() + task_status = result["output"]["task_status"] + print(f"Task status: {task_status}") + + if task_status == "SUCCEEDED": + # Task completed successfully + if "video_url" in result["output"]: + video_url = result["output"]["video_url"] + + # Download the video + video_response = requests.get(video_url) + video_response.raise_for_status() + + # Create a unique filename for the video + timestamp = datetime.now().strftime("%Y%m%d_%H%M%S") + video_filename = f"wan_vace_video_repainting_{timestamp}.mp4" + + # Handle output directory based on ComfyUI availability + if COMFYUI_AVAILABLE and not output_dir.startswith(("./", "/")): + # Use ComfyUI's output directory structure + if output_dir.endswith("/"): + output_dir = output_dir[:-1] + full_output_folder = folder_paths.get_output_directory() + output_path = os.path.join(full_output_folder, output_dir) + else: + # Resolve output directory path (existing logic) + if output_dir.startswith("./"): + # Relative to the node directory + output_path = os.path.join(os.path.dirname(__file__), output_dir[2:]) + else: + output_path = output_dir + + # Create output directory if it doesn't exist + os.makedirs(output_path, exist_ok=True) + + # Save video to file + video_path = os.path.join(output_path, video_filename) + with open(video_path, "wb") as f: + f.write(video_response.content) + + print(f"Video downloaded and saved to: {video_path}") + # Return path relative to ComfyUI output directory if using ComfyUI + if COMFYUI_AVAILABLE and not output_dir.startswith(("./", "/")): + return os.path.join(output_dir, video_filename) if output_dir != "./videos" else video_filename + else: + return video_path # Return full path + else: + raise ValueError(f"Unexpected API response format: {result}") + + elif task_status == "FAILED": + # Task failed + error_code = result["output"].get("code", "Unknown") + error_message = result["output"].get("message", "Unknown error") + raise RuntimeError(f"Task failed with code: {error_code}, message: {error_message}") + + elif task_status in ["PENDING", "RUNNING"]: + # Task still in progress, wait and retry + time.sleep(10) # Wait 10 seconds before retrying (video generation may take longer) + attempt += 1 + continue + + else: + raise ValueError(f"Unexpected task status: {task_status}") + + except requests.exceptions.RequestException as e: + raise RuntimeError(f"Failed to query task status: {str(e)}") + + # If we've reached here, we've exceeded max attempts + raise RuntimeError(f"Task did not complete within the expected time ({max_attempts} attempts)")