diff --git a/examples/i2v_workflow.json b/examples/i2v_workflow.json index be2e605..f5abcb1 100644 --- a/examples/i2v_workflow.json +++ b/examples/i2v_workflow.json @@ -1,8 +1,8 @@ { "id": "8e881096-4f7b-4633-b34b-6b9c0f8aa093", "revision": 0, - "last_node_id": 103, - "last_link_id": 70, + "last_node_id": 104, + "last_link_id": 72, "nodes": [ { "id": 86, @@ -78,42 +78,6 @@ "textEmbedding" ] }, - { - "id": 100, - "type": "SetNode", - "pos": [ - 2916.776611328125, - -288.08544921875 - ], - "size": [ - 210, - 60 - ], - "flags": {}, - "order": 14, - "mode": 0, - "inputs": [ - { - "name": "LIGHT_WAN_MODEL", - "type": "LIGHT_WAN_MODEL", - "link": 65 - } - ], - "outputs": [ - { - "name": "*", - "type": "*", - "links": null - } - ], - "title": "Set_WanModel", - "properties": { - "previousName": "WanModel" - }, - "widgets_values": [ - "WanModel" - ] - }, { "id": 99, "type": "GetNode", @@ -364,7 +328,7 @@ 154 ], "flags": {}, - "order": 10, + "order": 9, "mode": 0, "inputs": [ { @@ -397,62 +361,6 @@ "" ] }, - { - "id": 88, - "type": "Lightx2vWanVideoModelLoader", - "pos": [ - 2540.527587890625, - -307.2720642089844 - ], - "size": [ - 315, - 274 - ], - "flags": {}, - "order": 9, - "mode": 0, - "inputs": [ - { - "name": "teacache_args", - "shape": 7, - "type": "LIGHT_TEACACHEARGS", - "link": null - }, - { - "name": "model_dir", - "shape": 7, - "type": "STRING", - "widget": { - "name": "model_dir" - }, - "link": 69 - } - ], - "outputs": [ - { - "name": "wan_model", - "type": "LIGHT_WAN_MODEL", - "links": [ - 65 - ] - } - ], - "properties": { - "Node name for S&R": "Lightx2vWanVideoModelLoader" - }, - "widgets_values": [ - "", - "i2v", - "bf16", - "cuda", - "flash_attn3", - false, - "Default", - "", - 1.0700000000000003, - "/mnt/aigc/users/lijiaqi2/wan_model/Wan2.1-I2V-14B-480P" - ] - }, { "id": 87, "type": "Lightx2vWanVideoT5EncoderLoader", @@ -537,57 +445,6 @@ }, "widgets_values": [] }, - { - "id": 90, - "type": "Lightx2vWanVideoSampler", - "pos": [ - 2430.891845703125, - 129.44773864746094 - ], - "size": [ - 327.5999755859375, - 194 - ], - "flags": {}, - "order": 6, - "mode": 0, - "inputs": [ - { - "name": "model", - "type": "LIGHT_WAN_MODEL", - "link": 66 - }, - { - "name": "text_embeddings", - "type": "LIGHT_TEXT_EMBEDDINGS", - "link": 64 - }, - { - "name": "image_embeddings", - "type": "LIGHT_IMAGE_EMBEDDINGS", - "link": 62 - } - ], - "outputs": [ - { - "name": "latent", - "type": "LIGHT_LATENT", - "links": [ - 55 - ] - } - ], - "properties": { - "Node name for S&R": "Lightx2vWanVideoSampler" - }, - "widgets_values": [ - 20, - 8, - 1, - 42, - "fixed" - ] - }, { "id": 85, "type": "Lightx2vWanVideoImageEncoder", @@ -600,7 +457,7 @@ 146 ], "flags": {}, - "order": 15, + "order": 14, "mode": 0, "inputs": [ { @@ -700,7 +557,7 @@ "hidden": false, "paused": false, "params": { - "filename": "AnimateDiff_00019.mp4", + "filename": "AnimateDiff_00565.mp4", "subfolder": "", "type": "output", "format": "video/h264-mp4", @@ -710,40 +567,6 @@ } } }, - { - "id": 102, - "type": "Lightx2vWanVideoModelDir", - "pos": [ - 348.7199401855469, - -279.3098449707031 - ], - "size": [ - 351.8695373535156, - 146.87998962402344 - ], - "flags": {}, - "order": 4, - "mode": 0, - "inputs": [], - "outputs": [ - { - "name": "STRING", - "type": "STRING", - "links": [ - 67, - 68, - 69, - 70 - ] - } - ], - "properties": { - "Node name for S&R": "Lightx2vWanVideoModelDir" - }, - "widgets_values": [ - "/mnt/aigc/users/lijiaqi2/wan_model/Wan2.1-I2V-14B-720P-cfg" - ] - }, { "id": 93, "type": "LoadImage", @@ -756,7 +579,7 @@ 314.0000305175781 ], "flags": {}, - "order": 5, + "order": 4, "mode": 0, "inputs": [], "outputs": [ @@ -780,6 +603,184 @@ "img_0.jpg", "image" ] + }, + { + "id": 100, + "type": "SetNode", + "pos": [ + 2525.461669921875, + -330.6775207519531 + ], + "size": [ + 210, + 60 + ], + "flags": {}, + "order": 15, + "mode": 0, + "inputs": [ + { + "name": "LIGHT_WAN_MODEL", + "type": "LIGHT_WAN_MODEL", + "link": 72 + } + ], + "outputs": [ + { + "name": "*", + "type": "*", + "links": null + } + ], + "title": "Set_WanModel", + "properties": { + "previousName": "WanModel" + }, + "widgets_values": [ + "WanModel" + ] + }, + { + "id": 104, + "type": "Lightx2vWanVideoModelLoader", + "pos": [ + 2172.787841796875, + -329.7485656738281 + ], + "size": [ + 282.2554626464844, + 298 + ], + "flags": {}, + "order": 10, + "mode": 0, + "inputs": [ + { + "name": "teacache_args", + "shape": 7, + "type": "LIGHT_TEACACHEARGS", + "link": null + }, + { + "name": "model_dir", + "shape": 7, + "type": "STRING", + "widget": { + "name": "model_dir" + }, + "link": 71 + } + ], + "outputs": [ + { + "name": "wan_model", + "type": "LIGHT_WAN_MODEL", + "links": [ + 72 + ] + } + ], + "properties": { + "Node name for S&R": "Lightx2vWanVideoModelLoader" + }, + "widgets_values": [ + "", + "i2v", + "bf16", + "cuda", + "flash_attn3", + false, + "phase", + "", + "", + 1, + "/mnt/aigc/users/lijiaqi2/wan_model/Wan2.1-I2V-14B-480P" + ] + }, + { + "id": 102, + "type": "Lightx2vWanVideoModelDir", + "pos": [ + 348.7199401855469, + -279.3098449707031 + ], + "size": [ + 351.8695373535156, + 146.87998962402344 + ], + "flags": {}, + "order": 5, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "model_dir", + "type": "STRING", + "links": [ + 67, + 68, + 70, + 71 + ] + } + ], + "properties": { + "Node name for S&R": "Lightx2vWanVideoModelDir" + }, + "widgets_values": [ + "/mnt/aigc/users/lijiaqi2/wan_model/Wan2.1-I2V-14B-720P-cfg" + ] + }, + { + "id": 90, + "type": "Lightx2vWanVideoSampler", + "pos": [ + 2430.891845703125, + 129.44773864746094 + ], + "size": [ + 327.5999755859375, + 194 + ], + "flags": {}, + "order": 6, + "mode": 0, + "inputs": [ + { + "name": "model", + "type": "LIGHT_WAN_MODEL", + "link": 66 + }, + { + "name": "text_embeddings", + "type": "LIGHT_TEXT_EMBEDDINGS", + "link": 64 + }, + { + "name": "image_embeddings", + "type": "LIGHT_IMAGE_EMBEDDINGS", + "link": 62 + } + ], + "outputs": [ + { + "name": "latent", + "type": "LIGHT_LATENT", + "links": [ + 55 + ] + } + ], + "properties": { + "Node name for S&R": "Lightx2vWanVideoSampler" + }, + "widgets_values": [ + 20, + 8, + 1, + 42, + "decrement" + ] } ], "links": [ @@ -879,14 +880,6 @@ 1, "LIGHT_TEXT_EMBEDDINGS" ], - [ - 65, - 88, - 0, - 100, - 0, - "*" - ], [ 66, 101, @@ -911,14 +904,6 @@ 0, "STRING" ], - [ - 69, - 102, - 0, - 88, - 1, - "STRING" - ], [ 70, 102, @@ -926,6 +911,22 @@ 83, 0, "STRING" + ], + [ + 71, + 102, + 0, + 104, + 1, + "STRING" + ], + [ + 72, + 104, + 0, + 100, + 0, + "LIGHT_WAN_MODEL" ] ], "groups": [ @@ -985,10 +986,10 @@ "config": {}, "extra": { "ds": { - "scale": 0.7513148009015778, + "scale": 0.8264462809917354, "offset": [ - -41.12401565861029, - -19.674610850016407 + -1038.6861717665897, + 160.02657251659417 ] }, "frontendVersion": "1.19.9" diff --git a/lightx2v_refactored/__init__.py b/lightx2v_nodes/__init__.py similarity index 89% rename from lightx2v_refactored/__init__.py rename to lightx2v_nodes/__init__.py index f719866..21b47a9 100644 --- a/lightx2v_refactored/__init__.py +++ b/lightx2v_nodes/__init__.py @@ -1,4 +1,6 @@ # Refactored LightX2V module +from ..lightx2v.lightx2v.common.ops import * # noqa: F401, F403 for import global register + from .config import LightX2VConfig from .factory import LightX2VFactory from .models import ( @@ -25,7 +27,7 @@ __all__ = [ "LightX2VConfig", "LightX2VFactory", "LightX2VT5Encoder", - "LightX2VClipVisionEncoder", + "LightX2VClipVisionEncoder", "LightX2VVae", "LightX2VModel", "Lightx2vWanVideoModelDir", @@ -39,4 +41,4 @@ __all__ = [ "Lightx2vWanVideoModelLoader", "Lightx2vWanVideoSampler", "WanVideoTeaCache", -] \ No newline at end of file +] diff --git a/lightx2v_refactored/config.py b/lightx2v_nodes/config.py similarity index 82% rename from lightx2v_refactored/config.py rename to lightx2v_nodes/config.py index d51fc1e..c3e7c10 100644 --- a/lightx2v_refactored/config.py +++ b/lightx2v_nodes/config.py @@ -1,4 +1,5 @@ """Configuration management for LightX2V.""" + from dataclasses import dataclass, field from typing import Optional, Dict, Any, List, Tuple from pathlib import Path @@ -9,6 +10,7 @@ from easydict import EasyDict @dataclass class TeaCacheConfig: """Configuration for TeaCache optimization.""" + rel_l1_thresh: float = 0.26 start_percent: float = 0.1 end_percent: float = 1.0 @@ -21,20 +23,22 @@ class TeaCacheConfig: @dataclass class VideoConfig: """Video generation configuration.""" - width: int = 832 - height: int = 480 - num_frames: int = 81 + + target_width: int = 832 + target_height: int = 480 + target_video_length: int = 81 vae_stride: Tuple[int, int, int] = (4, 8, 8) patch_size: Tuple[int, int, int] = (1, 2, 2) - + @property def max_area(self) -> int: - return self.height * self.width + return self.target_height * self.target_width @dataclass class ModelConfig: """Model loading and inference configuration.""" + model_path: Path model_type: str = "i2v" # "t2v" or "i2v" precision: str = "bf16" # "bf16", "fp16", "fp32" @@ -42,18 +46,19 @@ class ModelConfig: attention_type: str = "flash_attn3" cpu_offload: bool = False offload_granularity: str = "phase" # "block" or "phase" - + # Optional configurations lora_path: Optional[Path] = None lora_strength: float = 1.0 mm_config: Dict[str, Any] = field(default_factory=dict) - + # Inference settings steps: int = 20 shift: float = 5.0 cfg_scale: float = 5.0 seed: int = 42 - + feature_caching: str = "NoCaching" + def to_dtype(self) -> torch.dtype: """Convert precision string to torch dtype.""" dtype_map = { @@ -62,7 +67,7 @@ class ModelConfig: "fp32": torch.float32, } return dtype_map[self.precision] - + def to_device(self) -> torch.device: """Get torch device.""" if self.device == "cuda": @@ -73,20 +78,21 @@ class ModelConfig: @dataclass class EncoderConfig: """Encoder configuration.""" + model_path: Path dtype: torch.dtype device: torch.device - + # T5 specific text_len: int = 512 tokenizer_path: Optional[Path] = None cpu_offload: bool = False - + # CLIP specific clip_quantized: bool = False clip_quantized_ckpt: Optional[Path] = None quant_scheme: Optional[str] = None - + # VAE specific z_dim: int = 16 parallel: bool = False @@ -95,26 +101,23 @@ class EncoderConfig: @dataclass class LightX2VConfig: """Main configuration container for LightX2V.""" + model: ModelConfig video: VideoConfig teacache: Optional[TeaCacheConfig] = None - + @classmethod def from_dict(cls, config_dict: Dict[str, Any]) -> "LightX2VConfig": """Create configuration from dictionary.""" model_config = ModelConfig(**config_dict.get("model", {})) video_config = VideoConfig(**config_dict.get("video", {})) - + teacache_config = None if "teacache" in config_dict: teacache_config = TeaCacheConfig(**config_dict["teacache"]) - - return cls( - model=model_config, - video=video_config, - teacache=teacache_config - ) - + + return cls(model=model_config, video=video_config, teacache=teacache_config) + def to_easydict(self) -> EasyDict: """Convert to EasyDict for legacy compatibility.""" config_dict = { @@ -125,9 +128,9 @@ class LightX2VConfig: "attention_type": self.model.attention_type, "cpu_offload": self.model.cpu_offload, "offload_granularity": self.model.offload_granularity, - "target_height": self.video.height, - "target_width": self.video.width, - "target_video_length": self.video.num_frames, + "target_height": self.video.target_height, + "target_width": self.video.target_width, + "target_video_length": self.video.target_video_length, "vae_stride": self.video.vae_stride, "patch_size": self.video.patch_size, "infer_steps": self.model.steps, @@ -136,20 +139,23 @@ class LightX2VConfig: "seed": self.model.seed, "enable_cfg": self.model.cfg_scale != 1.0, "mm_config": self.model.mm_config, + "feature_caching": self.model.feature_caching, } - + if self.teacache: - config_dict.update({ - "feature_caching": "Tea", - "teacache_thresh": self.teacache.rel_l1_thresh, - "use_ret_steps": self.teacache.use_ret_steps, - "coefficients": self.teacache.coefficients, - }) + config_dict.update( + { + "feature_caching": "Tea", + "teacache_thresh": self.teacache.rel_l1_thresh, + "use_ret_steps": self.teacache.use_ret_steps, + "coefficients": self.teacache.coefficients, + } + ) else: config_dict["feature_caching"] = "NoCaching" - + if self.model.lora_path: config_dict["lora_path"] = str(self.model.lora_path) config_dict["strength_model"] = self.model.lora_strength - - return EasyDict(config_dict) \ No newline at end of file + + return EasyDict(config_dict) diff --git a/lightx2v_refactored/factory.py b/lightx2v_nodes/factory.py similarity index 86% rename from lightx2v_refactored/factory.py rename to lightx2v_nodes/factory.py index 4244cae..e061d92 100644 --- a/lightx2v_refactored/factory.py +++ b/lightx2v_nodes/factory.py @@ -1,4 +1,5 @@ """Factory pattern for creating LightX2V components.""" + from pathlib import Path from typing import Optional, Dict, Any, Union import torch @@ -22,7 +23,7 @@ from ..lightx2v.lightx2v.models.networks.wan.lora_adapter import WanLoraWrapper class LightX2VFactory: """Factory for creating LightX2V components with proper configuration.""" - + @staticmethod def create_t5_encoder( model_path: Union[str, Path], @@ -33,13 +34,13 @@ class LightX2VFactory: ) -> LightX2VT5Encoder: """Create a T5 encoder with configuration.""" model_path = Path(model_path) - + # Auto-detect tokenizer path if not provided if tokenizer_path is None: tokenizer_path = model_path.parent / "google" / "umt5-xxl" if not tokenizer_path.exists(): raise ValueError(f"Tokenizer not found at {tokenizer_path}") - + config = EncoderConfig( model_path=model_path, tokenizer_path=Path(tokenizer_path), @@ -47,7 +48,7 @@ class LightX2VFactory: device=device or torch.device("cuda"), cpu_offload=cpu_offload, ) - + # Create underlying T5 model t5_model = T5EncoderModel( text_len=config.text_len, @@ -58,9 +59,9 @@ class LightX2VFactory: shard_fn=None, cpu_offload=config.cpu_offload, ) - + return LightX2VT5Encoder(t5_model, config) - + @staticmethod def create_clip_vision_encoder( model_path: Union[str, Path], @@ -79,7 +80,7 @@ class LightX2VFactory: clip_quantized_ckpt=Path(clip_quantized_ckpt) if clip_quantized_ckpt else None, quant_scheme=quant_scheme, ) - + # Create underlying CLIP model clip_model = ClipVisionModel( dtype=config.dtype, @@ -89,9 +90,9 @@ class LightX2VFactory: clip_quantized_ckpt=str(config.clip_quantized_ckpt) if config.clip_quantized_ckpt else None, quant_scheme=config.quant_scheme, ) - + return LightX2VClipVisionEncoder(clip_model, config) - + @staticmethod def create_vae( model_path: Union[str, Path], @@ -108,7 +109,7 @@ class LightX2VFactory: parallel=parallel, z_dim=z_dim, ) - + # Create underlying VAE model vae_model = WanVAE( z_dim=config.z_dim, @@ -117,9 +118,9 @@ class LightX2VFactory: device=config.device, parallel=config.parallel, ) - + return LightX2VVae(vae_model, config) - + @staticmethod def create_model( config: ModelConfig, @@ -129,14 +130,15 @@ class LightX2VFactory: # Load config.json if it exists config_json_path = config.model_path / "config.json" config_json = {} - + if config_json_path.exists(): import json + with open(config_json_path, "r") as f: config_json = json.load(f) else: logging.warning(f"Config file not found at {config_json_path}") - + # Create model configuration dict model_config_dict = { "model_path": str(config.model_path), @@ -152,33 +154,36 @@ class LightX2VFactory: "parallel_attn_type": None, "parallel_vae": False, "use_bfloat16": config.to_dtype() == torch.bfloat16, + "feature_caching": config.feature_caching, + "self_attn_1_type": "flash_attn3", + "cross_attn_1_type": "flash_attn3", + "cross_attn_2_type": "flash_attn3", } - + # Add video config if provided if video_config: - model_config_dict.update({ - "target_height": video_config.height, - "target_width": video_config.width, - "target_video_length": video_config.num_frames, - "vae_stride": video_config.vae_stride, - "patch_size": video_config.patch_size, - "max_area": video_config.max_area, - }) - + model_config_dict.update( + { + "target_height": video_config.target_height, + "target_width": video_config.target_width, + "target_video_length": video_config.target_video_length, + "vae_stride": video_config.vae_stride, + "patch_size": video_config.patch_size, + "max_area": video_config.max_area, + } + ) + # Merge with config.json model_config_dict.update(config_json) - + # Create EasyDict for compatibility from easydict import EasyDict + easydict_config = EasyDict(model_config_dict) - + # Create underlying model - wan_model = WanModel( - str(config.model_path), - easydict_config, - config.to_device() - ) - + wan_model = WanModel(str(config.model_path), easydict_config, config.to_device()) + # Apply LoRA if specified if config.lora_path and config.lora_path.exists(): logging.info(f"Applying LoRA from {config.lora_path}") @@ -186,35 +191,24 @@ class LightX2VFactory: lora_name = lora_wrapper.load_lora(str(config.lora_path)) lora_wrapper.apply_lora(lora_name, config.lora_strength) logging.info(f"LoRA {lora_name} applied successfully") - + return LightX2VModel(wan_model, config, easydict_config) - + @staticmethod def create_from_paths( - model_dir: Union[str, Path], - model_name: str, - model_type: str = "i2v", - precision: str = "bf16", - device: str = "cuda", - **kwargs + model_dir: Union[str, Path], model_name: str, model_type: str = "i2v", precision: str = "bf16", device: str = "cuda", **kwargs ) -> Dict[str, Any]: """Convenience method to create all components from a model directory.""" model_dir = Path(model_dir) - + # Create model config - model_config = ModelConfig( - model_path=model_dir / model_name, - model_type=model_type, - precision=precision, - device=device, - **kwargs - ) - + model_config = ModelConfig(model_path=model_dir / model_name, model_type=model_type, precision=precision, device=device, **kwargs) + # Create components components = { "model": LightX2VFactory.create_model(model_config), } - + # Try to create encoders if paths exist t5_path = model_dir / "models_t5_umt5-xxl-enc-bf16.pth" if t5_path.exists(): @@ -223,7 +217,7 @@ class LightX2VFactory: dtype=model_config.to_dtype(), device=model_config.to_device(), ) - + clip_path = model_dir / "models_clip_open-clip-xlm-roberta-large-vit-huge-14.pth" if clip_path.exists(): components["clip_encoder"] = LightX2VFactory.create_clip_vision_encoder( @@ -231,7 +225,7 @@ class LightX2VFactory: dtype=torch.float16, # CLIP typically uses fp16 device=model_config.to_device(), ) - + vae_path = model_dir / "Wan2.1_VAE.pth" if vae_path.exists(): components["vae"] = LightX2VFactory.create_vae( @@ -239,5 +233,5 @@ class LightX2VFactory: dtype=torch.float16, # VAE typically uses fp16 device=model_config.to_device(), ) - - return components \ No newline at end of file + + return components diff --git a/lightx2v_refactored/models.py b/lightx2v_nodes/models.py similarity index 74% rename from lightx2v_refactored/models.py rename to lightx2v_nodes/models.py index 9217953..6ba67d4 100644 --- a/lightx2v_refactored/models.py +++ b/lightx2v_nodes/models.py @@ -1,18 +1,18 @@ """Model wrappers for LightX2V components.""" -from typing import Any, Dict, List, Optional, Tuple, Union + +from typing import Any, Dict, List, Optional, Union import torch -from dataclasses import dataclass from abc import ABC, abstractmethod -from .config import EncoderConfig, ModelConfig, VideoConfig +from .config import EncoderConfig, ModelConfig, VideoConfig, TeaCacheConfig class BaseModel(ABC): """Base class for all LightX2V models.""" - + def __init__(self, config: Union[EncoderConfig, ModelConfig]): self.config = config - + @abstractmethod def to(self, device: torch.device) -> "BaseModel": """Move model to device.""" @@ -21,34 +21,27 @@ class BaseModel(ABC): class LightX2VT5Encoder(BaseModel): """Wrapper for T5 text encoder.""" - + def __init__(self, t5_model: Any, config: EncoderConfig): super().__init__(config) self._model = t5_model - + def encode(self, prompts: List[str]) -> Dict[str, torch.Tensor]: """Encode text prompts.""" context = self._model.infer(prompts) return {"context": context} - - def encode_with_negative( - self, - prompt: str, - negative_prompt: Optional[str] = None - ) -> Dict[str, torch.Tensor]: + + def encode_with_negative(self, prompt: str, negative_prompt: Optional[str] = None) -> Dict[str, torch.Tensor]: """Encode prompt with negative prompt.""" context = self._model.infer([prompt]) context_null = self._model.infer([negative_prompt if negative_prompt else ""]) - return { - "context": context, - "context_null": context_null - } - + return {"context": context, "context_null": context_null} + def to(self, device: torch.device) -> "LightX2VT5Encoder": """Move encoder to device.""" # T5 model handles device internally return self - + @property def device(self) -> torch.device: """Get current device.""" @@ -57,42 +50,38 @@ class LightX2VT5Encoder(BaseModel): class LightX2VClipVisionEncoder(BaseModel): """Wrapper for CLIP vision encoder.""" - + def __init__(self, clip_model: Any, config: EncoderConfig): super().__init__(config) self._model = clip_model - - def encode( - self, - images: torch.Tensor, - video_config: Optional[VideoConfig] = None - ) -> torch.Tensor: + + def encode(self, images: torch.Tensor, video_config: Optional[VideoConfig] = None) -> torch.Tensor: """Encode images with CLIP.""" if video_config: # Convert VideoConfig to dict format expected by CLIP config_dict = { - "target_height": video_config.height, - "target_width": video_config.width, - "target_video_length": video_config.num_frames, + "target_height": video_config.target_height, + "target_width": video_config.target_width, + "target_video_length": video_config.target_video_length, "vae_stride": video_config.vae_stride, "patch_size": video_config.patch_size, } else: config_dict = {} - + # Ensure images are in correct format [B, C, T, H, W] if images.dim() == 3: # [C, H, W] images = images.unsqueeze(0).unsqueeze(2) # [1, C, 1, H, W] elif images.dim() == 4: # [B, C, H, W] images = images.unsqueeze(2) # [B, C, 1, H, W] - + return self._model.visual(images, config_dict) - + def to(self, device: torch.device) -> "LightX2VClipVisionEncoder": """Move encoder to device.""" # CLIP model handles device internally return self - + @property def device(self) -> torch.device: """Get current device.""" @@ -101,49 +90,43 @@ class LightX2VClipVisionEncoder(BaseModel): class LightX2VVae(BaseModel): """Wrapper for VAE.""" - + def __init__(self, vae_model: Any, config: EncoderConfig): super().__init__(config) self._model = vae_model - - def encode( - self, - videos: List[torch.Tensor], - video_config: Optional[VideoConfig] = None, - cpu_offload: bool = False - ) -> List[torch.Tensor]: + + def encode(self, videos: List[torch.Tensor], video_config: Optional[VideoConfig] = None, cpu_offload: bool = False) -> List[torch.Tensor]: """Encode videos to latent space.""" config_dict = {"cpu_offload": cpu_offload} - + if video_config: - config_dict.update({ - "target_height": video_config.height, - "target_width": video_config.width, - "target_video_length": video_config.num_frames, - "vae_stride": video_config.vae_stride, - "patch_size": video_config.patch_size, - }) - + config_dict.update( + { + "target_height": video_config.target_height, + "target_width": video_config.target_width, + "target_video_length": video_config.target_video_length, + "vae_stride": video_config.vae_stride, + "patch_size": video_config.patch_size, + } + ) + from easydict import EasyDict + return self._model.encode(videos, EasyDict(config_dict)) - - def decode( - self, - latents: torch.Tensor, - generator: Optional[torch.Generator] = None, - cpu_offload: bool = False - ) -> torch.Tensor: + + def decode(self, latents: torch.Tensor, generator: Optional[torch.Generator] = None, cpu_offload: bool = False) -> torch.Tensor: """Decode latents to video.""" from easydict import EasyDict + config = EasyDict({"cpu_offload": cpu_offload}) - + return self._model.decode(latents, generator=generator, config=config) - + def to(self, device: torch.device) -> "LightX2VVae": """Move VAE to device.""" # VAE model handles device internally return self - + @property def device(self) -> torch.device: """Get current device.""" @@ -152,22 +135,22 @@ class LightX2VVae(BaseModel): class LightX2VModel(BaseModel): """Wrapper for main LightX2V model.""" - + def __init__(self, wan_model: Any, config: ModelConfig, easydict_config: Any): super().__init__(config) self._model = wan_model self._easydict_config = easydict_config self._scheduler = None - + def set_scheduler(self, scheduler: Any): """Set the scheduler for the model.""" self._scheduler = scheduler self._model.set_scheduler(scheduler) - + def infer(self, inputs: Dict[str, Any]): """Run inference.""" return self._model.infer(inputs) - + def prepare_inputs( self, text_embeddings: Dict[str, torch.Tensor], @@ -179,18 +162,18 @@ class LightX2VModel(BaseModel): "image_encoder_output": image_embeddings or {}, } return inputs - + def to(self, device: torch.device) -> "LightX2VModel": """Move model to device.""" # Model handles device internally return self - + @property def device(self) -> torch.device: """Get current device.""" return self.config.to_device() - + @property def easydict_config(self) -> Any: """Get EasyDict config for compatibility.""" - return self._easydict_config \ No newline at end of file + return self._easydict_config diff --git a/lightx2v_refactored/nodes.py b/lightx2v_nodes/nodes.py similarity index 97% rename from lightx2v_refactored/nodes.py rename to lightx2v_nodes/nodes.py index 8eef694..bcb82b6 100644 --- a/lightx2v_refactored/nodes.py +++ b/lightx2v_nodes/nodes.py @@ -16,12 +16,19 @@ from comfy.utils import ProgressBar from .config import LightX2VConfig, ModelConfig, VideoConfig, TeaCacheConfig from .factory import LightX2VFactory -from .models import LightX2VT5Encoder, LightX2VClipVisionEncoder, LightX2VVae, LightX2VModel +from .models import ( + LightX2VT5Encoder, + LightX2VClipVisionEncoder, + LightX2VVae, + LightX2VModel, +) # Import original LightX2V modules from ..lightx2v.lightx2v.utils.profiler import ProfilingContext from ..lightx2v.lightx2v.models.schedulers.wan.scheduler import WanScheduler -from ..lightx2v.lightx2v.models.schedulers.wan.feature_caching.scheduler import WanSchedulerTeaCaching +from ..lightx2v.lightx2v.models.schedulers.wan.feature_caching.scheduler import ( + WanSchedulerTeaCaching, +) # Coefficient values for TeaCache @@ -503,9 +510,9 @@ class Lightx2vWanVideoImageEncoder(BaseNode): # Create video configuration video_config = VideoConfig( - width=width, - height=height, - num_frames=num_frames, + target_width=width, + target_height=height, + target_video_length=num_frames, ) # Convert image format @@ -584,7 +591,10 @@ class Lightx2vWanVideoEmptyEmbeds(BaseNode): "required": { "width": ("INT", {"default": 832, "min": 64, "max": 2048, "step": 8}), "height": ("INT", {"default": 480, "min": 64, "max": 2048, "step": 8}), - "num_frames": ("INT", {"default": 81, "min": 1, "max": 10000, "step": 4}), + "num_frames": ( + "INT", + {"default": 81, "min": 1, "max": 10000, "step": 4}, + ), } } @@ -595,9 +605,9 @@ class Lightx2vWanVideoEmptyEmbeds(BaseNode): def process(self, num_frames: int, width: int, height: int) -> Tuple[Dict[str, Any]]: """Create empty image embeddings for T2V.""" video_config = VideoConfig( - width=width, - height=height, - num_frames=num_frames, + target_width=width, + target_height=height, + target_video_length=num_frames, ) return ({"config": video_config.to_easydict()},) @@ -683,6 +693,7 @@ class Lightx2vWanVideoModelLoader(BaseNode): lora_path=Path(lora_path) if lora_path and lora_path.strip() else None, lora_strength=lora_strength, mm_config=mm_config, + feature_caching="Tea" if teacache_args else "NoCaching", ) # Create model @@ -691,7 +702,6 @@ class Lightx2vWanVideoModelLoader(BaseNode): # Add TeaCache config if provided easydict_config = model.easydict_config if teacache_args: - easydict_config.feature_caching = "Tea" easydict_config.teacache_thresh = teacache_args.rel_l1_thresh easydict_config.use_ret_steps = teacache_args.use_ret_steps easydict_config.coefficients = teacache_args.coefficients @@ -721,7 +731,10 @@ class Lightx2vWanVideoSampler(BaseNode): "FLOAT", {"default": 5, "min": 1, "max": 20.0, "step": 0.1}, ), - "seed": ("INT", {"default": 42, "min": 0, "max": 0xFFFFFFFFFFFFFFFF, "step": 1}), + "seed": ( + "INT", + {"default": 42, "min": 0, "max": 0xFFFFFFFFFFFFFFFF, "step": 1}, + ), } } diff --git a/nodes.py b/nodes.py index 8e3d8b7..2fc52d2 100644 --- a/nodes.py +++ b/nodes.py @@ -1,908 +1,8 @@ -import os -import torch -import gc -from typing import cast, Any, no_type_check -import logging -import json -import numpy as np -import comfy.model_management as comfy_mm -from comfy.utils import ProgressBar -from pathlib import Path -import math - -# import folder_paths -from tqdm import tqdm -from easydict import EasyDict - - -# Import LightX2V modules -from .lightx2v.lightx2v.utils.profiler import ProfilingContext -from .lightx2v.lightx2v.models.input_encoders.hf.t5.model import T5EncoderModel -from .lightx2v.lightx2v.models.input_encoders.hf.xlm_roberta.model import ( - CLIPModel as ClipVisionModel, -) -from .lightx2v.lightx2v.models.video_encoders.hf.wan.vae import WanVAE -from .lightx2v.lightx2v.models.networks.wan.model import WanModel -from .lightx2v.lightx2v.models.networks.wan.lora_adapter import WanLoraWrapper -from .lightx2v.lightx2v.models.schedulers.wan.scheduler import WanScheduler -from .lightx2v.lightx2v.models.schedulers.wan.feature_caching.scheduler import ( - WanSchedulerTeaCaching, +# Import refactored modules +from .lightx2v_nodes.nodes import ( + NODE_CLASS_MAPPINGS, + NODE_DISPLAY_NAME_MAPPINGS, ) -from .lightx2v.lightx2v.common.ops import * # noqa: F401, F403 for import global register - - -class LightX2VEncoderFactory: - """编码器工厂类,统一管理所有编码器的创建""" - - @staticmethod - def create_t5_encoder(model_path, tokenizer_path, dtype, device, cpu_offload=False): - return T5EncoderModel( - text_len=512, - dtype=dtype, - device=device, - checkpoint_path=model_path, - tokenizer_path=tokenizer_path, - shard_fn=None, - cpu_offload=cpu_offload, - ) - - @staticmethod - def create_clip_vision_encoder(model_path, dtype, device): - return ClipVisionModel( - dtype=dtype, - device=device, - checkpoint_path=model_path, - clip_quantized=False, - clip_quantized_ckpt=None, - quant_scheme=None, - ) - - @staticmethod - def create_vae(model_path, dtype, device, parallel=False): - return WanVAE( - z_dim=16, - vae_pth=str(model_path), - dtype=dtype, - device=device, - parallel=parallel, - ) - - -def convert_dtype(dtype_str: str): - dtype_map = { - "bf16": torch.bfloat16, - "fp16": torch.float16, - "fp32": torch.float32, - } - return dtype_map[dtype_str] - - -class WanVideoTeaCache: - @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "rel_l1_thresh": ( - "FLOAT", - { - "default": 0.26, - "min": 0.0, - "max": 10.0, - "step": 0.001, - "tooltip": "Threshold for to determine when to apply the cache, compromise between speed and accuracy. When using coefficients a good value range is something between 0.2-0.4 for all but 1.3B model, which should be about 10 times smaller, same as when not using coefficients.", - }, - ), - "start_percent": ( - "FLOAT", - { - "default": 0.1, - "min": 0.0, - "max": 1.0, - "step": 0.01, - "tooltip": "The start percentage of the steps to use with TeaCache.", - }, - ), - "end_percent": ( - "FLOAT", - { - "default": 1.0, - "min": 0.0, - "max": 1.0, - "step": 0.01, - "tooltip": "The end percentage of the steps to use with TeaCache.", - }, - ), - "cache_device": ( - ["main_device", "offload_device"], - {"default": "offload_device", "tooltip": "Device to cache to"}, - ), - "coefficients": ( - [ - "i2v-14B-720p", - "i2v-14B-480p", - "t2v-1.3B", - "t2v-14B", - ], - { - "default": "i2v-14B-720p", - "tooltip": "Use coefficients for TeaCache. 'i2v-14B-720p' will use the default coefficients", - }, - ), - "use_ret_steps": ("BOOLEAN", {"default": False}), - }, - "optional": { - "mode": ( - ["e", "e0"], - { - "default": "e", - "tooltip": "Choice between using e (time embeds, default) or e0 (modulated time embeds)", - }, - ), - }, - } - - RETURN_TYPES = ("LIGHT_TEACACHEARGS",) - RETURN_NAMES = ("teacache_args",) - FUNCTION = "process" - CATEGORY = "LightX2V" - - EXPERIMENTAL = True - - def process( - self, - rel_l1_thresh: float, - start_percent: float, - end_percent: float, - cache_device: str, - coefficients: str, - use_ret_steps: bool, - mode="e", - ): - if cache_device == "main_device": - teacache_device = comfy_mm.get_torch_device() - else: - teacache_device = comfy_mm.unet_offload_device() - - # use_ret_steps = True is [0] - # use_ret_steps = False is [1] - coeff_values_map = { - "i2v-14B-480p": [ - [2.57151496e05, -3.54229917e04, 1.40286849e03, -1.35890334e01, 1.32517977e-01], - [-3.02331670e02, 2.23948934e02, -5.25463970e01, 5.87348440e00, -2.01973289e-01], - ], - "i2v-14B-720p": [ - [8.10705460e03, 2.13393892e03, -3.72934672e02, 1.66203073e01, -4.17769401e-02], - [-114.36346466, 65.26524496, -18.82220707, 4.91518089, -0.23412683], - ], - "t2v-1.3B": [ - [-5.21862437e04, 9.23041404e03, -5.28275948e02, 1.36987616e01, -4.99875664e-02], - [2.39676752e03, -1.31110545e03, 2.01331979e02, -8.29855975e00, 1.37887774e-01], - ], - "t2v-14B": [ - [-3.03318725e05, 4.90537029e04, -2.65530556e03, 5.87365115e01, -3.15583525e-01], - [-5784.54975374, 5449.50911966, -1811.16591783, 256.27178429, -13.02252404], - ], - } - - teacache_args = { - "rel_l1_thresh": rel_l1_thresh, - "start_percent": start_percent, - "end_percent": end_percent, - "cache_device": teacache_device, - "coefficients": coeff_values_map[coefficients], - "use_ret_steps": use_ret_steps, - "mode": mode, - } - return (teacache_args,) - - -class Lightx2vWanVideoModelDir: - @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "model_dir": ( - "STRING", - {"default": "/mnt/aigc/users/lijiaqi2/wan_model/Wan2.1-I2V-14B-480P"}, - ) - } - } - - RETURN_TYPES = ("STRING",) - RETURN_NAMES = ("STRING",) - FUNCTION = "process" - CATEGORY = "LightX2V" - - def process(self, model_dir): - assert Path(model_dir).exists(), f"Model directory {model_dir} does not exist." - return (model_dir,) - - -class Lightx2vWanVideoT5EncoderLoader: - @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "model_name": ( - "STRING", - {"default": "models_t5_umt5-xxl-enc-bf16.pth"}, - ), - "precision": (["bf16", "fp16", "fp32"], {"default": "bf16"}), - "device": (["cuda", "cpu"], {"default": "cuda"}), - }, - "optional": { - "model_dir": ("STRING", {"default": None}), - }, - } - - RETURN_TYPES = ("LIGHT_T5_ENCODER",) - RETURN_NAMES = ("t5_encoder",) - FUNCTION = "load_t5_encoder" - CATEGORY = "LightX2V" - - def load_t5_encoder( - self, - model_name: str, - precision: str, - device: str, - model_dir: str | None = None, - ): - dtype = convert_dtype(precision) - - if model_dir: - model_path = Path(model_dir) - model_path = model_path / model_name - else: - model_path = Path(model_name) - assert model_path.exists(), f"T5 model path {model_path} does not exist. Please provide a valid model path or set model_dir." - - tokenizer_path = model_path.parent / "google" / "umt5-xxl" - assert tokenizer_path.exists(), f"Tokenizer path {tokenizer_path} does not exist. Please provide a valid tokenizer path or set model_dir." - - if device == "cuda": - init_device = comfy_mm.get_torch_device() - cpu_offload = False - else: - init_device = torch.device("cpu") - cpu_offload = True - - t5_encoder = LightX2VEncoderFactory.create_t5_encoder(model_path, tokenizer_path, dtype, init_device, cpu_offload) - - return (t5_encoder,) - - -class Lightx2vWanVideoT5Encoder: - @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "t5_encoder": ("LIGHT_T5_ENCODER",), - "prompt": ( - "STRING", - { - "multiline": True, - "default": "Summer beach vacation style, a white cat wearing sunglasses sits on a surfboard. The fluffy-furred feline gazes directly at the camera with a relaxed expression. Blurred beach scenery forms the background featuring crystal-clear waters, distant green hills, and a blue sky dotted with white clouds. The cat assumes a naturally relaxed posture, as if savoring the sea breeze and warm sunlight. A close-up shot highlights the feline's intricate details and the refreshing atmosphere of the seaside.", - }, - ), - "negative_prompt": ( - "STRING", - { - "multiline": True, - "default": "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走", - }, - ), - } - } - - RETURN_TYPES = ("LIGHT_TEXT_EMBEDDINGS",) - RETURN_NAMES = ("text_embeddings",) - FUNCTION = "encode_text" - CATEGORY = "LightX2V" - - def encode_text(self, t5_encoder: T5EncoderModel, prompt: str, negative_prompt: str | None = None): - context = t5_encoder.infer([prompt]) - context_null = t5_encoder.infer([negative_prompt if negative_prompt else ""]) - text_embeddings = {"context": context, "context_null": context_null} - return (text_embeddings,) - - -class Lightx2vWanVideoVaeLoader: - @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "model_name": ( - "STRING", - {"default": "Wan2.1_VAE.pth"}, - ), - "precision": (["bf16", "fp16", "fp32"], {"default": "fp16"}), - "device": (["cuda", "cpu"], {"default": "cuda"}), - "parallel": ("BOOLEAN", {"default": False}), - }, - "optional": { - "model_dir": ("STRING", {"default": None}), - }, - } - - RETURN_TYPES = ("LIGHT_WAN_VAE",) - RETURN_NAMES = ("wan_vae",) - FUNCTION = "load_vae" - CATEGORY = "LightX2V" - - def load_vae(self, model_name: str, precision: str, device: str, parallel: bool, model_dir: str | None = None): - dtype = convert_dtype(precision) - - if model_dir: - model_path = Path(model_dir) - model_path = model_path / model_name - else: - model_path = Path(model_name) - - assert model_path.exists(), f"VAE model path {model_path} does not exist. Please provide a valid model path or set model_dir." - - if device == "cuda": - init_device = comfy_mm.get_torch_device() - else: - init_device = torch.device("cpu") - - vae = LightX2VEncoderFactory.create_vae(model_path, dtype, init_device, parallel) - - vae_result = {"vae_cls": vae, "device": device} - return (vae_result,) - - -class Lightx2vWanVideoVaeDecoder: - @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "wan_vae": ("LIGHT_WAN_VAE",), - "latent": ("LIGHT_LATENT",), - } - } - - RETURN_TYPES = ("IMAGE",) - RETURN_NAMES = ("images",) - FUNCTION = "decode_latent" - CATEGORY = "LightX2V" - - def decode_latent(self, wan_vae: dict[str, WanVAE | str], latent: dict[str, Any]): - wan_vae_instance: WanVAE = cast(WanVAE, wan_vae["vae_cls"]) - config = EasyDict({"cpu_offload": wan_vae["device"] == "cpu"}) - - latents = latent["samples"] - generator = latent["generator"] - - with torch.no_grad(): - with ProfilingContext("*decoded images*"): - decoded_images = wan_vae_instance.decode(latents, generator=generator, config=config) - - # 将像素值从 [-1, 1] 归一化到 [0, 1] - images = (decoded_images + 1) / 2 - - # 重新排列维度为ComfyUI标准的图像格式 [T, H, W, C] - # 从 [1, C, T, H, W] 转换为 [T, H, W, C] - images = images.squeeze(0).permute(1, 2, 3, 0).cpu() - - # 确保像素值在有效范围内 - images = torch.clamp(images, 0, 1) - - # 清理缓存以释放GPU内存 - torch.cuda.empty_cache() - gc.collect() - - return (images,) - - -class Lightx2vWanVideoClipVisionEncoderLoader: - @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "model_name": ( - "STRING", - {"default": "models_clip_open-clip-xlm-roberta-large-vit-huge-14.pth"}, - ), - "tokenizer_path": ( - "STRING", - {"default": "xlm-roberta-large"}, - ), - "precision": (["fp16", "fp32"], {"default": "fp16"}), - "device": (["cuda", "cpu"], {"default": "cuda"}), - }, - "optional": { - "model_dir": ("STRING", {"default": None}), - }, - } - - RETURN_TYPES = ("LIGHT_CLIP_VISION_ENCODER",) - RETURN_NAMES = ("clip_vision_encoder",) - FUNCTION = "load_clip_vision_encoder" - CATEGORY = "LightX2V" - - def load_clip_vision_encoder(self, model_name: str, tokenizer_path: str, precision: str, device: str, model_dir: str | None = None): - dtype = convert_dtype(precision) - - if model_dir: - model_path = Path(model_dir) - model_path = model_path / model_name - else: - model_path = Path(model_name) - assert model_path.exists(), f"CLIP model path {model_path} does not exist. Please provide a valid model path or set model_dir." - - if device == "cuda": - init_device = comfy_mm.get_torch_device() - else: - init_device = torch.device("cpu") - - clip_vision_encoder = LightX2VEncoderFactory.create_clip_vision_encoder(model_path, dtype, init_device) - - return (clip_vision_encoder,) - - -class Lightx2vWanVideoImageEncoder: - @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "vae": ("LIGHT_WAN_VAE",), - "clip_vision_encoder": ("LIGHT_CLIP_VISION_ENCODER",), - "image": ("IMAGE",), - "width": ( - "INT", - { - "default": 832, - "min": 64, - "max": 2048, - "step": 8, - "tooltip": "Width of the image to encode", - }, - ), - "height": ( - "INT", - { - "default": 480, - "min": 64, - "max": 29048, - "step": 8, - "tooltip": "Height of the image to encode", - }, - ), - "num_frames": ( - "INT", - { - "default": 81, - "min": 1, - "max": 10000, - "step": 4, - "tooltip": "Number of frames to encode", - }, - ), - } - } - - RETURN_TYPES = ("LIGHT_IMAGE_EMBEDDINGS",) - RETURN_NAMES = ("image_embeddings",) - FUNCTION = "encode_image" - CATEGORY = "LightX2V" - - def encode_image( - self, - vae: dict[str, WanVAE | str], - clip_vision_encoder: ClipVisionModel, - image: torch.Tensor, - width: int, - height: int, - num_frames: int, - ): - vae_instance: WanVAE = cast(WanVAE, vae["vae_cls"]) - - config = EasyDict( - { - "cpu_offload": True if vae["device"] == "cpu" else False, - "target_height": height, - "target_width": width, - "target_video_length": num_frames, - "vae_stride": (4, 8, 8), - "patch_size": (1, 2, 2), - } - ) - # skip lint - config = cast(Any, config) - - # 将图像转换为期望的张量格式 - device = comfy_mm.get_torch_device() - img = image[0].permute(2, 0, 1).to(device) # [C, H, W] - img = img.sub_(0.5).div_(0.5) # 归一化到 [-1, 1] - - # 使用CLIP视觉编码器编码图像 - with ProfilingContext("*clip encoder*"): - clip_encoder_out = clip_vision_encoder.visual([img[:, None, :, :]], config).squeeze(0).to(torch.bfloat16) - - # 计算宽高比和尺寸 - h, w = img.shape[1:] - aspect_ratio = h / w - max_area = config.target_height * config.target_width - lat_h = round(np.sqrt(max_area * aspect_ratio) // config.vae_stride[1] // config.patch_size[1] * config.patch_size[1]) - lat_w = round(np.sqrt(max_area / aspect_ratio) // config.vae_stride[2] // config.patch_size[2] * config.patch_size[2]) - - # XXX: trick - config.lat_h = lat_h - config.lat_w = lat_w - - h = lat_h * config.vae_stride[1] - w = lat_w * config.vae_stride[2] - - msk = torch.ones(1, config.target_video_length, lat_h, lat_w, device=torch.device("cuda")) - msk[:, 1:] = 0 - msk = torch.concat([torch.repeat_interleave(msk[:, 0:1], repeats=4, dim=1), msk[:, 1:]], dim=1) - msk = msk.view(1, msk.shape[1] // 4, 4, lat_h, lat_w) - msk = msk.transpose(1, 2)[0] - with ProfilingContext("*vae encoder*"): - vae_encode_out: torch.Tensor = vae_instance.encode( - [ - torch.concat( - [ - torch.nn.functional.interpolate(img[None].cpu(), size=(h, w), mode="bicubic").transpose(0, 1), - torch.zeros(3, config.target_video_length - 1, h, w), - ], - dim=1, - ).cuda() - ], # type: ignore - config, - )[0] - - vae_encode_out = torch.concat([msk, vae_encode_out]).to(torch.bfloat16) - - image_embeddings = { - "clip_encoder_out": clip_encoder_out, - "vae_encode_out": vae_encode_out, - "config": config, - } - - print(f"Image Encoder Output Shape: {clip_encoder_out.shape}") - print(f"VAE Encoder Output Shape: {vae_encode_out.shape}") - print(f"Latent Height: {lat_h}, Latent Width: {lat_w}") - print(f"Image Shape: {img.shape}") - print(f"Configuration: {config}") - - return (image_embeddings,) - - -class Lightx2vWanVideoEmptyEmbeds: - @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "width": ( - "INT", - { - "default": 832, - "min": 64, - "max": 2048, - "step": 8, - "tooltip": "Width of the image to encode", - }, - ), - "height": ( - "INT", - { - "default": 480, - "min": 64, - "max": 29048, - "step": 8, - "tooltip": "Height of the image to encode", - }, - ), - "num_frames": ( - "INT", - { - "default": 81, - "min": 1, - "max": 10000, - "step": 4, - "tooltip": "Number of frames to encode", - }, - ), - } - } - - RETURN_TYPES = ("LIGHT_IMAGE_EMBEDDINGS",) - RETURN_NAMES = ("image_embeddings",) - FUNCTION = "process" - CATEGORY = "LightX2V" - - def process(self, num_frames: int, width: int, height: int): - config = EasyDict( - { - "target_height": height, - "target_width": width, - "target_video_length": num_frames, - "vae_stride": (4, 8, 8), - "patch_size": (1, 2, 2), - } - ) - - return ({"config": config},) - - -class Lightx2vWanVideoModelLoader: - @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "model_name": ( - "STRING", - {"default": ""}, - ), - "model_type": (["t2v", "i2v"], {"default": "i2v"}), - "precision": (["bf16", "fp16", "fp32"], {"default": "bf16"}), - "device": (["cuda", "cpu"], {"default": "cuda"}), - "attention_type": ( - ["sdpa", "flash_attn2", "flash_attn3"], - {"default": "flash_attn3"}, - ), - "cpu_offload": ("BOOLEAN", {"default": False}), - "offload_granularity": ( - ["block", "phase"], - {"default": "phase"}, - ), - }, - "optional": { - "mm_type": ("STRING", {"default": None}), - "teacache_args": ("LIGHT_TEACACHEARGS", {"default": None}), - "lora_path": ("STRING", {"default": None}), - "lora_strength": ( - "FLOAT", - {"default": 1.0, "min": 0.0, "max": 2.0, "step": 0.01}, - ), - "model_dir": ( - "STRING", - {"default": "/mnt/aigc/users/lijiaqi2/wan_model/Wan2.1-I2V-14B-480P"}, - ), - }, - } - - RETURN_TYPES = ("LIGHT_WAN_MODEL",) - RETURN_NAMES = ("wan_model",) - FUNCTION = "load_model" - CATEGORY = "LightX2V" - - def load_model( - self, - model_name: str, - model_type: str, - precision: str, - device: str, - attention_type: str, - offload_granularity: str, - mm_type: str | None = None, - lora_path: str | None = None, - lora_strength: float = 1.0, - cpu_offload: bool = False, - teacache_args: dict[str, Any] | None = None, - model_dir: str | None = None, - ): - dtype = convert_dtype(precision) - - if device == "cuda": - init_device = comfy_mm.get_torch_device() - else: - init_device = torch.device("cpu") - - if model_dir: - model_path = Path(model_dir) / model_name - else: - model_path = Path(model_name) - assert model_path.exists(), f"Model path {model_path} does not exist. Please provide a valid model path or set model_dir." - - if model_path.is_dir(): - config_json_path = model_path / "config.json" - else: - config_json_path = model_path.parent / "config.json" - - config_json = {} - if config_json_path.exists(): - with open(config_json_path, "r") as f: - config_json = json.load(f) - else: - logging.error(f"Config file not found at {config_json_path}") - raise FileNotFoundError(f"Config file not found at {config_json_path}") - - feature_caching = "Tea" if teacache_args is not None else "NoCaching" - teacache_thresh = teacache_args["rel_l1_thresh"] if teacache_args else 0.26 - use_ret_steps = teacache_args["use_ret_steps"] if teacache_args else False - coefficients = teacache_args["coefficients"] if teacache_args else [] - - mm_config = {} - try: - if mm_type: - mm_config = json.loads(mm_type) - except Exception as e: - logging.error(f"Invalid mm_type config {mm_type}, error:{e}") - mm_config = {} - - # 创建配置字典 - config = { - "do_mm_calib": False, - "cpu_offload": cpu_offload, - "parallel_attn_type": None, # [None, "ulysses", "ring"] - "parallel_vae": False, - "max_area": False, - "vae_stride": (4, 8, 8), - "patch_size": (1, 2, 2), - "feature_caching": feature_caching, # ["NoCaching", "TaylorSeer", "Tea"] - "teacache_thresh": teacache_thresh, - "use_ret_steps": use_ret_steps, - "coefficients": coefficients, - "use_bfloat16": dtype == torch.bfloat16, - "mm_config": mm_config, - "model_path": str(model_path), - "task": model_type, - "model_cls": "wan2.1", - "device": init_device, - "attention_type": attention_type, - "lora_path": lora_path if lora_path and lora_path.strip() else None, - "strength_model": lora_strength, - "offload_granularity": offload_granularity, - } - - config.update(**config_json) - config = EasyDict(config) # NOTE(xxx): adapt to Lightx2v - logging.info(f"Loaded config:\n {config}") - logging.info(f"Loading WanModel from {model_path} with type {model_type}") - model = WanModel(model_path, config, init_device) - - # 如果指定了LoRA路径,应用LoRA - if lora_path and os.path.exists(lora_path): - logging.info(f"Applying LoRA from {lora_path} with strength {lora_strength}") - lora_wrapper = WanLoraWrapper(model) - lora_name = lora_wrapper.load_lora(lora_path) - lora_wrapper.apply_lora(lora_name, lora_strength) - logging.info(f"LoRA {lora_name} applied successfully") - - wan_model = {"wan_model": model, "config": config} - - return (wan_model,) - - -class Lightx2vWanVideoSampler: - @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "model": ("LIGHT_WAN_MODEL",), - "text_embeddings": ("LIGHT_TEXT_EMBEDDINGS",), - "image_embeddings": ("LIGHT_IMAGE_EMBEDDINGS",), - "steps": ("INT", {"default": 20, "min": 1, "max": 100, "step": 1}), - "shift": ("FLOAT", {"default": 5.0}), - "cfg_scale": ( - "FLOAT", - {"default": 5, "min": 1, "max": 20.0, "step": 0.1}, - ), - "seed": ("INT", {"default": 42, "min": 0, "max": 0xFFFFFFFFFFFFFFFF, "step": 1}), - } - } - - RETURN_TYPES = ("LIGHT_LATENT",) - RETURN_NAMES = ("latent",) - FUNCTION = "sample" - CATEGORY = "LightX2V" - - @no_type_check - def sample( - self, - model: dict[str, WanModel | dict[str, Any]], - text_embeddings: dict[str, Any], - steps: int, - shift: float, - cfg_scale: float, - seed: int, - image_embeddings: dict[str, Any], - ): - model_config = cast(EasyDict, model.get("config")) - model_config.update(image_embeddings.get("config", {})) - - wan_model = cast(WanModel, model.get("wan_model")) - clip_encoder_out = image_embeddings.get("clip_encoder_out", None) - vae_encode_out = image_embeddings.get("vae_encode_out", None) - - if model_config.task == "i2v" and (clip_encoder_out is None or vae_encode_out is None): # type: ignore - raise ValueError("clip_encoder_out must be provided for i2v task") - - model_config.infer_steps = steps - model_config.sample_shift = shift - model_config.sample_guide_scale = cfg_scale - model_config.seed = seed - - model_config.enable_cfg = False if math.isclose(cfg_scale, 1.0) else True - logging.info(f"Loaded update config:\n {model_config}") - - # wan_runner.set_target_shape - num_channels_latents = model_config.get("num_channels_latents", 16) - - if model_config.task == "i2v": - model_config.target_shape = ( - num_channels_latents, - (model_config.target_video_length - 1) // model_config.vae_stride[0] + 1, - model_config.lat_h, - model_config.lat_w, - ) - elif model_config.task == "t2v": # type: ignore - model_config.target_shape = ( - 16, - (model_config.target_video_length - 1) // 4 + 1, - int(model_config.target_height) // model_config.vae_stride[1], - int(model_config.target_width) // model_config.vae_stride[2], - ) - - # wan_runner.init_scheduler - if model_config.feature_caching == "NoCaching": - scheduler = WanScheduler(model_config) - elif model_config.feature_caching == "Tea": - scheduler = WanSchedulerTeaCaching(model_config) - else: - raise NotImplementedError( - f"Unsupported feature_caching type: {model_config.feature_caching}" # type:ignore - ) - - # setup scheduler - wan_model.set_scheduler(scheduler) - - # Set up inputs - inputs = { - "text_encoder_output": text_embeddings, - "image_encoder_output": image_embeddings, - } - - # Prepare for sampling - scheduler.prepare(inputs.get("image_encoder_output")) - - # Run sampling - progress = ProgressBar(steps) - for step_index in tqdm(range(scheduler.infer_steps), desc="inference", unit="step"): - scheduler.step_pre(step_index=step_index) - with ProfilingContext("model.infer"): - wan_model.infer(inputs) - scheduler.step_post() - - progress.update(1) - - latents, generator = scheduler.latents, scheduler.generator - scheduler.clear() - del inputs, scheduler, text_embeddings, image_embeddings - torch.cuda.empty_cache() - - return ({"samples": latents, "generator": generator},) - - -# Register the nodes -NODE_CLASS_MAPPINGS = { - "Lightx2vWanVideoModelDir": Lightx2vWanVideoModelDir, - "Lightx2vWanVideoT5EncoderLoader": Lightx2vWanVideoT5EncoderLoader, - "Lightx2vWanVideoT5Encoder": Lightx2vWanVideoT5Encoder, - "Lightx2vWanVideoClipVisionEncoderLoader": Lightx2vWanVideoClipVisionEncoderLoader, - "Lightx2vWanVideoVaeLoader": Lightx2vWanVideoVaeLoader, - "Lightx2vTeaCache": WanVideoTeaCache, - "Lightx2vWanVideoEmptyEmbeds": Lightx2vWanVideoEmptyEmbeds, - "Lightx2vWanVideoImageEncoder": Lightx2vWanVideoImageEncoder, - "Lightx2vWanVideoVaeDecoder": Lightx2vWanVideoVaeDecoder, - "Lightx2vWanVideoModelLoader": Lightx2vWanVideoModelLoader, - "Lightx2vWanVideoSampler": Lightx2vWanVideoSampler, -} - -NODE_DISPLAY_NAME_MAPPINGS = { - "Lightx2vWanVideoModelDir": "LightX2V WAN Model Directory", - "Lightx2vWanVideoT5EncoderLoader": "LightX2V WAN T5 Encoder Loader", - "Lightx2vWanVideoT5Encoder": "LightX2V WAN T5 Encoder", - "Lightx2vWanVideoClipVisionEncoderLoader": "LightX2V WAN CLIP Vision Encoder Loader", - "Lightx2vWanVideoClipVisionEncoder": "LightX2V WAN CLIP Vision Encoder", - "Lightx2vWanVideoVaeLoader": "LightX2V WAN VAE Loader", - "Lightx2vWanVideoImageEncoder": "LightX2V WAN Image Encoder", - "Lightx2vWanVideoVaeDecoder": "LightX2V WAN VAE Decoder", - "Lightx2vWanVideoModelLoader": "LightX2V WAN Model Loader", - "Lightx2vWanVideoSampler": "LightX2V WAN Video Sampler", - "Lightx2vTeaCache": "LightX2V WAN Tea Cache", - "Lightx2vWanVideoEmptyEmbeds": "LightX2V WAN Video Empty Embeds", -} +# Export the mappings +__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"] diff --git a/nodes_refactored.py b/nodes_refactored.py deleted file mode 100644 index e9be2e4..0000000 --- a/nodes_refactored.py +++ /dev/null @@ -1,11 +0,0 @@ -""" -Refactored nodes.py that uses the new modular structure while maintaining backward compatibility. -""" -# Import refactored modules -from .lightx2v_refactored.nodes import ( - NODE_CLASS_MAPPINGS, - NODE_DISPLAY_NAME_MAPPINGS, -) - -# Export the mappings -__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"] \ No newline at end of file