feat: add talk_objects configuration option in LightX2VDefaultConfig; refactor TalkObjectConfigBuilder to streamline audio handling and remove unused source_type attribute; enhance TalkObjectsConfig methods for better data representation

This commit is contained in:
gaclove
2025-09-26 07:02:25 +00:00
parent 06e2791795
commit 098a38a5e5
4 changed files with 13 additions and 43 deletions
+1 -1
View File
@@ -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 = {
+2 -26
View File
@@ -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())
+2 -9
View File
@@ -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
+8 -7
View File
@@ -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)