diff --git a/bridge.py b/bridge.py index 80fd75c..776b023 100644 --- a/bridge.py +++ b/bridge.py @@ -160,6 +160,7 @@ class LightX2VDefaultConfig: "cfg_parallel": False, "audio_sr": 16000, "return_video": True, + "talk_objects": None, } @@ -300,7 +301,6 @@ class ModularConfigManager: updates[update_key] = config[config_key] def apply_inference_config(self, config: Dict[str, Any]) -> Dict[str, Any]: - """Apply basic inference configuration.""" updates = {} basic_mappings = { diff --git a/config_builder.py b/config_builder.py index ac74482..f428bfd 100644 --- a/config_builder.py +++ b/config_builder.py @@ -183,23 +183,19 @@ class TalkObjectConfigBuilder: mask: Optional[Any] = None, save_to_input: bool = True, ) -> TalkObject: - """Build talk object from input data.""" if audio is None: return None talk_object = TalkObject(name=name) - # Process audio if save_to_input and audio is not None: audio_path = self._save_audio_to_input(name, audio) if audio_path: talk_object.audio = audio_path - talk_object.source_type = "file" + else: talk_object.audio = audio - talk_object.source_type = "data" - # Process mask if mask is not None: if save_to_input: mask_path = self._save_mask_to_input(name, mask) @@ -228,7 +224,6 @@ class TalkObjectConfigBuilder: name=obj_data.get("name", "unknown"), audio=obj_data["audio"], mask=obj_data.get("mask"), - source_type="path", ) config.add_object(talk_obj) @@ -259,14 +254,12 @@ class TalkObjectConfigBuilder: name=name_list[i] if i < len(name_list) else f"person_{i + 1}", audio=audio_file, mask=mask_list[i] if i < len(mask_list) else None, - source_type="file", ) config.add_object(talk_obj) return config def _save_audio_to_input(self, name: str, audio_data: Any) -> Optional[str]: - """Save audio data to input directory.""" try: filename = f"{name}_audio_{uuid.uuid4().hex[:8]}.wav" return self.resolver.save_to_input(audio_data, filename, self.audio_handler) @@ -275,7 +268,6 @@ class TalkObjectConfigBuilder: return None def _save_mask_to_input(self, name: str, mask_data: Any) -> Optional[str]: - """Save mask data to input directory.""" try: filename = f"{name}_mask_{uuid.uuid4().hex[:8]}.png" return self.resolver.save_to_input(mask_data, filename, self.mask_handler) @@ -300,8 +292,6 @@ class ConfigBuilder: lora_chain: Optional[List[Dict[str, Any]]] = None, talk_objects_config: Optional[TalkObjectsConfig] = None, ) -> EasyDict: - """Combine all configurations into final config.""" - # Create combined config object combined = CombinedConfig( inference=inference_config, teacache=teacache_config, @@ -318,26 +308,12 @@ class ConfigBuilder: ) 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, - "memory": memory_config.to_dict() if memory_config else None, - } - - # Filter out None values + configs_dict = combined.to_dict() configs_dict = {k: v for k, v in configs_dict.items() if v is not None} - - # Use existing manager to build config final_config = self.manager.build_final_config(configs_dict) - # Add additional configs if lora_chain: final_config.lora_configs = lora_chain - if talk_objects_config: final_config.update(talk_objects_config.to_dict()) diff --git a/data_models.py b/data_models.py index a0a2c6f..ae16cb6 100644 --- a/data_models.py +++ b/data_models.py @@ -14,7 +14,6 @@ class TalkObject: name: str audio: Optional[Union[str, Dict[str, Any], torch.Tensor, np.ndarray]] = None mask: Optional[Union[str, torch.Tensor, np.ndarray]] = None - source_type: str = "data" # "data", "file", or "path" def to_dict(self) -> Dict[str, Any]: """Convert to dictionary for pipeline.""" @@ -30,7 +29,6 @@ class TalkObject: elif self.mask is not None: result["mask_data"] = self.mask - result["source_type"] = self.source_type return result @@ -149,21 +147,16 @@ class LoRAConfig: @dataclass class TalkObjectsConfig: - """Configuration for multiple talk objects.""" - talk_objects: List[TalkObject] = field(default_factory=list) def add_object(self, talk_object: TalkObject): - """Add a talk object to the configuration.""" self.talk_objects.append(talk_object) def to_dict(self) -> Dict[str, Any]: - """Convert to dictionary for pipeline.""" - return {"talk_objects": [obj for obj in self.talk_objects]} + return {"talk_objects": [obj.to_dict() for obj in self.talk_objects]} def to_list(self) -> List[Dict[str, Any]]: - """Get list of talk object dictionaries.""" - return [obj for obj in self.talk_objects] + return [obj.to_dict() for obj in self.talk_objects] @dataclass diff --git a/nodes.py b/nodes.py index 65cb557..b19e133 100644 --- a/nodes.py +++ b/nodes.py @@ -563,20 +563,21 @@ class TalkObjectInput: ) if talk_object: - return (talk_object.to_dict(),) + return (talk_object,) return (None,) class TalkObjectsCombiner: + PREDEFINED_SLOTS = 16 + @classmethod def INPUT_TYPES(cls): inputs = {"required": {}, "optional": {}} - # Pre-defined 10 TALK_OBJECT input slots - for i in range(1, 11): - inputs["optional"][f"talk_object_{i}"] = ( + for i in range(cls.PREDEFINED_SLOTS): + inputs["optional"][f"talk_object_{i + 1}"] = ( "TALK_OBJECT", - {"tooltip": f"talk object {i}"}, + {"tooltip": f"talk object {i + 1}"}, ) return inputs @@ -589,8 +590,8 @@ class TalkObjectsCombiner: def combine_talk_objects(self, **kwargs): config = TalkObjectsConfig() - for i in range(1, 11): - talk_obj = kwargs.get(f"talk_object_{i}") + for i in range(self.PREDEFINED_SLOTS): + talk_obj = kwargs.get(f"talk_object_{i + 1}") if talk_obj is not None: config.add_object(talk_obj)