From 06e27917954e8a50666b9aa17435ca55c9f76c69 Mon Sep 17 00:00:00 2001 From: gaclove Date: Fri, 26 Sep 2025 06:01:34 +0000 Subject: [PATCH] refactor: enhance code readability by formatting multi-line parameters and comments; improve configuration handling in various classes for better maintainability --- bridge.py | 98 ++++++++++++++++++++------- config_builder.py | 57 +++++++++++++--- data_models.py | 6 +- file_handlers.py | 31 +++++++-- nodes.py | 168 +++++++++++++++++++++++++++++++++------------- 5 files changed, 271 insertions(+), 89 deletions(-) diff --git a/bridge.py b/bridge.py index 2a4bc92..80fd75c 100644 --- a/bridge.py +++ b/bridge.py @@ -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: diff --git a/config_builder.py b/config_builder.py index f137ab0..ac74482 100644 --- a/config_builder.py +++ b/config_builder.py @@ -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 = [] diff --git a/data_models.py b/data_models.py index ab2416b..a0a2c6f 100644 --- a/data_models.py +++ b/data_models.py @@ -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 diff --git a/file_handlers.py b/file_handlers.py index ddf29ef..30d1592 100644 --- a/file_handlers.py +++ b/file_handlers.py @@ -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) diff --git a/nodes.py b/nodes.py index 0011f83..65cb557 100644 --- a/nodes.py +++ b/nodes.py @@ -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