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:
gaclove
2025-09-26 06:01:34 +00:00
parent b311e23c0e
commit 06e2791795
5 changed files with 271 additions and 89 deletions
+73 -25
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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)
+123 -45
View File
@@ -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