refactor: consolidate and refactor LightX2V nodes into a new modular structure, enhancing organization and maintainability
This commit is contained in:
+203
-202
@@ -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},
|
||||
),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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"]
|
||||
|
||||
@@ -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"]
|
||||
Reference in New Issue
Block a user