refactor: enhance code readability by formatting multi-line parameters and comments; improve configuration handling in various classes for better maintainability
This commit is contained in:
@@ -42,7 +42,6 @@ def is_module_installed(module_name):
|
||||
|
||||
|
||||
def get_available_ops(op_mapping):
|
||||
"""通用的操作可用性检查函数"""
|
||||
available_ops = []
|
||||
for op_name, module_name in op_mapping.items():
|
||||
is_available = is_module_installed(module_name)
|
||||
@@ -51,13 +50,20 @@ def get_available_ops(op_mapping):
|
||||
|
||||
|
||||
def get_available_quant_ops():
|
||||
quant_mapping = {"sgl": "sgl_kernel", "vllm": "vllm", "q8f": "q8_kernels", "torchao": "torchao"}
|
||||
quant_mapping = {
|
||||
"sgl": "sgl_kernel",
|
||||
"vllm": "vllm",
|
||||
"q8f": "q8_kernels",
|
||||
"torchao": "torchao",
|
||||
}
|
||||
|
||||
available_ops = get_available_ops(quant_mapping)
|
||||
|
||||
# Ada架构GPU优先使用q8f
|
||||
# Prefer q8f for Ada architecture GPUs
|
||||
if is_ada_architecture_gpu():
|
||||
q8f_available = next((op for op in available_ops if op[0] == "q8f" and op[1]), None)
|
||||
q8f_available = next(
|
||||
(op for op in available_ops if op[0] == "q8f" and op[1]), None
|
||||
)
|
||||
if q8f_available:
|
||||
available_ops.remove(q8f_available)
|
||||
available_ops.insert(0, q8f_available)
|
||||
@@ -66,7 +72,12 @@ def get_available_quant_ops():
|
||||
|
||||
|
||||
def get_available_attn_ops():
|
||||
attn_mapping = {"sage_attn2": "sageattention", "flash_attn3": "flash_attn_interface", "flash_attn2": "flash_attn", "torch_sdpa": "torch"}
|
||||
attn_mapping = {
|
||||
"sage_attn2": "sageattention",
|
||||
"flash_attn3": "flash_attn_interface",
|
||||
"flash_attn2": "flash_attn",
|
||||
"torch_sdpa": "torch",
|
||||
}
|
||||
|
||||
return get_available_ops(attn_mapping)
|
||||
|
||||
@@ -74,10 +85,21 @@ def get_available_attn_ops():
|
||||
class LightX2VDefaultConfig:
|
||||
"""Central default configuration for LightX2V."""
|
||||
|
||||
# 分组常量
|
||||
DEFAULT_ATTENTION_TYPE = "flash_attn3"
|
||||
DEFAULT_QUANTIZATION_SCHEMES = {"dit": "bf16", "t5": "bf16", "clip": "fp16", "adapter": "bf16"}
|
||||
DEFAULT_VIDEO_PARAMS = {"height": 480, "width": 832, "length": 81, "fps": 16, "vae_stride": [4, 8, 8], "patch_size": [1, 2, 2]}
|
||||
DEFAULT_QUANTIZATION_SCHEMES = {
|
||||
"dit": "bf16",
|
||||
"t5": "bf16",
|
||||
"clip": "fp16",
|
||||
"adapter": "bf16",
|
||||
}
|
||||
DEFAULT_VIDEO_PARAMS = {
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"length": 81,
|
||||
"fps": 16,
|
||||
"vae_stride": [4, 8, 8],
|
||||
"patch_size": [1, 2, 2],
|
||||
}
|
||||
|
||||
DEFAULT_CONFIG = {
|
||||
# Model Configuration
|
||||
@@ -246,8 +268,9 @@ class ModularConfigManager:
|
||||
self._available_attn_ops = None
|
||||
self._available_quant_ops = None
|
||||
|
||||
def _get_available_ops(self, ops_list: List[Tuple[str, bool]], fallback: str = None) -> List[str]:
|
||||
"""从操作列表中提取可用的操作"""
|
||||
def _get_available_ops(
|
||||
self, ops_list: List[Tuple[str, bool]], fallback: str = None
|
||||
) -> List[str]:
|
||||
available = [op_name for op_name, is_available in ops_list if is_available]
|
||||
if fallback and fallback not in available:
|
||||
available.append(fallback)
|
||||
@@ -267,8 +290,9 @@ class ModularConfigManager:
|
||||
self._available_quant_ops = get_available_quant_ops()
|
||||
return self._get_available_ops(self._available_quant_ops)
|
||||
|
||||
def _update_from_config(self, updates: Dict, config: Dict, mappings: Dict[str, str]) -> None:
|
||||
"""通用配置更新方法"""
|
||||
def _update_from_config(
|
||||
self, updates: Dict, config: Dict, mappings: Dict[str, str]
|
||||
) -> None:
|
||||
for config_key, update_key in mappings.items():
|
||||
if config_key in config:
|
||||
if config_key == "seed" and config[config_key] == -1:
|
||||
@@ -279,7 +303,6 @@ class ModularConfigManager:
|
||||
"""Apply basic inference configuration."""
|
||||
updates = {}
|
||||
|
||||
# 基础映射配置
|
||||
basic_mappings = {
|
||||
"model_cls": "model_cls",
|
||||
"model_path": "model_path",
|
||||
@@ -311,17 +334,31 @@ class ModularConfigManager:
|
||||
if "wan2.2" in config["model_cls"]:
|
||||
updates["use_image_encoder"] = False
|
||||
|
||||
attention_type = config.get("attention_type", LightX2VDefaultConfig.DEFAULT_ATTENTION_TYPE)
|
||||
for attn_key in ["attention_type", "self_attn_1_type", "cross_attn_1_type", "cross_attn_2_type"]:
|
||||
attention_type = config.get(
|
||||
"attention_type", LightX2VDefaultConfig.DEFAULT_ATTENTION_TYPE
|
||||
)
|
||||
for attn_key in [
|
||||
"attention_type",
|
||||
"self_attn_1_type",
|
||||
"cross_attn_1_type",
|
||||
"cross_attn_2_type",
|
||||
]:
|
||||
updates[attn_key] = attention_type
|
||||
|
||||
# TinyVAE配置
|
||||
if config.get("use_tiny_vae", False):
|
||||
updates.update({"use_tiny_vae": True, "tiny_vae": True, "tiny_vae_path": os.path.join(config["model_path"], "taew2_1.pth")})
|
||||
updates.update(
|
||||
{
|
||||
"use_tiny_vae": True,
|
||||
"tiny_vae": True,
|
||||
"tiny_vae_path": os.path.join(config["model_path"], "taew2_1.pth"),
|
||||
}
|
||||
)
|
||||
|
||||
return updates
|
||||
|
||||
def apply_teacache_config(self, config: Dict[str, Any], model_info: Dict[str, Any]) -> Dict[str, Any]:
|
||||
def apply_teacache_config(
|
||||
self, config: Dict[str, Any], model_info: Dict[str, Any]
|
||||
) -> Dict[str, Any]:
|
||||
"""Apply TeaCache configuration."""
|
||||
updates = {}
|
||||
|
||||
@@ -337,7 +374,9 @@ class ModularConfigManager:
|
||||
model_info.get("target_height", 480),
|
||||
)
|
||||
|
||||
coeffs = CoefficientCalculator.get_coefficients(task, model_size, resolution, updates["use_ret_steps"])
|
||||
coeffs = CoefficientCalculator.get_coefficients(
|
||||
task, model_size, resolution, updates["use_ret_steps"]
|
||||
)
|
||||
updates["coefficients"] = coeffs
|
||||
else:
|
||||
updates["feature_caching"] = "NoCaching"
|
||||
@@ -345,7 +384,6 @@ class ModularConfigManager:
|
||||
return updates
|
||||
|
||||
def _get_mm_type(self, dit_scheme: str, quant_backend: str) -> str:
|
||||
"""获取mm_type配置"""
|
||||
if dit_scheme == "bf16":
|
||||
return "Default"
|
||||
|
||||
@@ -413,21 +451,29 @@ class ModularConfigManager:
|
||||
}
|
||||
|
||||
for config_key, update_key in direct_mappings.items():
|
||||
updates[update_key] = config.get(config_key, config.get("cpu_offload", False))
|
||||
updates[update_key] = config.get(
|
||||
config_key, config.get("cpu_offload", False)
|
||||
)
|
||||
|
||||
if updates.get("rotary_chunk"):
|
||||
updates["rotary_chunk_size"] = config.get("rotary_chunk_size", 100)
|
||||
|
||||
if updates.get("cpu_offload"):
|
||||
updates.update({"offload_granularity": config.get("offload_granularity", "phase"), "offload_ratio": config.get("offload_ratio", 1.0)})
|
||||
updates.update(
|
||||
{
|
||||
"offload_granularity": config.get("offload_granularity", "phase"),
|
||||
"offload_ratio": config.get("offload_ratio", 1.0),
|
||||
}
|
||||
)
|
||||
|
||||
if updates.get("t5_cpu_offload"):
|
||||
updates["t5_offload_granularity"] = config.get("t5_offload_granularity", "model")
|
||||
updates["t5_offload_granularity"] = config.get(
|
||||
"t5_offload_granularity", "model"
|
||||
)
|
||||
|
||||
return updates
|
||||
|
||||
def _load_model_config(self, model_path: str) -> Dict[str, Any]:
|
||||
"""加载模型配置文件"""
|
||||
config_path = os.path.join(model_path, "config.json")
|
||||
if not os.path.exists(config_path):
|
||||
return {}
|
||||
@@ -454,7 +500,9 @@ class ModularConfigManager:
|
||||
final_config.update(memory_updates)
|
||||
|
||||
if "teacache" in configs:
|
||||
teacache_updates = self.apply_teacache_config(configs["teacache"], final_config)
|
||||
teacache_updates = self.apply_teacache_config(
|
||||
configs["teacache"], final_config
|
||||
)
|
||||
final_config.update(teacache_updates)
|
||||
|
||||
if "quantization" in configs:
|
||||
|
||||
+46
-11
@@ -134,7 +134,9 @@ class InferenceConfigBuilder:
|
||||
|
||||
return config
|
||||
|
||||
def _apply_optional_params(self, config: InferenceConfig, optional_params: Dict[str, Any]):
|
||||
def _apply_optional_params(
|
||||
self, config: InferenceConfig, optional_params: Dict[str, Any]
|
||||
):
|
||||
"""Apply optional parameters to config."""
|
||||
# Handle denoising steps
|
||||
if "denoising_steps" in optional_params:
|
||||
@@ -148,7 +150,13 @@ class InferenceConfigBuilder:
|
||||
logging.warning(f"Invalid denoising steps: {steps_str}")
|
||||
|
||||
# Handle other optional params
|
||||
for param in ["resize_mode", "fixed_area", "segment_length", "prev_frame_length", "use_tiny_vae"]:
|
||||
for param in [
|
||||
"resize_mode",
|
||||
"fixed_area",
|
||||
"segment_length",
|
||||
"prev_frame_length",
|
||||
"use_tiny_vae",
|
||||
]:
|
||||
if param in optional_params:
|
||||
setattr(config, param, optional_params[param])
|
||||
|
||||
@@ -168,7 +176,13 @@ class TalkObjectConfigBuilder:
|
||||
self.mask_handler = MaskFileHandler()
|
||||
self.resolver = ComfyUIFileResolver()
|
||||
|
||||
def build_from_input(self, name: str, audio: Optional[Any] = None, mask: Optional[Any] = None, save_to_input: bool = True) -> TalkObject:
|
||||
def build_from_input(
|
||||
self,
|
||||
name: str,
|
||||
audio: Optional[Any] = None,
|
||||
mask: Optional[Any] = None,
|
||||
save_to_input: bool = True,
|
||||
) -> TalkObject:
|
||||
"""Build talk object from input data."""
|
||||
if audio is None:
|
||||
return None
|
||||
@@ -210,7 +224,12 @@ class TalkObjectConfigBuilder:
|
||||
if not isinstance(obj_data, dict) or "audio" not in obj_data:
|
||||
continue
|
||||
|
||||
talk_obj = TalkObject(name=obj_data.get("name", "unknown"), audio=obj_data["audio"], mask=obj_data.get("mask"), source_type="path")
|
||||
talk_obj = TalkObject(
|
||||
name=obj_data.get("name", "unknown"),
|
||||
audio=obj_data["audio"],
|
||||
mask=obj_data.get("mask"),
|
||||
source_type="path",
|
||||
)
|
||||
config.add_object(talk_obj)
|
||||
|
||||
return config if config.talk_objects else None
|
||||
@@ -218,13 +237,19 @@ class TalkObjectConfigBuilder:
|
||||
except json.JSONDecodeError as e:
|
||||
logging.error(f"Failed to parse JSON: {e}")
|
||||
|
||||
def build_from_files(self, audio_files: str, mask_files: str = "", names: str = "") -> Optional[TalkObjectsConfig]:
|
||||
def build_from_files(
|
||||
self, audio_files: str, mask_files: str = "", names: str = ""
|
||||
) -> Optional[TalkObjectsConfig]:
|
||||
"""Build talk objects configuration from file lists."""
|
||||
audio_list = [f.strip() for f in audio_files.split("\n") if f.strip()]
|
||||
if not audio_list:
|
||||
return None
|
||||
|
||||
mask_list = [f.strip() for f in mask_files.split("\n") if f.strip()] if mask_files else []
|
||||
mask_list = (
|
||||
[f.strip() for f in mask_files.split("\n") if f.strip()]
|
||||
if mask_files
|
||||
else []
|
||||
)
|
||||
name_list = [n.strip() for n in names.split("\n") if n.strip()] if names else []
|
||||
|
||||
config = TalkObjectsConfig()
|
||||
@@ -288,14 +313,18 @@ class ConfigBuilder:
|
||||
# Add LoRA configs
|
||||
if lora_chain:
|
||||
for lora_dict in lora_chain:
|
||||
lora_config = LoRAConfig(path=lora_dict["path"], strength=lora_dict.get("strength", 1.0))
|
||||
lora_config = LoRAConfig(
|
||||
path=lora_dict["path"], strength=lora_dict.get("strength", 1.0)
|
||||
)
|
||||
combined.lora_configs.append(lora_config)
|
||||
|
||||
# Build final config using existing manager
|
||||
configs_dict = {
|
||||
"inference": inference_config.to_dict() if inference_config else {},
|
||||
"teacache": teacache_config.to_dict() if teacache_config else None,
|
||||
"quantization": quantization_config.to_dict() if quantization_config else None,
|
||||
"quantization": quantization_config.to_dict()
|
||||
if quantization_config
|
||||
else None,
|
||||
"memory": memory_config.to_dict() if memory_config else None,
|
||||
}
|
||||
|
||||
@@ -333,8 +362,12 @@ class ConfigBuilder:
|
||||
"offload_ratio": getattr(config, "offload_ratio", None),
|
||||
"t5_cpu_offload": getattr(config, "t5_cpu_offload", False),
|
||||
"t5_offload_granularity": getattr(config, "t5_offload_granularity", None),
|
||||
"audio_encoder_cpu_offload": getattr(config, "audio_encoder_cpu_offload", False),
|
||||
"audio_adapter_cpu_offload": getattr(config, "audio_adapter_cpu_offload", False),
|
||||
"audio_encoder_cpu_offload": getattr(
|
||||
config, "audio_encoder_cpu_offload", False
|
||||
),
|
||||
"audio_adapter_cpu_offload": getattr(
|
||||
config, "audio_adapter_cpu_offload", False
|
||||
),
|
||||
"vae_cpu_offload": getattr(config, "vae_cpu_offload", False),
|
||||
"use_tiling_vae": getattr(config, "use_tiling_vae", False),
|
||||
"unload_after_inference": getattr(config, "unload_after_inference", False),
|
||||
@@ -359,7 +392,9 @@ class LoRAChainBuilder:
|
||||
"""Builder for LoRA chain configurations."""
|
||||
|
||||
@staticmethod
|
||||
def build_chain(lora_name: str, strength: float, existing_chain: Optional[List[Dict]] = None) -> List[Dict]:
|
||||
def build_chain(
|
||||
lora_name: str, strength: float, existing_chain: Optional[List[Dict]] = None
|
||||
) -> List[Dict]:
|
||||
"""Build or extend a LoRA chain."""
|
||||
if existing_chain is None:
|
||||
chain = []
|
||||
|
||||
+5
-1
@@ -86,7 +86,11 @@ class TeaCacheConfig:
|
||||
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}
|
||||
return {
|
||||
"enable": self.enable,
|
||||
"threshold": self.threshold,
|
||||
"use_ret_steps": self.use_ret_steps,
|
||||
}
|
||||
|
||||
|
||||
@dataclass
|
||||
|
||||
+24
-7
@@ -33,7 +33,12 @@ class AudioFileHandler(FileHandler):
|
||||
def __init__(self):
|
||||
self.supported_formats = [".wav", ".mp3", ".flac", ".m4a"]
|
||||
|
||||
def save(self, audio_data: Union[Dict, torch.Tensor, np.ndarray, Tuple], path: str, sample_rate: Optional[int] = None) -> str:
|
||||
def save(
|
||||
self,
|
||||
audio_data: Union[Dict, torch.Tensor, np.ndarray, Tuple],
|
||||
path: str,
|
||||
sample_rate: Optional[int] = None,
|
||||
) -> str:
|
||||
"""Save audio data to file.
|
||||
|
||||
Args:
|
||||
@@ -80,7 +85,9 @@ class AudioFileHandler(FileHandler):
|
||||
sample_rate, waveform = wavfile.read(path)
|
||||
return waveform, sample_rate
|
||||
|
||||
def _extract_audio_data(self, audio_data: Any, sample_rate: Optional[int] = None) -> Tuple[np.ndarray, int]:
|
||||
def _extract_audio_data(
|
||||
self, audio_data: Any, sample_rate: Optional[int] = None
|
||||
) -> Tuple[np.ndarray, int]:
|
||||
"""Extract waveform and sample rate from various audio formats.
|
||||
|
||||
Handles three main sources:
|
||||
@@ -98,7 +105,9 @@ class AudioFileHandler(FileHandler):
|
||||
if isinstance(waveform, torch.Tensor):
|
||||
if waveform.dim() == 3: # [batch, channels, samples]
|
||||
waveform = waveform[0] # Take first batch
|
||||
if waveform.dim() == 2 and waveform.shape[0] <= 2: # [channels, samples]
|
||||
if (
|
||||
waveform.dim() == 2 and waveform.shape[0] <= 2
|
||||
): # [channels, samples]
|
||||
waveform = waveform.transpose(0, 1) # -> [samples, channels]
|
||||
waveform = waveform.cpu().numpy()
|
||||
else:
|
||||
@@ -173,7 +182,9 @@ class ImageFileHandler(FileHandler):
|
||||
def __init__(self):
|
||||
self.supported_formats = [".png", ".jpg", ".jpeg", ".bmp", ".tiff"]
|
||||
|
||||
def save(self, image_data: Union[torch.Tensor, np.ndarray, Image.Image], path: str) -> str:
|
||||
def save(
|
||||
self, image_data: Union[torch.Tensor, np.ndarray, Image.Image], path: str
|
||||
) -> str:
|
||||
"""Save image data to file.
|
||||
|
||||
Args:
|
||||
@@ -254,7 +265,9 @@ class TempFileManager:
|
||||
self.temp_files: List[str] = []
|
||||
|
||||
@contextmanager
|
||||
def temp_file(self, suffix: str = "", prefix: str = "lightx2v_", delete: bool = True):
|
||||
def temp_file(
|
||||
self, suffix: str = "", prefix: str = "lightx2v_", delete: bool = True
|
||||
):
|
||||
"""Context manager for temporary file creation.
|
||||
|
||||
Args:
|
||||
@@ -265,7 +278,9 @@ class TempFileManager:
|
||||
Yields:
|
||||
Path to temporary file
|
||||
"""
|
||||
temp_file = tempfile.NamedTemporaryFile(suffix=suffix, prefix=prefix, delete=False)
|
||||
temp_file = tempfile.NamedTemporaryFile(
|
||||
suffix=suffix, prefix=prefix, delete=False
|
||||
)
|
||||
temp_path = temp_file.name
|
||||
temp_file.close()
|
||||
|
||||
@@ -287,7 +302,9 @@ class TempFileManager:
|
||||
Returns:
|
||||
Path to temporary file
|
||||
"""
|
||||
with tempfile.NamedTemporaryFile(suffix=suffix, prefix=prefix, delete=False) as tmp:
|
||||
with tempfile.NamedTemporaryFile(
|
||||
suffix=suffix, prefix=prefix, delete=False
|
||||
) as tmp:
|
||||
temp_path = tmp.name
|
||||
|
||||
self.temp_files.append(temp_path)
|
||||
|
||||
@@ -150,7 +150,14 @@ class LightX2VInferenceConfig:
|
||||
},
|
||||
),
|
||||
"resize_mode": (
|
||||
["adaptive", "keep_ratio_fixed_area", "fixed_min_area", "fixed_max_area", "fixed_shape", "fixed_min_side"],
|
||||
[
|
||||
"adaptive",
|
||||
"keep_ratio_fixed_area",
|
||||
"fixed_min_area",
|
||||
"fixed_max_area",
|
||||
"fixed_shape",
|
||||
"fixed_min_side",
|
||||
],
|
||||
{
|
||||
"default": "adaptive",
|
||||
"tooltip": "Adaptive resize input image to target aspect ratio",
|
||||
@@ -282,7 +289,9 @@ class LightX2VTeaCache:
|
||||
|
||||
def create_config(self, enable, threshold, use_ret_steps):
|
||||
"""Create TeaCache configuration."""
|
||||
config = TeaCacheConfig(enable=enable, threshold=threshold, use_ret_steps=use_ret_steps)
|
||||
config = TeaCacheConfig(
|
||||
enable=enable, threshold=threshold, use_ret_steps=use_ret_steps
|
||||
)
|
||||
return (config.to_dict(),)
|
||||
|
||||
|
||||
@@ -514,7 +523,9 @@ class LightX2VLoRALoader:
|
||||
|
||||
def load_lora(self, lora_name, strength, lora_chain=None):
|
||||
"""Load and chain LoRA configurations."""
|
||||
chain = LoRAChainBuilder.build_chain(lora_name=lora_name, strength=strength, existing_chain=lora_chain)
|
||||
chain = LoRAChainBuilder.build_chain(
|
||||
lora_name=lora_name, strength=strength, existing_chain=lora_chain
|
||||
)
|
||||
return (chain,)
|
||||
|
||||
|
||||
@@ -523,12 +534,18 @@ class TalkObjectInput:
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"name": ("STRING", {"default": "person_1", "tooltip": "说话人名称标识"}),
|
||||
"name": (
|
||||
"STRING",
|
||||
{"default": "person_1", "tooltip": "speaker name identifier"},
|
||||
),
|
||||
},
|
||||
"optional": {
|
||||
"audio": ("AUDIO", {"tooltip": "上传的音频文件"}),
|
||||
"mask": ("MASK", {"tooltip": "上传的遮罩图像(可选)"}),
|
||||
"save_to_input": ("BOOLEAN", {"default": True, "tooltip": "是否保存到input文件夹"}),
|
||||
"audio": ("AUDIO", {"tooltip": "uploaded audio file"}),
|
||||
"mask": ("MASK", {"tooltip": "uploaded mask image (optional)"}),
|
||||
"save_to_input": (
|
||||
"BOOLEAN",
|
||||
{"default": True, "tooltip": "save to input folder"},
|
||||
),
|
||||
},
|
||||
}
|
||||
|
||||
@@ -541,7 +558,9 @@ class TalkObjectInput:
|
||||
"""Create a talk object from input data."""
|
||||
builder = TalkObjectConfigBuilder()
|
||||
|
||||
talk_object = builder.build_from_input(name=name, audio=audio, mask=mask, save_to_input=save_to_input)
|
||||
talk_object = builder.build_from_input(
|
||||
name=name, audio=audio, mask=mask, save_to_input=save_to_input
|
||||
)
|
||||
|
||||
if talk_object:
|
||||
return (talk_object.to_dict(),)
|
||||
@@ -549,15 +568,16 @@ class TalkObjectInput:
|
||||
|
||||
|
||||
class TalkObjectsCombiner:
|
||||
"""组合多个谈话对象为配置"""
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
inputs = {"required": {}, "optional": {}}
|
||||
|
||||
# 预定义10个TALK_OBJECT输入槽
|
||||
# Pre-defined 10 TALK_OBJECT input slots
|
||||
for i in range(1, 11):
|
||||
inputs["optional"][f"talk_object_{i}"] = ("TALK_OBJECT", {"tooltip": f"谈话对象{i}"})
|
||||
inputs["optional"][f"talk_object_{i}"] = (
|
||||
"TALK_OBJECT",
|
||||
{"tooltip": f"talk object {i}"},
|
||||
)
|
||||
|
||||
return inputs
|
||||
|
||||
@@ -591,7 +611,7 @@ class TalkObjectsFromJSON:
|
||||
{
|
||||
"multiline": True,
|
||||
"default": '[{"name": "person1", "audio": "/path/to/audio1.wav", "mask": "/path/to/mask1.png"}]',
|
||||
"tooltip": "JSON格式的谈话对象配置",
|
||||
"tooltip": "JSON format talk objects configuration",
|
||||
},
|
||||
),
|
||||
},
|
||||
@@ -613,11 +633,32 @@ class TalkObjectsFromFiles:
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"audio_files": ("STRING", {"multiline": True, "default": "audio1.wav\naudio2.wav", "tooltip": "音频文件列表(每行一个)"}),
|
||||
"audio_files": (
|
||||
"STRING",
|
||||
{
|
||||
"multiline": True,
|
||||
"default": "audio1.wav\naudio2.wav",
|
||||
"tooltip": "audio file list (one per line)",
|
||||
},
|
||||
),
|
||||
},
|
||||
"optional": {
|
||||
"mask_files": ("STRING", {"multiline": True, "default": "mask1.png\nmask2.png", "tooltip": "遮罩文件列表(每行一个,可选)"}),
|
||||
"names": ("STRING", {"multiline": True, "default": "person1\nperson2", "tooltip": "人物名称列表(每行一个,可选)"}),
|
||||
"mask_files": (
|
||||
"STRING",
|
||||
{
|
||||
"multiline": True,
|
||||
"default": "mask1.png\nmask2.png",
|
||||
"tooltip": "mask file list (one per line, optional)",
|
||||
},
|
||||
),
|
||||
"names": (
|
||||
"STRING",
|
||||
{
|
||||
"multiline": True,
|
||||
"default": "person1\nperson2",
|
||||
"tooltip": "talk object name list (one per line, optional)",
|
||||
},
|
||||
),
|
||||
},
|
||||
}
|
||||
|
||||
@@ -659,7 +700,10 @@ class LightX2VConfigCombiner:
|
||||
{"tooltip": "Memory optimization configuration"},
|
||||
),
|
||||
"lora_chain": ("LORA_CHAIN", {"tooltip": "LoRA chain configuration"}),
|
||||
"talk_objects_config": ("TALK_OBJECTS_CONFIG", {"tooltip": "Multi-person talk objects configuration"}),
|
||||
"talk_objects_config": (
|
||||
"TALK_OBJECTS_CONFIG",
|
||||
{"tooltip": "Multi-person talk objects configuration"},
|
||||
),
|
||||
},
|
||||
}
|
||||
|
||||
@@ -681,10 +725,26 @@ class LightX2VConfigCombiner:
|
||||
# Convert dict configs back to objects if needed
|
||||
|
||||
# Create objects from dicts
|
||||
inf_config = InferenceConfig(**inference_config) if isinstance(inference_config, dict) else None
|
||||
tea_config = TeaCacheConfig(**teacache_config) if teacache_config and isinstance(teacache_config, dict) else None
|
||||
quant_config = QuantizationConfig(**quantization_config) if quantization_config and isinstance(quantization_config, dict) else None
|
||||
mem_config = MemoryOptimizationConfig(**memory_config) if memory_config and isinstance(memory_config, dict) else None
|
||||
inf_config = (
|
||||
InferenceConfig(**inference_config)
|
||||
if isinstance(inference_config, dict)
|
||||
else None
|
||||
)
|
||||
tea_config = (
|
||||
TeaCacheConfig(**teacache_config)
|
||||
if teacache_config and isinstance(teacache_config, dict)
|
||||
else None
|
||||
)
|
||||
quant_config = (
|
||||
QuantizationConfig(**quantization_config)
|
||||
if quantization_config and isinstance(quantization_config, dict)
|
||||
else None
|
||||
)
|
||||
mem_config = (
|
||||
MemoryOptimizationConfig(**memory_config)
|
||||
if memory_config and isinstance(memory_config, dict)
|
||||
else None
|
||||
)
|
||||
|
||||
config = self.config_builder.combine_configs(
|
||||
inference_config=inf_config,
|
||||
@@ -703,11 +763,11 @@ class LightX2VModularInference:
|
||||
_current_config_hash = None
|
||||
|
||||
def __init__(self):
|
||||
if not hasattr(self.__class__, '_current_runner'):
|
||||
if not hasattr(self.__class__, "_current_runner"):
|
||||
self.__class__._current_runner = None
|
||||
if not hasattr(self.__class__, '_current_config_hash'):
|
||||
if not hasattr(self.__class__, "_current_config_hash"):
|
||||
self.__class__._current_config_hash = None
|
||||
|
||||
|
||||
self.config_builder = ConfigBuilder()
|
||||
self.temp_manager = TempFileManager()
|
||||
self.image_handler = ImageFileHandler()
|
||||
@@ -778,7 +838,11 @@ class LightX2VModularInference:
|
||||
logging.info(f"Image saved to {temp_path}")
|
||||
|
||||
# Handle audio input for seko models
|
||||
if audio is not None and hasattr(config, "model_cls") and "seko" in config.model_cls:
|
||||
if (
|
||||
audio is not None
|
||||
and hasattr(config, "model_cls")
|
||||
and "seko" in config.model_cls
|
||||
):
|
||||
temp_path = self.temp_manager.create_temp_file(suffix=".wav")
|
||||
self.audio_handler.save(audio, temp_path)
|
||||
config.audio_path = temp_path
|
||||
@@ -805,37 +869,52 @@ class LightX2VModularInference:
|
||||
for obj in processed_talk_objects:
|
||||
if "audio" in obj and obj["audio"]:
|
||||
audio_path = obj["audio"]
|
||||
if not os.path.isabs(audio_path) and not audio_path.startswith("/tmp"):
|
||||
if not os.path.isabs(audio_path) and not audio_path.startswith(
|
||||
"/tmp"
|
||||
):
|
||||
obj["audio"] = self.resolver.resolve_input_path(audio_path)
|
||||
logging.info(f"Resolved audio path: {audio_path} -> {obj['audio']}")
|
||||
logging.info(
|
||||
f"Resolved audio path: {audio_path} -> {obj['audio']}"
|
||||
)
|
||||
if not os.path.exists(obj["audio"]):
|
||||
logging.warning(f"Audio file not found: {obj['audio']}")
|
||||
|
||||
if "mask" in obj and obj["mask"]:
|
||||
mask_path = obj["mask"]
|
||||
if not os.path.isabs(mask_path) and not mask_path.startswith("/tmp"):
|
||||
if not os.path.isabs(mask_path) and not mask_path.startswith(
|
||||
"/tmp"
|
||||
):
|
||||
obj["mask"] = self.resolver.resolve_input_path(mask_path)
|
||||
logging.info(f"Resolved mask path: {mask_path} -> {obj['mask']}")
|
||||
logging.info(
|
||||
f"Resolved mask path: {mask_path} -> {obj['mask']}"
|
||||
)
|
||||
if not os.path.exists(obj["mask"]):
|
||||
logging.warning(f"Mask file not found: {obj['mask']}")
|
||||
|
||||
if processed_talk_objects:
|
||||
config.talk_objects = processed_talk_objects
|
||||
logging.info(f"Processed {len(processed_talk_objects)} talk objects")
|
||||
logging.info(
|
||||
f"Processed {len(processed_talk_objects)} talk objects"
|
||||
)
|
||||
|
||||
logging.info("lightx2v config: " + json.dumps(config, indent=2, ensure_ascii=False))
|
||||
|
||||
config_hash = self._get_config_hash(config)
|
||||
|
||||
# 安全地访问类属性
|
||||
current_runner = getattr(self.__class__, '_current_runner', None)
|
||||
current_config_hash = getattr(self.__class__, '_current_config_hash', None)
|
||||
|
||||
needs_reinit = (
|
||||
current_runner is None or current_config_hash != config_hash or getattr(config, "lazy_load", False)
|
||||
logging.info(
|
||||
"lightx2v config: " + json.dumps(config, indent=2, ensure_ascii=False)
|
||||
)
|
||||
|
||||
logging.info(f"Needs reinit: {needs_reinit}, old config hash: {current_config_hash}, new config hash: {config_hash}")
|
||||
config_hash = self._get_config_hash(config)
|
||||
|
||||
current_runner = getattr(self.__class__, "_current_runner", None)
|
||||
current_config_hash = getattr(self.__class__, "_current_config_hash", None)
|
||||
|
||||
needs_reinit = (
|
||||
current_runner is None
|
||||
or current_config_hash != config_hash
|
||||
or getattr(config, "lazy_load", False)
|
||||
)
|
||||
|
||||
logging.info(
|
||||
f"Needs reinit: {needs_reinit}, old config hash: {current_config_hash}, new config hash: {config_hash}"
|
||||
)
|
||||
if needs_reinit:
|
||||
if current_runner is not None:
|
||||
# current_runner.end_run()
|
||||
@@ -856,9 +935,8 @@ class LightX2VModularInference:
|
||||
def update_progress(current_step, _total):
|
||||
progress.update_absolute(current_step)
|
||||
|
||||
# 重新获取当前runner,因为可能在reinit过程中发生了变化
|
||||
current_runner = getattr(self.__class__, '_current_runner', None)
|
||||
|
||||
current_runner = getattr(self.__class__, "_current_runner", None)
|
||||
|
||||
if hasattr(current_runner, "set_progress_callback"):
|
||||
current_runner.set_progress_callback(update_progress)
|
||||
|
||||
@@ -867,7 +945,7 @@ class LightX2VModularInference:
|
||||
audio = result_dict.get("audio", None)
|
||||
|
||||
if getattr(config, "unload_after_inference", False):
|
||||
if hasattr(self.__class__, '_current_runner'):
|
||||
if hasattr(self.__class__, "_current_runner"):
|
||||
del self.__class__._current_runner
|
||||
self.__class__._current_runner = None
|
||||
self.__class__._current_config_hash = None
|
||||
|
||||
Reference in New Issue
Block a user