Files
ModelTC-ComfyUI-Lightx2vWra…/data_models.py
T

196 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[int]] = 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."""
quant_op: str = "none"
dit_quant_scheme: str = "bf16"
t5_quant_scheme: str = "bf16"
clip_quant_scheme: str = "fp16"
adapter_quant_scheme: str = "bf16"
def to_dict(self) -> Dict[str, Any]:
return {
"quant_op": self.quant_op,
"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