WAN2.2文生视频ComfyUI工作流稳定性优化:OOM防护、超时熔断、重试机制

1. 引言

如果你用过WAN2.2文生视频工作流,大概率遇到过这种情况:满怀期待地输入一段精彩的提示词,点击“执行”按钮,然后……要么是漫长的等待后弹出一个内存不足的错误,要么是进度条卡在某个节点一动不动,最后只能无奈地重启整个流程。

文生视频本身就是一个计算密集型的任务,尤其是结合了SDXL Prompt风格化之后,对显存和算力的要求更是水涨船高。在ComfyUI中直接运行复杂工作流,稳定性问题就像一颗定时炸弹,随时可能让我们的创作过程戛然而止。

今天,我们不谈如何写出更惊艳的提示词,也不谈如何调出更美的画面。我们来聊聊一个更基础、但同样重要的话题:如何让你的WAN2.2文生视频工作流跑得更稳、更可靠。我将分享一套经过实战检验的稳定性优化方案,核心围绕三个机制展开:OOM(内存溢出)防护、超时熔断和智能重试。这套方案能显著降低任务失败率,让你从频繁的“重启-重试”循环中解脱出来,把更多精力投入到创意本身。

2. 理解WAN2.2工作流的稳定性挑战

在深入解决方案之前,我们有必要先搞清楚,为什么这个工作流如此“脆弱”。

2.1 资源消耗的“黑洞”

文生视频模型,尤其是像WAN2.2这样追求高质量输出的模型,本质上是一个巨大的参数海洋。生成每一帧画面都需要进行复杂的数学运算和大量的数据吞吐。

  • 显存压力:模型加载、中间特征图、视频帧缓存都需要占用显存。当视频分辨率提高、时长增加时,显存消耗几乎是呈指数级增长。一个常见的1080p、4秒视频生成任务,峰值显存占用轻松突破10GB。
  • 计算时长:不同于文生图可以秒级出结果,文生视频的推理过程漫长且不可预测。单个节点的计算卡顿(例如VAE编码、运动模块推理)会导致整个流水线阻塞。

2.2 ComfyUI工作流的“木桶效应”

ComfyUI的节点式工作流非常直观,但也带来了独特的稳定性问题。整个工作流就像一串多米诺骨牌,任何一个节点失败,都会导致后续所有节点失效。更棘手的是,某些节点的失败是“沉默”的——它不会立即报错,而是陷入一种假死状态,持续占用资源却不产出结果,直到你手动干预或系统资源耗尽。

2.3 外部环境的不确定性

即使你的本地硬件足够强大,也难免受到干扰:

  • 其他应用程序:一个后台突然启动的软件更新、一次不小心的多任务切换,都可能抢占宝贵的GPU资源。
  • 驱动与框架:CUDA驱动、PyTorch版本等底层环境的微小差异,有时会引发难以预料的兼容性问题,导致推理中断。

理解了这些挑战,我们就能有的放矢地构建我们的防御体系。接下来,我们将从三个维度,层层加固我们的工作流。

3. 核心防御机制一:OOM防护与资源管理

内存溢出是导致工作流崩溃的头号杀手。我们的目标不是无限扩充显存,而是在有限的资源内,聪明地完成任务。

3.1 动态显存监控与预警

与其在OOM发生后才手忙脚乱,不如提前感知风险。我们可以创建一个简单的守护脚本,在ComfyUI运行时持续监控GPU状态。

# gpu_monitor.py
import pynvml
import time
import threading
from queue import Queue

class GPUMonitor:
    def __init__(self, warning_threshold=0.85, check_interval=2):
        """
        初始化GPU监控器
        :param warning_threshold: 显存使用率警告阈值(0.85代表85%)
        :param check_interval: 检查间隔(秒)
        """
        pynvml.nvmlInit()
        self.device_count = pynvml.nvmlDeviceGetCount()
        self.warning_threshold = warning_threshold
        self.check_interval = check_interval
        self.alert_queue = Queue()
        
    def get_gpu_status(self):
        """获取所有GPU的实时状态"""
        status_list = []
        for i in range(self.device_count):
            handle = pynvml.nvmlDeviceGetHandleByIndex(i)
            mem_info = pynvml.nvmlDeviceGetMemoryInfo(handle)
            util = pynvml.nvmlDeviceGetUtilizationRates(handle)
            
            status = {
                'gpu_id': i,
                'total_mem': mem_info.total / 1024**3,  # 转换为GB
                'used_mem': mem_info.used / 1024**3,
                'free_mem': mem_info.free / 1024**3,
                'mem_util': mem_info.used / mem_info.total,
                'gpu_util': util.gpu
            }
            status_list.append(status)
        return status_list
    
    def monitor_loop(self):
        """监控循环,在独立线程中运行"""
        while True:
            try:
                status_list = self.get_gpu_status()
                for status in status_list:
                    if status['mem_util'] > self.warning_threshold:
                        warning_msg = f"[警告] GPU {status['gpu_id']} 显存使用率过高: {status['mem_util']:.1%}"
                        self.alert_queue.put(warning_msg)
                        # 这里可以触发自动降级策略,例如降低视频分辨率
                        self.trigger_degradation(status['gpu_id'])
                time.sleep(self.check_interval)
            except Exception as e:
                print(f"监控循环出错: {e}")
                time.sleep(5)
    
    def trigger_degradation(self, gpu_id):
        """触发降级策略(示例:通过API通知ComfyUI调整参数)"""
        # 在实际应用中,这里可以调用ComfyUI的API
        # 动态修改工作流中KSampler节点的分辨率参数
        print(f"正在为GPU{gpu_id}触发降级策略...")
        # 例如,将生成分辨率从1024x576降低到768x432

# 启动监控
monitor = GPUMonitor()
monitor_thread = threading.Thread(target=monitor.monitor_loop, daemon=True)
monitor_thread.start()

这个监控器会在后台运行,一旦发现显存使用率超过85%,就会发出警告,并可以自动触发预设的降级策略。

3.2 工作流参数化与自适应调整

最有效的OOM防护是“治未病”。我们可以通过参数化设计,让工作流根据可用资源自动调整。

1. 分辨率与时长联动调整: 不要固定使用1024x576的分辨率。我们可以创建一个配置表,将显存容量与推荐参数关联起来:

可用显存 (GB)推荐分辨率最大推荐时长 (秒)批处理大小
< 8768x43221
8 - 12896x50441
12 - 161024x57661
> 161152x64881

2. 在ComfyUI中实现参数注入: 我们可以修改工作流,使其能够接收外部传入的参数。在关键的Empty Latent Image节点和Video Linear CFG节点处,使用{{resolution}}和{{duration}}这样的占位符,然后通过一个前置的Python脚本动态替换这些值。

# workflow_param_injector.py
import json
import os

def adapt_workflow_to_memory(workflow_path, available_vram_gb):
    """根据可用显存自适应调整工作流参数"""
    with open(workflow_path, 'r', encoding='utf-8') as f:
        workflow = json.load(f)
    
    # 根据显存决定参数
    if available_vram_gb < 8:
        resolution = "768x432"
        duration_frames = 16  # 假设8fps,2秒视频
    elif available_vram_gb < 12:
        resolution = "896x504"
        duration_frames = 32  # 4秒视频
    elif available_vram_gb < 16:
        resolution = "1024x576"
        duration_frames = 48  # 6秒视频
    else:
        resolution = "1152x648"
        duration_frames = 64  # 8秒视频
    
    # 遍历工作流节点,替换参数(这里需要根据实际节点ID调整)
    for node_id, node in workflow.items():
        if node.get('_meta', {}).get('title') == 'Empty Latent Image':
            # 假设该节点有width和height属性
            width, height = map(int, resolution.split('x'))
            node['inputs']['width'] = width
            node['inputs']['height'] = height
        elif node.get('_meta', {}).get('title') == 'Video Linear CFG':
            # 调整视频时长相关参数
            node['inputs']['duration_frames'] = duration_frames
    
    # 保存调整后的工作流
    adapted_path = workflow_path.replace('.json', '_adapted.json')
    with open(adapted_path, 'w', encoding='utf-8') as f:
        json.dump(workflow, f, indent=2)
    
    print(f"工作流已适配,分辨率: {resolution}, 时长帧数: {duration_frames}")
    return adapted_path

# 使用示例
if __name__ == "__main__":
    # 假设我们检测到有10GB可用显存
    adapted_workflow = adapt_workflow_to_memory("wan2.2_workflow.json", 10)

3.3 显存碎片整理与缓存清理

长时间运行多个任务后,即使任务结束,PyTorch的显存缓存可能也不会立即释放,导致可用显存越来越小。我们可以在任务间隙插入清理例程。

# memory_cleaner.py
import torch
import gc

def aggressive_memory_cleanup():
    """
    执行激进的显存和内存清理
    在每次工作流执行前后调用此函数
    """
    if torch.cuda.is_available():
        torch.cuda.empty_cache()  # 清空PyTorch的CUDA缓存
        torch.cuda.synchronize()  # 等待所有CUDA操作完成
        
        # 尝试更彻底的清理(针对某些特定情况)
        for obj in gc.get_objects():
            try:
                if torch.is_tensor(obj) and obj.is_cuda:
                    del obj
            except:
                pass
    
    # 强制进行垃圾回收
    gc.collect()
    
    print("显存缓存已清理")

# 在工作流执行前和执行后调用
aggressive_memory_cleanup()

通过这三层防护,我们可以将OOM的风险降到最低。但稳定性不仅仅是内存问题,接下来我们看看如何应对“卡死”的情况。

4. 核心防御机制二:超时熔断与进程监控

当一个节点计算时间过长时,与其无限期等待,不如及时“熔断”,释放资源并尝试恢复。这就是超时熔断机制的核心思想。

4.1 基于子进程的执行封装

ComfyUI通常以服务器形式运行,我们不易直接控制其内部节点的执行。一个实用的方法是将整个工作流的执行封装在一个独立的子进程中,这样我们就可以控制它的生命周期。

# timeout_wrapper.py
import subprocess
import threading
import time
import signal
import os

class TimeoutProcess:
    def __init__(self, cmd, timeout_seconds=300, work_dir=None):
        """
        带超时控制的进程封装
        :param cmd: 要执行的命令列表
        :param timeout_seconds: 超时时间(秒)
        :param work_dir: 工作目录
        """
        self.cmd = cmd
        self.timeout = timeout_seconds
        self.work_dir = work_dir
        self.process = None
        self.timed_out = False
        self.output = ""
        self.error = ""
    
    def run(self):
        """运行命令,如果超时则终止进程"""
        def target():
            try:
                self.process = subprocess.Popen(
                    self.cmd,
                    cwd=self.work_dir,
                    stdout=subprocess.PIPE,
                    stderr=subprocess.PIPE,
                    text=True,
                    bufsize=1,
                    universal_newlines=True
                )
                
                # 实时读取输出
                while True:
                    line = self.process.stdout.readline()
                    if line:
                        self.output += line
                        print(line.strip())  # 实时打印到控制台
                    elif self.process.poll() is not None:
                        # 进程已结束,读取剩余输出
                        remaining_out, remaining_err = self.process.communicate()
                        self.output += remaining_out
                        self.error += remaining_err
                        break
                    time.sleep(0.1)
                        
            except Exception as e:
                self.error = str(e)
        
        # 在独立线程中运行目标函数
        thread = threading.Thread(target=target)
        thread.start()
        
        # 等待线程完成或超时
        thread.join(self.timeout)
        
        if thread.is_alive():
            # 超时发生
            self.timed_out = True
            print(f"进程执行超时({self.timeout}秒),正在终止...")
            
            # 尝试优雅终止
            if self.process:
                try:
                    self.process.terminate()
                    time.sleep(2)
                    if self.process.poll() is None:
                        self.process.kill()
                except:
                    pass
            
            # 强制结束线程
            return False, "执行超时"
        
        # 检查进程返回码
        if self.process and self.process.returncode != 0:
            return False, f"进程异常退出,返回码: {self.process.returncode}\n错误信息: {self.error}"
        
        return True, self.output

# 使用示例:封装ComfyUI工作流执行
def execute_comfyui_workflow(workflow_file, prompt, timeout=600):
    """
    执行ComfyUI工作流,带超时控制
    """
    # 构建执行命令(假设通过ComfyUI API触发)
    # 这里需要替换为实际的ComfyUI API调用命令
    cmd = [
        "python", "comfyui_api_client.py",
        "--workflow", workflow_file,
        "--prompt", prompt,
        "--timeout", str(timeout)
    ]
    
    executor = TimeoutProcess(cmd, timeout_seconds=timeout)
    success, result = executor.run()
    
    if success:
        print("工作流执行成功!")
        # 解析结果,获取生成的视频文件路径等
        return True, result
    else:
        print(f"工作流执行失败: {result}")
        return False, result

4.2 节点级超时检测

对于更精细的控制,我们可以尝试监控工作流中特定节点的执行时间。虽然ComfyUI没有直接提供节点级超时API,但我们可以通过分析执行日志来近似实现。

# node_timeout_detector.py
import re
import time
from datetime import datetime

class NodeExecutionMonitor:
    """通过日志分析监控节点执行时间"""
    
    def __init__(self, log_file_path, node_timeout_map=None):
        """
        :param log_file_path: ComfyUI日志文件路径
        :param node_timeout_map: 节点名称到超时时间(秒)的映射
        """
        self.log_file = log_file_path
        self.node_timeout_map = node_timeout_map or {
            'KSampler': 120,          # 采样器最多执行2分钟
            'VAE Encode': 60,         # VAE编码最多1分钟
            'Video Linear CFG': 180,  # 视频CFG最多3分钟
        }
        self.node_start_times = {}
        self.active_nodes = {}
        
    def tail_log(self):
        """模拟tail -f命令,持续读取日志新增内容"""
        with open(self.log_file, 'r', encoding='utf-8') as f:
            # 移动到文件末尾
            f.seek(0, 2)
            
            while True:
                line = f.readline()
                if not line:
                    time.sleep(0.1)
                    continue
                
                yield line
    
    def parse_log_line(self, line):
        """解析日志行,提取节点执行信息"""
        # 示例日志格式: [时间] 节点开始/结束执行
        # 实际正则表达式需要根据ComfyUI的日志格式调整
        start_pattern = r'Executing node: (.+?) \(ID: (.+?)\)'
        end_pattern = r'Node (.+?) \(ID: (.+?)\) finished'
        
        start_match = re.search(start_pattern, line)
        if start_match:
            node_name, node_id = start_match.groups()
            return 'start', node_name, node_id, datetime.now()
        
        end_match = re.search(end_pattern, line)
        if end_match:
            node_name, node_id = end_match.groups()
            return 'end', node_name, node_id, datetime.now()
        
        return None
    
    def monitor_loop(self):
        """监控循环,检测节点执行超时"""
        for line in self.tail_log():
            parsed = self.parse_log_line(line)
            if not parsed:
                continue
            
            event_type, node_name, node_id, timestamp = parsed
            
            if event_type == 'start':
                # 记录节点开始时间
                self.node_start_times[node_id] = timestamp
                self.active_nodes[node_id] = {
                    'name': node_name,
                    'start_time': timestamp
                }
                print(f"节点开始执行: {node_name} (ID: {node_id})")
                
            elif event_type == 'end':
                # 节点执行结束,清理记录
                if node_id in self.active_nodes:
                    elapsed = (timestamp - self.active_nodes[node_id]['start_time']).total_seconds()
                    print(f"节点执行完成: {node_name} (ID: {node_id}), 耗时: {elapsed:.1f}秒")
                    del self.active_nodes[node_id]
                    del self.node_start_times[node_id]
            
            # 检查是否有节点超时
            self.check_timeouts()
    
    def check_timeouts(self):
        """检查所有活跃节点是否超时"""
        current_time = datetime.now()
        for node_id, node_info in list(self.active_nodes.items()):
            node_name = node_info['name']
            start_time = node_info['start_time']
            
            elapsed = (current_time - start_time).total_seconds()
            timeout = self.node_timeout_map.get(node_name, 300)  # 默认5分钟
            
            if elapsed > timeout:
                print(f"[警报] 节点 {node_name} (ID: {node_id}) 已执行 {elapsed:.1f}秒,超过超时阈值 {timeout}秒!")
                # 这里可以触发熔断动作,例如通过API终止工作流
                self.trigger_circuit_breaker(node_id)
    
    def trigger_circuit_breaker(self, node_id):
        """触发熔断机制"""
        # 实际实现中,这里应该调用ComfyUI的API来终止特定节点或整个工作流
        print(f"正在对节点 {node_id} 执行熔断...")
        # 示例:发送HTTP请求到ComfyUI的中断端点
        # requests.post('http://localhost:8188/interrupt', json={'node_id': node_id})

# 启动监控
if __name__ == "__main__":
    monitor = NodeExecutionMonitor(
        log_file_path="comfyui.log",
        node_timeout_map={
            'KSampler': 120,
            'VAE Encode': 60,
            'Video Linear CFG': 180,
        }
    )
    
    # 在独立线程中运行监控
    import threading
    monitor_thread = threading.Thread(target=monitor.monitor_loop, daemon=True)
    monitor_thread.start()

4.3 熔断后的优雅降级

当熔断发生时,直接失败并不是唯一的选择。我们可以设计降级策略,尝试用更轻量的方式完成任务。

熔断降级策略示例:

  1. 分辨率降级:如果KSampler节点超时,自动将分辨率降低一档后重试。
  2. 跳过风格化:如果SDXL Prompt Styler节点处理时间过长,可以回退到不使用风格化的普通提示词。
  3. 缩短时长:如果视频生成整体超时,减少视频帧数,生成更短的视频。
# circuit_breaker_with_fallback.py

class CircuitBreakerWithFallback:
    """带降级策略的熔断器"""
    
    def __init__(self):
        self.fallback_strategies = {
            'high_resolution': self.fallback_to_medium_resolution,
            'sdxl_styling': self.fallback_to_basic_prompt,
            'long_video': self.fallback_to_short_video,
        }
    
    def execute_with_fallback(self, workflow_config, original_params):
        """
        执行工作流,如果失败则尝试降级策略
        """
        strategies_to_try = ['original'] + list(self.fallback_strategies.keys())
        
        for strategy in strategies_to_try:
            print(f"尝试策略: {strategy}")
            
            if strategy == 'original':
                # 原始参数
                params = original_params.copy()
            else:
                # 应用降级策略
                params = self.fallback_strategies[strategy](original_params.copy())
            
            # 执行工作流
            success, result = self.execute_workflow(workflow_config, params)
            
            if success:
                print(f"策略 '{strategy}' 执行成功")
                return True, result, strategy
            
            print(f"策略 '{strategy}' 执行失败,尝试下一个策略...")
        
        return False, "所有降级策略均失败", None
    
    def fallback_to_medium_resolution(self, params):
        """降级到中等分辨率"""
        if params.get('resolution') == '1024x576':
            params['resolution'] = '896x504'
        elif params.get('resolution') == '896x504':
            params['resolution'] = '768x432'
        return params
    
    def fallback_to_basic_prompt(self, params):
        """降级到基础提示词(跳过SDXL风格化)"""
        # 移除风格化相关的参数
        if 'style_preset' in params:
            del params['style_preset']
        params['prompt'] = params.get('raw_prompt', params.get('prompt', ''))
        return params
    
    def fallback_to_short_video(self, params):
        """降级到短视频"""
        original_frames = params.get('duration_frames', 48)
        params['duration_frames'] = max(16, original_frames // 2)  # 至少16帧,最多减半
        return params
    
    def execute_workflow(self, workflow_config, params):
        """执行工作流(实际实现需要调用ComfyUI API)"""
        # 这里应该是调用ComfyUI API的实际代码
        # 为示例简化,假设50%成功率
        import random
        success = random.random() > 0.5
        return success, "模拟执行结果" if success else "模拟执行失败"

有了超时熔断机制,我们就能及时止损,避免资源被无限占用。但有时候,失败只是暂时的,重试一下可能就成功了。

5. 核心防御机制三:智能重试与状态恢复

不是所有的失败都需要人工干预。智能重试机制可以自动处理临时性故障,如短暂的资源竞争、网络波动或随机计算错误。

5.1 指数退避重试策略

简单的立即重试可能会加重系统负担。指数退避策略在每次重试前等待越来越长的时间,既给了系统恢复的机会,又避免了雪崩效应。

# retry_with_backoff.py
import time
import random
from functools import wraps
from typing import Callable, Any, Tuple

class ExponentialBackoffRetry:
    """
    指数退避重试装饰器
    支持根据异常类型选择性地重试
    """
    
    def __init__(self, 
                 max_retries: int = 3,
                 base_delay: float = 1.0,
                 max_delay: float = 60.0,
                 retry_exceptions: Tuple = (Exception,)):
        """
        :param max_retries: 最大重试次数
        :param base_delay: 基础延迟(秒)
        :param max_delay: 最大延迟(秒)
        :param retry_exceptions: 需要重试的异常类型
        """
        self.max_retries = max_retries
        self.base_delay = base_delay
        self.max_delay = max_delay
        self.retry_exceptions = retry_exceptions
    
    def __call__(self, func: Callable) -> Callable:
        @wraps(func)
        def wrapper(*args, **kwargs) -> Any:
            last_exception = None
            
            for attempt in range(self.max_retries + 1):  # +1 包含第一次尝试
                try:
                    return func(*args, **kwargs)
                
                except self.retry_exceptions as e:
                    last_exception = e
                    
                    # 检查是否达到最大重试次数
                    if attempt >= self.max_retries:
                        print(f"达到最大重试次数 ({self.max_retries}),放弃重试")
                        raise
                    
                    # 计算退避延迟(指数退避 + 随机抖动)
                    delay = min(
                        self.base_delay * (2 ** attempt) + random.uniform(0, 1),
                        self.max_delay
                    )
                    
                    print(f"尝试 {func.__name__} 失败 (尝试 {attempt + 1}/{self.max_retries + 1})")
                    print(f"异常: {type(e).__name__}: {str(e)[:100]}...")
                    print(f"等待 {delay:.1f} 秒后重试...")
                    
                    time.sleep(delay)
            
            # 理论上不会执行到这里
            raise last_exception
        
        return wrapper

# 使用示例:装饰工作流执行函数
@ExponentialBackoffRetry(
    max_retries=3,
    base_delay=2.0,
    max_delay=30.0,
    retry_exceptions=(ConnectionError, TimeoutError, RuntimeError)
)
def execute_wan22_workflow(workflow_file, prompt, resolution="1024x576"):
    """
    执行WAN2.2工作流,自带重试机制
    """
    # 模拟可能失败的操作
    print(f"执行工作流: {workflow_file}, 提示词: {prompt[:50]}...")
    
    # 这里应该是实际的ComfyUI API调用
    # 为演示,我们随机模拟成功或失败
    import random
    if random.random() < 0.6:  # 60%成功率
        return {"status": "success", "video_path": "/path/to/generated/video.mp4"}
    else:
        # 模拟不同类型的失败
        failures = [
            ConnectionError("无法连接到ComfyUI服务器"),
            TimeoutError("API响应超时"),
            RuntimeError("CUDA内存不足"),
            RuntimeError("模型加载失败")
        ]
        raise random.choice(failures)

# 测试重试机制
try:
    result = execute_wan22_workflow(
        workflow_file="wan2.2_workflow.json",
        prompt="一只可爱的猫咪在草地上追逐蝴蝶,阳光明媚,风格为电影感",
        resolution="1024x576"
    )
    print(f"执行成功: {result}")
except Exception as e:
    print(f"所有重试均失败: {type(e).__name__}: {e}")

5.2 检查点与状态恢复

对于长时间运行的任务,从头开始重试成本太高。我们可以实现轻量级的检查点机制,保存中间状态,从失败点附近恢复。

# checkpoint_recovery.py
import json
import os
import pickle
from datetime import datetime

class WorkflowCheckpointManager:
    """工作流检查点管理器"""
    
    def __init__(self, checkpoint_dir="checkpoints"):
        self.checkpoint_dir = checkpoint_dir
        os.makedirs(checkpoint_dir, exist_ok=True)
    
    def save_checkpoint(self, workflow_id, node_id, node_output, metadata=None):
        """
        保存检查点
        :param workflow_id: 工作流唯一标识
        :param node_id: 当前完成的节点ID
        :param node_output: 节点输出数据
        :param metadata: 额外元数据
        """
        checkpoint_file = os.path.join(
            self.checkpoint_dir, 
            f"{workflow_id}_checkpoint_{node_id}.pkl"
        )
        
        checkpoint_data = {
            'workflow_id': workflow_id,
            'node_id': node_id,
            'node_output': node_output,
            'timestamp': datetime.now().isoformat(),
            'metadata': metadata or {}
        }
        
        # 使用pickle保存(对于复杂对象)
        with open(checkpoint_file, 'wb') as f:
            pickle.dump(checkpoint_data, f)
        
        # 同时保存JSON版本(用于调试)
        json_file = checkpoint_file.replace('.pkl', '.json')
        json_data = checkpoint_data.copy()
        # 尝试序列化node_output,如果不能序列化则保存类型信息
        try:
            json.dumps(json_data)
        except:
            json_data['node_output'] = f"<{type(node_output).__name__} object>"
        
        with open(json_file, 'w', encoding='utf-8') as f:
            json.dump(json_data, f, indent=2, default=str)
        
        print(f"检查点已保存: {checkpoint_file}")
        return checkpoint_file
    
    def load_checkpoint(self, workflow_id, node_id=None):
        """
        加载检查点
        :param workflow_id: 工作流唯一标识
        :param node_id: 要加载的节点ID,如果为None则加载最新的
        :return: 检查点数据,如果没有找到则返回None
        """
        if node_id:
            # 加载特定节点检查点
            checkpoint_file = os.path.join(
                self.checkpoint_dir, 
                f"{workflow_id}_checkpoint_{node_id}.pkl"
            )
            if os.path.exists(checkpoint_file):
                with open(checkpoint_file, 'rb') as f:
                    return pickle.load(f)
        else:
            # 查找最新的检查点
            pattern = f"{workflow_id}_checkpoint_*.pkl"
            checkpoints = []
            for fname in os.listdir(self.checkpoint_dir):
                if fname.startswith(f"{workflow_id}_checkpoint_"):
                    filepath = os.path.join(self.checkpoint_dir, fname)
                    mtime = os.path.getmtime(filepath)
                    checkpoints.append((mtime, filepath))
            
            if checkpoints:
                # 按修改时间排序,取最新的
                checkpoints.sort(reverse=True)
                latest_file = checkpoints[0][1]
                with open(latest_file, 'rb') as f:
                    return pickle.load(f)
        
        return None
    
    def resume_from_checkpoint(self, workflow_id, comfyui_api):
        """
        从检查点恢复工作流执行
        :param workflow_id: 工作流唯一标识
        :param comfyui_api: ComfyUI API客户端实例
        :return: 恢复后的执行结果
        """
        checkpoint = self.load_checkpoint(workflow_id)
        if not checkpoint:
            print(f"未找到工作流 {workflow_id} 的检查点,从头开始执行")
            return self.execute_full_workflow(workflow_id, comfyui_api)
        
        print(f"从检查点恢复: 节点 {checkpoint['node_id']}")
        
        # 获取工作流定义
        workflow_def = self.get_workflow_definition(workflow_id)
        
        # 找到检查点之后的节点
        nodes_to_resume = self.get_downstream_nodes(workflow_def, checkpoint['node_id'])
        
        # 设置已完成的节点输出
        comfyui_api.set_node_output(
            checkpoint['node_id'], 
            checkpoint['node_output']
        )
        
        # 继续执行后续节点
        results = {}
        for node_id in nodes_to_resume:
            try:
                print(f"恢复执行节点: {node_id}")
                result = comfyui_api.execute_node(node_id)
                results[node_id] = result
                
                # 保存新的检查点
                self.save_checkpoint(workflow_id, node_id, result)
                
            except Exception as e:
                print(f"节点 {node_id} 执行失败: {e}")
                # 可以在这里触发重试或熔断
                raise
        
        return results
    
    def get_downstream_nodes(self, workflow_def, start_node_id):
        """获取指定节点的下游节点(简化实现)"""
        # 这里需要根据实际工作流图结构实现
        # 简化示例:假设我们知道节点执行顺序
        all_nodes = ['load_checkpoint', 'clip_encode', 'ksampler', 'vae_decode', 'save_video']
        
        try:
            start_index = all_nodes.index(start_node_id)
            return all_nodes[start_index + 1:]
        except ValueError:
            return all_nodes  # 如果找不到,从头开始
    
    def execute_full_workflow(self, workflow_id, comfyui_api):
        """完整执行工作流(示例)"""
        print(f"完整执行工作流: {workflow_id}")
        # 实际实现应调用ComfyUI API
        return {"status": "completed"}

# 使用示例
checkpoint_mgr = WorkflowCheckpointManager()

# 模拟工作流执行
workflow_id = "wan22_workflow_001"
nodes = [
    {'id': 'load_checkpoint', 'name': '加载模型'},
    {'id': 'clip_encode', 'name': 'CLIP编码'},
    {'id': 'ksampler', 'name': 'K采样器'},
    {'id': 'vae_decode', 'name': 'VAE解码'},
    {'id': 'save_video', 'name': '保存视频'},
]

# 模拟执行到第三个节点时失败
for i, node in enumerate(nodes[:3]):  # 执行前三个节点
    print(f"执行节点: {node['name']} ({node['id']})")
    # 模拟节点输出
    node_output = {"status": "success", "data": f"output_of_{node['id']}"}
    checkpoint_mgr.save_checkpoint(workflow_id, node['id'], node_output)

print("\n模拟第三个节点后发生故障...\n")
print("系统重启后,从检查点恢复...")

# 从检查点恢复
last_checkpoint = checkpoint_mgr.load_checkpoint(workflow_id)
if last_checkpoint:
    print(f"成功加载检查点,最后完成的节点: {last_checkpoint['node_id']}")
    print(f"检查点时间: {last_checkpoint['timestamp']}")
    
    # 继续执行剩余节点
    remaining_nodes = [n for n in nodes if n['id'] not in ['load_checkpoint', 'clip_encode', 'ksampler']]
    for node in remaining_nodes:
        print(f"恢复执行: {node['name']} ({node['id']})")

5.3 基于失败类型的智能重试策略

不是所有错误都值得重试。我们可以根据错误类型决定重试策略:

# smart_retry_strategy.py

class SmartRetryStrategy:
    """基于错误类型的智能重试策略"""
    
    def __init__(self):
        self.strategy_map = {
            'CUDA out of memory': self.handle_oom,
            'Connection refused': self.handle_connection,
            'Timeout': self.handle_timeout,
            'Model not found': self.handle_model_error,
            'default': self.handle_generic_error
        }
    
    def should_retry(self, error_msg, attempt_count):
        """判断是否应该重试"""
        error_type = self.classify_error(error_msg)
        strategy = self.strategy_map.get(error_type, self.strategy_map['default'])
        return strategy(error_msg, attempt_count)
    
    def classify_error(self, error_msg):
        """根据错误信息分类错误类型"""
        error_msg_lower = error_msg.lower()
        
        if any(keyword in error_msg_lower for keyword in ['cuda', 'memory', 'oom']):
            return 'CUDA out of memory'
        elif any(keyword in error_msg_lower for keyword in ['connection', 'refused', 'reset']):
            return 'Connection refused'
        elif any(keyword in error_msg_lower for keyword in ['timeout', 'timed out']):
            return 'Timeout'
        elif any(keyword in error_msg_lower for keyword in ['model', 'weight', 'checkpoint']):
            return 'Model not found'
        else:
            return 'default'
    
    def handle_oom(self, error_msg, attempt_count):
        """处理OOM错误"""
        if attempt_count >= 2:  # OOM错误最多重试2次
            return False, "OOM错误重试次数过多,建议降低分辨率或批处理大小"
        
        # 第一次OOM,建议降低分辨率重试
        suggestion = "检测到显存不足,建议将分辨率从1024x576降低到896x504"
        return True, suggestion
    
    def handle_connection(self, error_msg, attempt_count):
        """处理连接错误"""
        if attempt_count >= 3:  # 连接错误最多重试3次
            return False, "连接失败,请检查ComfyUI服务是否正常运行"
        return True, "连接错误,等待后重试"
    
    def handle_timeout(self, error_msg, attempt_count):
        """处理超时错误"""
        if attempt_count >= 2:  # 超时错误最多重试2次
            return False, "执行超时,建议简化提示词或缩短视频时长"
        
        # 增加超时时间重试
        suggestion = "执行超时,将增加超时时间并重试"
        return True, suggestion
    
    def handle_model_error(self, error_msg, attempt_count):
        """处理模型错误"""
        # 模型错误通常无法通过重试解决
        return False, "模型文件错误,请检查模型路径和完整性"
    
    def handle_generic_error(self, error_msg, attempt_count):
        """处理通用错误"""
        if attempt_count >= 3:
            return False, "未知错误,重试次数过多"
        return True, "未知错误,尝试重试"

# 使用示例
retry_strategy = SmartRetryStrategy()

test_errors = [
    "RuntimeError: CUDA out of memory",
    "ConnectionRefusedError: [Errno 111] Connection refused",
    "TimeoutError: The read operation timed out",
    "FileNotFoundError: Model checkpoint not found",
    "Some random error"
]

for i, error in enumerate(test_errors):
    for attempt in range(1, 4):
        should_retry, suggestion = retry_strategy.should_retry(error, attempt)
        print(f"错误: {error[:50]}...")
        print(f"  尝试 {attempt}: 重试? {should_retry}, 建议: {suggestion}")
        if not should_retry:
            break
    print()

6. 实战:构建完整的稳定性优化流水线

现在,让我们把所有的机制组合起来,构建一个完整的稳定性优化流水线。这个流水线将OOM防护、超时熔断和智能重试有机地结合在一起。

# stability_pipeline.py
import time
import threading
from queue import Queue, Empty
from typing import Dict, Any, Optional

class StableWAN22Pipeline:
    """稳定的WAN2.2文生视频流水线"""
    
    def __init__(self, config: Dict[str, Any]):
        """
        初始化稳定性流水线
        :param config: 配置字典,包含各种参数
        """
        self.config = config
        self.status_queue = Queue()  # 状态消息队列
        self.stop_event = threading.Event()  # 停止事件
        
        # 初始化各个组件
        from gpu_monitor import GPUMonitor
        from node_timeout_detector import NodeExecutionMonitor
        from checkpoint_recovery import WorkflowCheckpointManager
        from smart_retry_strategy import SmartRetryStrategy
        
        self.gpu_monitor = GPUMonitor(
            warning_threshold=config.get('gpu_warning_threshold', 0.85)
        )
        self.node_monitor = NodeExecutionMonitor(
            log_file_path=config.get('log_file', 'comfyui.log'),
            node_timeout_map=config.get('node_timeouts', {})
        )
        self.checkpoint_mgr = WorkflowCheckpointManager(
            checkpoint_dir=config.get('checkpoint_dir', 'checkpoints')
        )
        self.retry_strategy = SmartRetryStrategy()
        
        # 执行统计
        self.stats = {
            'total_executions': 0,
            'successful_executions': 0,
            'failed_executions': 0,
            'retry_attempts': 0,
            'fallback_executions': 0,
            'avg_execution_time': 0
        }
    
    def execute_workflow(self, 
                        workflow_file: str, 
                        prompt: str, 
                        style: Optional[str] = None,
                        resolution: str = "1024x576",
                        duration_seconds: int = 4) -> Dict[str, Any]:
        """
        执行工作流,包含完整的稳定性保障
        """
        execution_id = f"wan22_{int(time.time())}"
        self.status_queue.put(f"开始执行工作流: {execution_id}")
        
        # 1. 预检:检查系统资源
        if not self.preflight_check():
            self.status_queue.put("预检失败,资源不足")
            return {"success": False, "error": "系统资源不足"}
        
        # 2. 自适应参数调整
        adapted_params = self.adapt_parameters(resolution, duration_seconds)
        self.status_queue.put(f"使用自适应参数: {adapted_params}")
        
        # 3. 清理内存
        self.cleanup_memory()
        
        # 4. 带重试的执行
        max_retries = self.config.get('max_retries', 3)
        last_error = None
        
        for attempt in range(max_retries + 1):
            try:
                self.status_queue.put(f"执行尝试 {attempt + 1}/{max_retries + 1}")
                
                # 启动监控
                self.start_monitoring(execution_id)
                
                # 执行工作流
                start_time = time.time()
                result = self._execute_single_attempt(
                    workflow_file, prompt, style, adapted_params, execution_id
                )
                end_time = time.time()
                
                # 更新统计
                self._update_stats(True, end_time - start_time)
                
                # 停止监控
                self.stop_monitoring()
                
                self.status_queue.put(f"执行成功: {execution_id}")
                return {"success": True, "result": result, "execution_id": execution_id}
                
            except Exception as e:
                last_error = e
                self._update_stats(False, 0)
                
                # 检查是否应该重试
                should_retry, suggestion = self.retry_strategy.should_retry(
                    str(e), attempt
                )
                
                self.status_queue.put(f"执行失败: {type(e).__name__}: {str(e)[:100]}")
                self.status_queue.put(f"重试建议: {suggestion}")
                
                if not should_retry or attempt >= max_retries:
                    break
                
                # 应用重试建议(如果可能)
                if "降低分辨率" in suggestion and "896x504" in suggestion:
                    adapted_params['resolution'] = "896x504"
                    self.status_queue.put("应用建议:降低分辨率到896x504")
                
                # 指数退避
                delay = min(2.0 * (2 ** attempt), 30.0)
                self.status_queue.put(f"等待 {delay:.1f}秒后重试...")
                time.sleep(delay)
        
        # 所有重试都失败
        self.stop_monitoring()
        self.status_queue.put(f"所有重试均失败: {last_error}")
        
        return {
            "success": False, 
            "error": str(last_error),
            "execution_id": execution_id,
            "stats": self.stats.copy()
        }
    
    def preflight_check(self) -> bool:
        """执行前检查系统资源"""
        try:
            gpu_status = self.gpu_monitor.get_gpu_status()
            if not gpu_status:
                self.status_queue.put("未检测到GPU")
                return False
            
            # 检查显存
            for status in gpu_status:
                if status['mem_util'] > 0.9:  # 使用率超过90%
                    self.status_queue.put(f"GPU {status['gpu_id']} 显存不足: {status['mem_util']:.1%}")
                    return False
            
            # 检查磁盘空间(简化示例)
            import shutil
            disk_usage = shutil.disk_usage("/")
            if disk_usage.free < 2 * 1024**3:  # 小于2GB
                self.status_queue.put("磁盘空间不足")
                return False
            
            return True
            
        except Exception as e:
            self.status_queue.put(f"预检异常: {e}")
            return False
    
    def adapt_parameters(self, resolution: str, duration_seconds: int) -> Dict[str, Any]:
        """根据系统资源自适应调整参数"""
        # 获取GPU状态
        gpu_status = self.gpu_monitor.get_gpu_status()
        if not gpu_status:
            return {"resolution": resolution, "duration_seconds": duration_seconds}
        
        # 使用第一个GPU的状态
        status = gpu_status[0]
        free_mem_gb = status['free_mem']
        
        # 根据可用显存调整参数
        adapted = {
            "resolution": resolution,
            "duration_seconds": duration_seconds,
            "batch_size": 1
        }
        
        if free_mem_gb < 6:
            adapted["resolution"] = "768x432"
            adapted["duration_seconds"] = min(duration_seconds, 2)
        elif free_mem_gb < 10:
            adapted["resolution"] = "896x504"
            adapted["duration_seconds"] = min(duration_seconds, 4)
        elif free_mem_gb < 14:
            adapted["resolution"] = "1024x576"
            adapted["duration_seconds"] = min(duration_seconds, 6)
        else:
            adapted["resolution"] = "1152x648"
            adapted["duration_seconds"] = min(duration_seconds, 8)
        
        return adapted
    
    def cleanup_memory(self):
        """清理内存"""
        import torch
        import gc
        
        if torch.cuda.is_available():
            torch.cuda.empty_cache()
            torch.cuda.synchronize()
        
        gc.collect()
        self.status_queue.put("内存清理完成")
    
    def start_monitoring(self, execution_id: str):
        """启动监控线程"""
        self.stop_event.clear()
        
        # 启动GPU监控
        self.gpu_monitor_thread = threading.Thread(
            target=self._gpu_monitoring_loop,
            args=(execution_id,),
            daemon=True
        )
        self.gpu_monitor_thread.start()
        
        # 启动节点监控
        self.node_monitor_thread = threading.Thread(
            target=self._node_monitoring_loop,
            args=(execution_id,),
            daemon=True
        )
        self.node_monitor_thread.start()
    
    def _gpu_monitoring_loop(self, execution_id: str):
        """GPU监控循环"""
        while not self.stop_event.is_set():
            try:
                status_list = self.gpu_monitor.get_gpu_status()
                for status in status_list:
                    if status['mem_util'] > 0.95:  # 超过95%使用率
                        self.status_queue.put(
                            f"[紧急] GPU {status['gpu_id']} 显存使用率: {status['mem_util']:.1%}"
                        )
                time.sleep(5)  # 每5秒检查一次
            except Exception as e:
                self.status_queue.put(f"GPU监控错误: {e}")
                time.sleep(10)
    
    def _node_monitoring_loop(self, execution_id: str):
        """节点监控循环(简化版)"""
        # 这里应该实现实际的节点监控逻辑
        # 为简化示例,我们只是模拟
        start_time = time.time()
        while not self.stop_event.is_set():
            elapsed = time.time() - start_time
            if elapsed > 300:  # 5分钟超时
                self.status_queue.put(f"[超时] 执行 {execution_id} 已运行 {elapsed:.0f}秒")
                # 这里可以触发熔断
                break
            time.sleep(10)
    
    def stop_monitoring(self):
        """停止所有监控"""
        self.stop_event.set()
        if hasattr(self, 'gpu_monitor_thread'):
            self.gpu_monitor_thread.join(timeout=2)
        if hasattr(self, 'node_monitor_thread'):
            self.node_monitor_thread.join(timeout=2)
    
    def _execute_single_attempt(self, workflow_file, prompt, style, params, execution_id):
        """单次执行尝试(实际应调用ComfyUI API)"""
        # 这里是实际调用ComfyUI API的地方
        # 为示例,我们模拟执行
        self.status_queue.put(f"执行工作流: {workflow_file}")
        self.status_queue.put(f"提示词: {prompt[:50]}...")
        self.status_queue.put(f"参数: {params}")
        
        # 模拟执行时间
        time.sleep(2)
        
        # 模拟随机成功/失败
        import random
        if random.random() < 0.8:  # 80%成功率
            return {
                "video_path": f"/output/{execution_id}.mp4",
                "resolution": params['resolution'],
                "duration": params['duration_seconds'],
                "file_size": f"{random.randint(50, 200)}MB"
            }
        else:
            raise RuntimeError("模拟执行失败")
    
    def _update_stats(self, success: bool, execution_time: float):
        """更新执行统计"""
        self.stats['total_executions'] += 1
        
        if success:
            self.stats['successful_executions'] += 1
            # 更新平均执行时间(移动平均)
            if self.stats['avg_execution_time'] == 0:
                self.stats['avg_execution_time'] = execution_time
            else:
                self.stats['avg_execution_time'] = (
                    0.9 * self.stats['avg_execution_time'] + 0.1 * execution_time
                )
        else:
            self.stats['failed_executions'] += 1
            self.stats['retry_attempts'] += 1
    
    def get_status_messages(self):
        """获取状态消息"""
        messages = []
        while True:
            try:
                messages.append(self.status_queue.get_nowait())
            except Empty:
                break
        return messages
    
    def get_statistics(self):
        """获取执行统计"""
        if self.stats['total_executions'] > 0:
            success_rate = (self.stats['successful_executions'] / 
                          self.stats['total_executions'] * 100)
        else:
            success_rate = 0
        
        stats = self.stats.copy()
        stats['success_rate'] = f"{success_rate:.1f}%"
        return stats

# 使用示例
if __name__ == "__main__":
    # 配置流水线
    config = {
        'gpu_warning_threshold': 0.85,
        'log_file': 'comfyui.log',
        'node_timeouts': {
            'KSampler': 120,
            'VAE Encode': 60,
            'Video Linear CFG': 180,
        },
        'checkpoint_dir': 'wan22_checkpoints',
        'max_retries': 3
    }
    
    # 创建流水线实例
    pipeline = StableWAN22Pipeline(config)
    
    # 执行工作流
    result = pipeline.execute_workflow(
        workflow_file="wan2.2_workflow.json",
        prompt="一只可爱的猫咪在草地上追逐蝴蝶,阳光明媚,风格为电影感",
        style="cinematic",
        resolution="1024x576",
        duration_seconds=4
    )
    
    # 打印结果
    print("\n" + "="*50)
    print("执行结果:")
    print(f"成功: {result['success']}")
    if result['success']:
        print(f"视频路径: {result['result']['video_path']}")
        print(f"分辨率: {result['result']['resolution']}")
        print(f"时长: {result['result']['duration']}秒")
    else:
        print(f"错误: {result['error']}")
    
    # 打印状态消息
    print("\n状态消息:")
    for msg in pipeline.get_status_messages()[-10:]:  # 最后10条消息
        print(f"  - {msg}")
    
    # 打印统计信息
    print("\n执行统计:")
    stats = pipeline.get_statistics()
    for key, value in stats.items():
        print(f"  {key}: {value}")

7. 总结

通过实施OOM防护、超时熔断和智能重试这三重稳定性优化机制,我们可以显著提升WAN2.2文生视频工作流的可靠性。让我们回顾一下关键要点:

7.1 核心机制回顾

  1. OOM防护与资源管理:通过动态显存监控、参数自适应调整和显存清理,预防内存溢出问题。关键是在问题发生前就采取行动,根据可用资源智能调整工作流参数。

  2. 超时熔断与进程监控:为长时间运行的任务设置合理的超时限制,通过子进程封装和节点级监控,及时终止卡住的任务,避免资源被无限占用。熔断后还可以尝试优雅降级,而不是直接失败。

  3. 智能重试与状态恢复:不是所有失败都是永久的。通过指数退避重试、检查点恢复和基于错误类型的智能重试策略,我们可以自动处理许多临时性故障,提高任务的整体成功率。

7.2 实践建议

在实际部署这套稳定性方案时,我有几个建议:

循序渐进地实施:不要一次性实现所有功能。先从最基本的OOM监控和简单重试开始,逐步添加更复杂的特性。这样更容易调试和排查问题。

监控与日志是关键:良好的监控和详细的日志是稳定性优化的基础。确保记录足够的上下文信息,这样当问题发生时,你才能快速定位原因。

参数需要调优:本文提供的参数(如超时时间、重试次数、显存阈值)都是示例值。你需要根据自己的硬件配置和工作流特点进行调整。建议先在测试环境中找到适合你的最佳参数。

保持简单有效:稳定性方案本身不应该成为新的不稳定因素。保持代码简洁,避免过度设计。每个机制都应该有明确的收益,而不是为了复杂而复杂。

7.3 展望未来

随着文生视频技术的快速发展,工作流会变得越来越复杂,对稳定性的要求也会越来越高。未来的优化方向可能包括:

  • 预测性资源调度:基于历史数据预测任务资源需求,提前进行调度和分配。
  • 更智能的降级策略:不仅仅是降低分辨率,还可以动态调整模型精度、跳过某些非关键步骤等。
  • 分布式执行支持:将工作流的不同节点分布到多台机器上执行,突破单机资源限制。
  • 自适应学习:系统能够从历史失败中学习,自动调整参数和策略。

稳定性不是一次性的工作,而是一个持续的过程。随着你对工作流的使用越来越深入,你会不断发现新的优化点。记住,最好的稳定性方案是那个能够真正解决你实际问题的方案。


获取更多AI镜像

想探索更多AI镜像和应用场景?访问 CSDN星图镜像广场,提供丰富的预置镜像,覆盖大模型推理、图像生成、视频生成、模型微调等多个领域,支持一键部署。

Logo

北京人形旗下天工造物具身智能开源社区,聚焦具身天工与慧思开物两大平台

更多推荐