refactor: consolidate and refactor LightX2V nodes into a new modular structure, enhancing organization and maintainability

This commit is contained in:
GACLove
2025-07-14 15:52:43 +08:00
parent e234343952
commit b1b663c439
8 changed files with 376 additions and 1288 deletions
+203 -202
View File
@@ -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"
@@ -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",
]
]
@@ -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)
return EasyDict(config_dict)
@@ -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
return components
@@ -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
return self._easydict_config
@@ -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},
),
}
}
+6 -906
View File
@@ -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"]
-11
View File
@@ -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"]