Celery Worker中Whisper模型按需加载卸载的实现方案问询
Celery调度Whisper模型的动态加载/卸载实现方案
问题根源
- 信号处理函数
taskReceivedHandler中未声明global model,导致赋值的是局部变量,全局model始终为None。 - 用Redis列表维护任务列表存在并发竞态,
llen和lpush非原子操作,可能导致模型加载判断错误。 - 任务执行与模型加载的时序无保障,可能出现任务开始执行时模型还未加载完成。
修正后的实现代码
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
相关产品推荐
相关产品推荐

