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:
@@ -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
@@ -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
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user