refactor: clean up comments in universal_bridge.py for improved readability and maintainability
This commit is contained in:
@@ -350,7 +350,6 @@ class ConfigManager:
|
||||
# Direct mapping for unknown parameters
|
||||
config[comfy_key] = value
|
||||
|
||||
# 检查并加载模型路径中的config.json
|
||||
config = cls._load_model_config(config)
|
||||
|
||||
return EasyDict(config)
|
||||
@@ -368,7 +367,6 @@ class ConfigManager:
|
||||
try:
|
||||
with open(model_config_path, "r") as f:
|
||||
model_config = json.load(f)
|
||||
# 合并模型配置,模型配置优先级更高
|
||||
config.update(model_config)
|
||||
print(f"Loaded model config from: {model_config_path}")
|
||||
except Exception as e:
|
||||
@@ -409,11 +407,9 @@ class LightX2VConfigBuilder:
|
||||
# Get available models and configs
|
||||
model_choices = bridge.model_registry
|
||||
|
||||
# 构建配置选项
|
||||
config_choices = ["custom"]
|
||||
config_descriptions = {}
|
||||
|
||||
# 添加普通配置
|
||||
for key, config_info in bridge.config_registry.items():
|
||||
config_choices.append(key)
|
||||
config_descriptions[key] = config_info["description"]
|
||||
@@ -440,7 +436,6 @@ class LightX2VConfigBuilder:
|
||||
),
|
||||
}
|
||||
|
||||
# 添加用户可配置参数
|
||||
for param_name, param_config in ConfigManager.USER_CONFIGURABLE_PARAMS.items():
|
||||
required_inputs[param_name] = (
|
||||
param_config["type"],
|
||||
@@ -529,7 +524,6 @@ class LightX2VConfigBuilder:
|
||||
):
|
||||
"""Build configuration for LightX2V inference."""
|
||||
|
||||
# 加载基础配置
|
||||
if config_preset == "custom":
|
||||
base_config = {}
|
||||
else:
|
||||
@@ -539,11 +533,9 @@ class LightX2VConfigBuilder:
|
||||
else:
|
||||
raise ValueError(f"配置预设 '{config_preset}' 未找到")
|
||||
|
||||
# 解析自定义配置
|
||||
custom_cfg = ConfigManager.parse_custom_config(custom_config)
|
||||
base_config.update(custom_cfg)
|
||||
|
||||
# 合并其他配置
|
||||
if quantization_config:
|
||||
base_config.update(quantization_config)
|
||||
if attention_config:
|
||||
@@ -553,7 +545,6 @@ class LightX2VConfigBuilder:
|
||||
if distill_config:
|
||||
base_config.update(distill_config)
|
||||
|
||||
# 准备覆盖参数
|
||||
overrides = {
|
||||
"seed": seed if seed != -1 else None,
|
||||
"steps": steps,
|
||||
@@ -566,7 +557,6 @@ class LightX2VConfigBuilder:
|
||||
"strength_model": strength_model,
|
||||
}
|
||||
|
||||
# 创建最终配置
|
||||
config = ConfigManager.create_config(model_cls, model_path, task, base_config, overrides)
|
||||
|
||||
return (config,)
|
||||
@@ -580,7 +570,6 @@ class LightX2VQuantizationConfig:
|
||||
"""Define inputs for quantization config."""
|
||||
bridge = LightX2VBridge()
|
||||
|
||||
# Get available quantization configs
|
||||
quant_choices = ["custom"]
|
||||
for key, config_info in bridge.quantization_registry.items():
|
||||
quant_choices.append(key)
|
||||
@@ -824,56 +813,42 @@ class LightX2VInference:
|
||||
audio=None,
|
||||
**kwargs,
|
||||
):
|
||||
"""Generate images using LightX2V."""
|
||||
|
||||
# 直接使用传入的config (已经是EasyDict)
|
||||
# 添加运行时需要的参数
|
||||
config.prompt = prompt
|
||||
config.negative_prompt = negative_prompt
|
||||
config.mode = "infer"
|
||||
|
||||
# 设置环境变量(从run_wan_i2v.sh中提取)
|
||||
os.environ["TOKENIZERS_PARALLELISM"] = "false"
|
||||
if "DTYPE" not in os.environ:
|
||||
os.environ["DTYPE"] = "BF16" # 可通过配置更改
|
||||
os.environ["DTYPE"] = "BF16"
|
||||
if "ENABLE_GRAPH_MODE" not in os.environ:
|
||||
os.environ["ENABLE_GRAPH_MODE"] = "false"
|
||||
if "ENABLE_PROFILING_DEBUG" not in os.environ:
|
||||
os.environ["ENABLE_PROFILING_DEBUG"] = "true"
|
||||
|
||||
# 临时文件列表,用于清理
|
||||
temp_files = []
|
||||
|
||||
try:
|
||||
# 处理i2v任务的图像输入
|
||||
if config.task == "i2v" and image is not None:
|
||||
# 保存图像到临时文件
|
||||
pil_image = self.converter.comfy_image_to_pil(image)
|
||||
with tempfile.NamedTemporaryFile(suffix=".png", delete=False) as tmp:
|
||||
pil_image.save(tmp.name)
|
||||
config.image_path = tmp.name
|
||||
temp_files.append(tmp.name)
|
||||
elif config.task == "i2v" and image is None:
|
||||
raise ValueError("i2v任务需要输入图像")
|
||||
raise ValueError("i2v task requires input image")
|
||||
|
||||
# 处理音频输入
|
||||
if audio is not None and "audio" in config.model_cls:
|
||||
audio_path = self.converter.comfy_audio_to_path(audio)
|
||||
config.audio_path = audio_path
|
||||
temp_files.append(audio_path)
|
||||
|
||||
# 获取或创建runner
|
||||
runner = self.bridge.get_runner(config)
|
||||
|
||||
if runner is None:
|
||||
raise RuntimeError("Failed to initialize runner")
|
||||
|
||||
# 运行生成 - 我们需要创建一个自定义的pipeline来获取latents
|
||||
# 由于原始的run_pipeline方法会删除latents和generator,我们需要分步执行
|
||||
if asyncio.iscoroutinefunction(runner.run_pipeline):
|
||||
# 在同步环境中运行异步方法
|
||||
try:
|
||||
try:
|
||||
# 使用现有的事件循环(如果存在)或创建新的
|
||||
loop = asyncio.get_event_loop()
|
||||
if loop.is_closed():
|
||||
loop = asyncio.new_event_loop()
|
||||
@@ -882,44 +857,40 @@ class LightX2VInference:
|
||||
loop = asyncio.new_event_loop()
|
||||
asyncio.set_event_loop(loop)
|
||||
|
||||
try:
|
||||
# 分步执行pipeline,在删除latents前获取它们
|
||||
# 这里我们使用try/except来捕获可能的属性访问错误
|
||||
try:
|
||||
# 尝试分步执行
|
||||
if hasattr(runner, "run_input_encoder"):
|
||||
runner.inputs = loop.run_until_complete(runner.run_input_encoder())
|
||||
if hasattr(runner, "set_target_shape"):
|
||||
kwargs = runner.set_target_shape()
|
||||
if hasattr(runner, "run_dit"):
|
||||
latents, generator = loop.run_until_complete(runner.run_dit(kwargs))
|
||||
images = loop.run_until_complete(runner.run_vae_decoder(latents, generator))
|
||||
# 直接解码latents为图像
|
||||
return self._decode_latents_to_images(images)
|
||||
except (AttributeError, TypeError):
|
||||
raise RuntimeError("Failed to get latents from generation")
|
||||
except Exception as e:
|
||||
print(f"Error during pipeline execution: {e}")
|
||||
raise
|
||||
else:
|
||||
try:
|
||||
if hasattr(runner, "run_input_encoder"):
|
||||
runner.inputs = runner.run_input_encoder()
|
||||
if hasattr(runner, "set_target_shape"):
|
||||
from ..lightx2v.lightx2v.models.runners.async_wrapper import AsyncWrapper
|
||||
|
||||
async def custom_pipeline():
|
||||
async with AsyncWrapper(runner) as wrapper:
|
||||
if runner.config.get("use_prompt_enhancer", False):
|
||||
runner.config["prompt_enhanced"] = await wrapper.run_prompt_enhancer()
|
||||
|
||||
runner.inputs = await wrapper.run_input_encoder()
|
||||
|
||||
kwargs = runner.set_target_shape()
|
||||
if hasattr(runner, "run_dit"):
|
||||
latents, generator = runner.run_dit(kwargs)
|
||||
# 直接解码latents为图像
|
||||
images = runner.run_vae_decoder(latents, generator)
|
||||
return self._decode_latents_to_images(images)
|
||||
except (AttributeError, TypeError):
|
||||
raise RuntimeError("Failed to get latents from generation")
|
||||
|
||||
latents, generator = await wrapper.run_dit(kwargs)
|
||||
|
||||
images = await wrapper.run_vae_decoder(latents, generator)
|
||||
|
||||
del latents, generator, runner.inputs
|
||||
|
||||
return images
|
||||
|
||||
images = loop.run_until_complete(custom_pipeline())
|
||||
|
||||
torch.cuda.empty_cache()
|
||||
gc.collect()
|
||||
|
||||
return self._decode_latents_to_images(images)
|
||||
|
||||
except Exception as e:
|
||||
print(f"Error during pipeline execution: {e}")
|
||||
raise
|
||||
|
||||
except Exception as e:
|
||||
print(f"Error in LightX2V generation: {e}")
|
||||
raise
|
||||
finally:
|
||||
# 清理临时文件
|
||||
for temp_file in temp_files:
|
||||
if os.path.exists(temp_file):
|
||||
try:
|
||||
@@ -930,10 +901,7 @@ class LightX2VInference:
|
||||
def _decode_latents_to_images(self, decoded_images):
|
||||
"""Decode latents to images using VAE."""
|
||||
|
||||
# 归一化从 [-1, 1] 到 [0, 1]
|
||||
images = (decoded_images + 1) / 2
|
||||
|
||||
# 重新排列维度为ComfyUI格式 [T, H, W, C]
|
||||
images = images.squeeze(0).permute(1, 2, 3, 0).cpu()
|
||||
images = torch.clamp(images, 0, 1)
|
||||
|
||||
@@ -1027,7 +995,6 @@ class LightX2VDistillConfig:
|
||||
|
||||
try:
|
||||
config["denoising_step_list"] = json.loads(denoising_steps)
|
||||
# 确保去噪步骤列表长度与推理步数匹配
|
||||
if len(config["denoising_step_list"]) != infer_steps:
|
||||
print(f"Warning: denoising_step_list length ({len(config['denoising_step_list'])}) doesn't match infer_steps ({infer_steps})")
|
||||
except json.JSONDecodeError:
|
||||
|
||||
Reference in New Issue
Block a user