194 lines
5.2 KiB
Python
194 lines
5.2 KiB
Python
"""Data models for LightX2V ComfyUI wrapper."""
|
|
|
|
from dataclasses import dataclass, field
|
|
from typing import Any, Dict, List, Optional, Union
|
|
|
|
import numpy as np
|
|
import torch
|
|
|
|
|
|
@dataclass
|
|
class TalkObject:
|
|
"""Single talk object containing audio and optional mask."""
|
|
|
|
name: str
|
|
audio: Optional[Union[str, Dict[str, Any], torch.Tensor, np.ndarray]] = None
|
|
mask: Optional[Union[str, torch.Tensor, np.ndarray]] = None
|
|
|
|
def to_dict(self) -> Dict[str, Any]:
|
|
"""Convert to dictionary for pipeline."""
|
|
result = {"name": self.name}
|
|
|
|
if isinstance(self.audio, str):
|
|
result["audio"] = self.audio
|
|
elif self.audio is not None:
|
|
result["audio_data"] = self.audio
|
|
|
|
if isinstance(self.mask, str):
|
|
result["mask"] = self.mask
|
|
elif self.mask is not None:
|
|
result["mask_data"] = self.mask
|
|
|
|
return result
|
|
|
|
|
|
@dataclass
|
|
class InferenceConfig:
|
|
"""Basic inference configuration."""
|
|
|
|
model_cls: str = "wan2.1"
|
|
model_path: str = ""
|
|
task: str = "i2v"
|
|
infer_steps: int = 4
|
|
seed: int = 42
|
|
cfg_scale: float = 5.0
|
|
cfg_scale2: float = 5.0
|
|
sample_shift: int = 5
|
|
height: int = 1280
|
|
width: int = 720
|
|
video_length: int = 81
|
|
fps: int = 16
|
|
video_duration: float = 5.0
|
|
attention_type: str = "torch_sdpa"
|
|
use_31_block: bool = True
|
|
|
|
# Optional parameters
|
|
denoising_step_list: Optional[List[float]] = None
|
|
resize_mode: str = "adaptive"
|
|
fixed_area: str = "720p"
|
|
segment_length: int = 81
|
|
prev_frame_length: int = 5
|
|
use_tiny_vae: bool = False
|
|
|
|
# Runtime parameters
|
|
prompt: str = ""
|
|
negative_prompt: str = ""
|
|
image_path: Optional[str] = None
|
|
audio_path: Optional[str] = None
|
|
|
|
def to_dict(self) -> Dict[str, Any]:
|
|
"""Convert to dictionary, excluding None values."""
|
|
result = {}
|
|
for key, value in self.__dict__.items():
|
|
if value is not None:
|
|
result[key] = value
|
|
return result
|
|
|
|
|
|
@dataclass
|
|
class TeaCacheConfig:
|
|
"""TeaCache configuration."""
|
|
|
|
enable: bool = False
|
|
threshold: float = 0.26
|
|
use_ret_steps: bool = False
|
|
|
|
def to_dict(self) -> Dict[str, Any]:
|
|
return {
|
|
"enable": self.enable,
|
|
"threshold": self.threshold,
|
|
"use_ret_steps": self.use_ret_steps,
|
|
}
|
|
|
|
|
|
@dataclass
|
|
class QuantizationConfig:
|
|
"""Quantization configuration."""
|
|
|
|
dit_quant_scheme: str = "Default"
|
|
t5_quant_scheme: str = "Default"
|
|
clip_quant_scheme: str = "Default"
|
|
adapter_quant_scheme: str = "Default"
|
|
|
|
def to_dict(self) -> Dict[str, Any]:
|
|
return {
|
|
"dit_quant_scheme": self.dit_quant_scheme,
|
|
"t5_quant_scheme": self.t5_quant_scheme,
|
|
"clip_quant_scheme": self.clip_quant_scheme,
|
|
"adapter_quant_scheme": self.adapter_quant_scheme,
|
|
}
|
|
|
|
|
|
@dataclass
|
|
class MemoryOptimizationConfig:
|
|
"""Memory optimization configuration."""
|
|
|
|
enable_rotary_chunk: bool = False
|
|
rotary_chunk_size: int = 100
|
|
clean_cuda_cache: bool = False
|
|
cpu_offload: bool = True
|
|
offload_granularity: str = "block"
|
|
offload_ratio: float = 1.0
|
|
t5_cpu_offload: bool = True
|
|
t5_offload_granularity: str = "model"
|
|
audio_encoder_cpu_offload: bool = True
|
|
audio_adapter_cpu_offload: bool = True
|
|
vae_cpu_offload: bool = True
|
|
use_tiling_vae: bool = True
|
|
lazy_load: bool = False
|
|
unload_after_inference: bool = False
|
|
|
|
def to_dict(self) -> Dict[str, Any]:
|
|
return self.__dict__.copy()
|
|
|
|
|
|
@dataclass
|
|
class LoRAConfig:
|
|
"""LoRA configuration."""
|
|
|
|
path: str
|
|
strength: float = 1.0
|
|
|
|
def to_dict(self) -> Dict[str, Any]:
|
|
return {"path": self.path, "strength": self.strength}
|
|
|
|
|
|
@dataclass
|
|
class TalkObjectsConfig:
|
|
talk_objects: List[TalkObject] = field(default_factory=list)
|
|
|
|
def add_object(self, talk_object: TalkObject):
|
|
self.talk_objects.append(talk_object)
|
|
|
|
def to_dict(self) -> Dict[str, Any]:
|
|
return {"talk_objects": [obj.to_dict() for obj in self.talk_objects]}
|
|
|
|
def to_list(self) -> List[Dict[str, Any]]:
|
|
return [obj.to_dict() for obj in self.talk_objects]
|
|
|
|
|
|
@dataclass
|
|
class CombinedConfig:
|
|
"""Combined configuration for all modules."""
|
|
|
|
inference: Optional[InferenceConfig] = None
|
|
teacache: Optional[TeaCacheConfig] = None
|
|
quantization: Optional[QuantizationConfig] = None
|
|
memory: Optional[MemoryOptimizationConfig] = None
|
|
lora_configs: List[LoRAConfig] = field(default_factory=list)
|
|
talk_objects: Optional[TalkObjectsConfig] = None
|
|
|
|
def to_dict(self) -> Dict[str, Any]:
|
|
"""Convert to dictionary for pipeline."""
|
|
result = {}
|
|
|
|
if self.inference:
|
|
result.update(self.inference.to_dict())
|
|
|
|
if self.teacache:
|
|
result["teacache"] = self.teacache.to_dict()
|
|
|
|
if self.quantization:
|
|
result["quantization"] = self.quantization.to_dict()
|
|
|
|
if self.memory:
|
|
result["memory"] = self.memory.to_dict()
|
|
|
|
if self.lora_configs:
|
|
result["lora_configs"] = [lora.to_dict() for lora in self.lora_configs]
|
|
|
|
if self.talk_objects:
|
|
result["talk_objects"] = self.talk_objects.to_list()
|
|
|
|
return result
|