You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

Celery Worker中Whisper模型按需加载卸载的实现方案问询

Celery调度Whisper模型的动态加载/卸载实现方案

问题根源

  1. 信号处理函数taskReceivedHandler中未声明global model,导致赋值的是局部变量,全局model始终为None。
  2. 用Redis列表维护任务列表存在并发竞态,llen和lpush非原子操作,可能导致模型加载判断错误。
  3. 任务执行与模型加载的时序无保障,可能出现任务开始执行时模型还未加载完成。

修正后的实现代码

from celery import Celery, Task
from celery.signals import task_received, celeryd_after_setup, task_success, task_failure
import redis
import torch
import whisper
from celery.utils.log import get_task_logger

logger = get_task_logger(__name__)

# 初始化Celery实例(需补充你的配置)
celery = Celery('whisper_tasks')
celery.config_from_object('your_celery_config')

# Redis连接(替换为你的hostname)
hostname = 'your_redis_host'
r = redis.Redis(host=hostname)

# 全局模型变量,每个Worker进程独立维护
model = None
# 任务计数的Redis键,格式为"worker:{worker_name}:active_tasks"
ACTIVE_TASKS_KEY = None

def get_worker_active_tasks_key(worker_name):
    return f"worker:{worker_name}:active_tasks"

def check_and_unload_model():
    global model
    # 原子获取当前任务数
    task_count = r.get(ACTIVE_TASKS_KEY)
    if task_count is None or int(task_count) == 0:
        if model is not None:
            logger.info("无剩余任务,卸载Whisper模型并释放GPU显存")
            del model
            model = None
            torch.cuda.empty_cache()

@celery.task(bind=True)
def myTask(self):
    global model
    # 确保模型已加载(处理时序问题,比如信号触发后模型还在加载的情况)
    while model is None:
        import time
        time.sleep(0.1)
    # 执行Whisper任务
    result = launchAnalyticsWhisper(model)
    # 返回任务ID方便后续处理(如果需要)
    return {"taskId": self.request.id, "result": result}

@task_received.connect(sender=myTask)
def task_received_handler(sender, request, **kwargs):
    global model
    worker_name = kwargs.get('sender') or request.hostname
    # 原子递增任务计数
    task_count = r.incr(get_worker_active_tasks_key(worker_name))
    # 任务计数为1时,说明是当前Worker的第一个任务,加载模型
    if task_count == 1:
        logger.info("收到首个任务,加载Whisper medium模型")
        model = whisper.load_model("medium")
        # 可选:将模型移到GPU(如果需要)
        if torch.cuda.is_available():
            model = model.to("cuda")

@task_success.connect(sender=myTask)
def task_success_handler(sender, result, **kwargs):
    worker_name = kwargs.get('sender') or result["taskId"].split('@')[1]
    # 原子递减任务计数
    r.decr(get_worker_active_tasks_key(worker_name))
    check_and_unload_model()

@task_failure.connect(sender=myTask)
def task_failure_handler(sender, task_id, exception, **kwargs):
    worker_name = kwargs.get('sender') or task_id.split('@')[1]
    # 原子递减任务计数
    r.decr(get_worker_active_tasks_key(worker_name))
    check_and_unload_model()

@celeryd_after_setup.connect
def init_worker_state(sender, instance, **kwargs):
    global ACTIVE_TASKS_KEY
    worker_name = instance.hostname
    ACTIVE_TASKS_KEY = get_worker_active_tasks_key(worker_name)
    # 初始化Worker的任务计数为0
    r.set(ACTIVE_TASKS_KEY, 0)
    logger.info(f"初始化Worker {worker_name} 的任务计数")

关键优化点

  • 全局变量作用域修正:在task_received_handler和check_and_unload_model中明确使用global model,确保赋值的是全局变量。
  • 原子任务计数:用Redis的incr/decr原子命令替代列表操作,避免并发下的竞态问题,准确统计当前Worker的活跃任务数。
  • 时序保障:在任务myTask中增加循环等待,确保模型加载完成后再执行任务逻辑。
  • Worker独立隔离:每个Worker使用独立的Redis键维护任务计数,避免多Worker之间的干扰。
  • 显存释放:卸载模型后调用torch.cuda.empty_cache()主动释放GPU显存,确保其他系统能正常占用。

注意事项

  • 确保Celery Worker以单进程运行(或每个进程独立维护模型),如果使用多进程Worker,每个进程会有自己的model变量,需根据实际部署调整。
  • Whisper模型加载耗时较长,首次任务的响应时间会包含模型加载时间,但后续任务无需重复加载。
  • 可根据实际需求调整模型加载的等待逻辑,比如设置超时时间,避免无限等待。

内容的提问来源于stack exchange,提问作者UnPapeur

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.06 00:26:01