尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

FastAPI+GPU推理并发控制:避免显存溢出与OOM的实战方案

FastAPI+GPU推理并发控制:避免显存溢出与OOM的实战方案 1. 从一次线上事故说起为什么并发控制是GPU推理服务的生死线去年冬天我帮一个团队排查他们AI绘画服务的线上故障。现象很典型单用户测试一切正常一旦推广到几十个用户同时使用服务就像被掐住脖子一样先是响应时间从2秒飙升到30秒接着大量请求超时最后整个进程直接崩溃退出。查看日志满屏都是CUDA out of memory。重启服务撑不过五分钟又挂。运维同学一度怀疑是显卡坏了甚至准备申请采购新卡。这个场景在FastAPI GPU推理的架构里太常见了。FastAPI本身是异步框架天生适合处理高并发IO密集型任务但GPU推理是典型的计算密集型任务而且显存是独占资源。当多个请求同时涌入每个请求都试图往显存里塞模型、塞中间张量显存瞬间被撑爆服务自然就崩了。这不是FastAPI的锅也不是GPU的锅而是架构设计时没有做好并发控制。这篇文章就是围绕这个核心问题展开在FastAPI中调用GPU进行模型推理时如何设计一套可靠的并发控制机制让服务在请求量上涨时依然稳定运行而不是一多就显存溢出。我会从整体设计思路讲到具体代码实现从参数计算讲到踩坑经验尽量把每个决策背后的“为什么”说清楚。无论你是刚接触FastAPI的新手还是已经在跑推理服务但被并发问题困扰的开发者都能从中找到可以直接复用的方案。2. 整体设计思路把GPU当成稀缺资源来调度2.1 核心矛盾异步框架与同步计算的冲突FastAPI基于Starlette底层是asyncio事件循环。它的优势在于处理大量并发连接时不需要为每个请求创建线程而是通过协程切换来充分利用CPU时间。但GPU推理调用比如PyTorch的model.generate()或model.predict()是同步阻塞操作一旦执行就会占住当前线程事件循环被卡住其他请求只能排队等待。更麻烦的是很多开发者习惯在FastAPI的async def路由里直接调用同步的推理函数。这会导致两个问题第一事件循环被阻塞FastAPI的并发能力完全发挥不出来第二多个请求的推理任务同时向GPU提交显存分配器来不及回收直接OOM。我见过最典型的错误写法是这样的app.post(/predict) async def predict(request: Request): data await request.json() result model(data) # 同步阻塞调用且没有并发限制 return {result: result}这段代码在单请求下没问题但并发一上来model(data)会被多个协程几乎同时调用GPU显存瞬间被吃光。2.2 设计原则串行化GPU访问 异步化请求处理解决思路其实不复杂把GPU推理变成串行队列把请求处理保持异步。具体来说所有推理请求不直接调用模型而是把任务提交到一个队列里由一个或多个固定的工作线程按顺序从队列取任务执行。这样GPU同一时间只处理一个推理任务显存占用可控。同时FastAPI的异步路由可以继续接收新请求把任务放入队列后立即返回一个任务ID客户端通过轮询或WebSocket获取结果。这个模式的好处很明显显存可控同一时刻只有一个推理任务在跑显存峰值就是单次推理的峰值不会叠加。请求不丢失队列可以设置最大长度超过时返回429或503而不是直接崩溃。可扩展如果单卡性能不够可以增加工作线程数但要注意显存或者部署多卡多实例。响应及时异步接口不会因为推理慢而阻塞其他请求的接收。当然这个方案也有代价吞吐量受限于单次推理时间。如果单次推理需要500ms那理论最大QPS就是2。要提升吞吐要么优化模型要么用批处理batching要么多卡并行。但在显存有限、请求量波动大的场景下串行化是最稳妥的起点。2.3 方案选型为什么不用信号量或线程池直接限制有人可能会问用asyncio.Semaphore限制并发数不就行了吗比如设置信号量为1这样同一时间只有一个协程能进入推理函数。这个思路方向是对的但实际用起来有几个坑。第一asyncio.Semaphore只能限制协程的进入但如果你在async def里调用同步阻塞函数事件循环还是会被卡住。信号量释放的时机取决于阻塞函数何时返回这期间其他协程虽然拿不到信号量但事件循环本身已经被阻塞了连拒绝请求都做不到。第二信号量无法处理排队超时和队列长度限制。如果100个请求同时进来99个在等信号量它们会一直挂着占用连接资源直到超时。你没法告诉第50个请求“前面排队太多你先回去吧”。第三信号量不便于监控和动态调整。你很难知道当前队列里积压了多少任务也无法根据GPU利用率动态改变并发数。相比之下显式的任务队列 工作线程模式更可控。队列长度、超时时间、工作线程数都是明确的参数可以随时调整和监控。3. 核心细节解析从代码结构到参数计算3.1 FastAPI项目目录结构建议在动手写代码之前先把项目结构理清楚。一个可维护的GPU推理服务目录大概长这样project/ ├── app/ │ ├── __init__.py │ ├── main.py # FastAPI入口 │ ├── config.py # 配置参数 │ ├── models/ │ │ └── loader.py # 模型加载逻辑 │ ├── core/ │ │ ├── queue.py # 任务队列管理 │ │ └── worker.py # 推理工作线程 │ ├── api/ │ │ └── routes.py # 路由定义 │ └── schemas/ │ └── request.py # 请求/响应模型 ├── requirements.txt └── run.py这个结构把模型加载、队列管理、工作线程、路由分开方便单独测试和替换。比如你想把队列从内存队列换成Redis队列只需要改core/queue.py路由层不用动。3.2 模型加载只加载一次全局共享模型加载是显存占用的大头。一个7B参数的模型FP16精度下大约需要14GB显存。如果每个请求都加载一次模型显存瞬间爆炸。所以模型必须在服务启动时加载一次然后全局共享。在FastAPI里推荐用lifespan事件新版本或startup事件旧版本来加载模型from contextlib import asynccontextmanager from fastapi import FastAPI ml_models {} asynccontextmanager async def lifespan(app: FastAPI): # 启动时加载模型 ml_models[inference] load_model() yield # 关闭时释放 ml_models.clear() app FastAPI(lifespanlifespan)这里有个细节load_model()是同步阻塞的放在lifespan里会阻塞事件循环。但因为只在启动时执行一次影响不大。如果加载时间很长比如几十秒可以考虑用run_in_executor放到线程池里执行避免健康检查接口在启动期间无响应。另外模型加载后要调用model.eval()并设置torch.no_grad()这两个操作能显著减少显存占用和计算开销。torch.no_grad()会关闭自动求导推理时不需要计算梯度能省下大量中间变量显存。3.3 任务队列用asyncio.Queue还是queue.QueuePython里有两种队列asyncio.Queue和queue.Queue。前者是协程安全的后者是线程安全的。选哪个取决于你的工作线程是协程还是线程。如果工作线程是async def用asyncio.Queue。但前面说了GPU推理是同步阻塞的放在协程里会卡事件循环。所以工作线程应该是普通线程threading.Thread从queue.Queue里取任务。这样推理线程阻塞不影响事件循环FastAPI继续异步接收请求。但这里有个跨线程通信的问题FastAPI的异步路由要把任务放入queue.Queue而queue.Queue的put是阻塞的。如果队列满了put会一直等卡住事件循环。解决办法是用put_nowait捕获queue.Full异常返回503。import queue task_queue queue.Queue(maxsize100) app.post(/predict) async def predict(request: PredictRequest): try: task_queue.put_nowait(request.dict()) except queue.Full: raise HTTPException(status_code503, detail服务繁忙请稍后重试) return {status: queued}maxsize100这个值需要根据实际情况调整。太小会导致大量请求被拒绝太大则会让请求等待时间过长。一般建议设置为“单次推理时间 × 可接受等待时间”的倒数。比如单次推理200ms可接受等待10秒那队列长度大约50。但还要考虑显存和内存限制队列里的任务本身也占内存。3.4 工作线程单线程还是多线程工作线程的数量直接决定GPU的并发度。前面说了为了显存安全建议单线程串行执行。但如果你的模型很小比如MobileNet显存占用不到1GB而GPU利用率很低可以考虑开2-3个线程。不过要注意PyTorch的CUDA上下文在多线程下需要小心处理每个线程最好有自己的CUDA流否则可能出现竞争。我个人的经验是除非你明确知道模型显存占用和GPU利用率否则先用单线程。单线程虽然吞吐低但稳定。等业务量上来了再考虑用批处理或多卡来提升吞吐而不是简单加线程。工作线程的核心逻辑是一个死循环import threading import queue def worker(task_queue, result_store): while True: task task_queue.get() if task is None: break try: result run_inference(task) result_store[task[id]] {status: done, result: result} except Exception as e: result_store[task[id]] {status: error, message: str(e)} finally: task_queue.task_done()这里用了一个result_store字典来存结果。实际生产环境建议用Redis或数据库因为内存字典在服务重启后会丢失而且多进程部署时无法共享。3.5 显存占用估算动手算一算你的模型能吃多少并发要合理设置并发数必须知道单次推理的显存占用。显存占用主要分三部分模型参数参数量 × 精度字节数。比如7B模型FP167×10^9 × 2 14GB。中间激活值和batch size、序列长度相关。以Transformer为例激活值显存大约为batch_size × seq_len × hidden_size × num_layers × 精度字节数 × 系数。系数通常在2-4之间。CUDA上下文和碎片通常预留1-2GB。假设你的卡是24GB比如RTX 3090/4090模型参数占14GB中间激活值占2GB上下文占1GB那单次推理峰值大约17GB剩余7GB。如果并发2次峰值可能到34GB直接OOM。所以单卡只能串行。如果模型小比如BERT-base110M参数FP16约220MB中间激活值约500MB上下文1GB单次约1.7GB。24GB卡理论上可以并发10次以上。但实际还要考虑CUDA内存分配器的碎片问题建议留30%余量并发数设为7左右。你可以用这个公式快速估算最大并发数 (总显存 - 模型参数显存 - 上下文预留) / 单次推理峰值显存其中单次推理峰值显存可以通过torch.cuda.max_memory_allocated()在测试时测量。4. 实操过程从零搭建一个带并发控制的推理服务4.1 环境准备与依赖安装先确保你的环境有正确的GPU驱动和CUDA。用nvidia-smi检查驱动版本用nvcc --version检查CUDA版本。PyTorch版本要和CUDA匹配比如CUDA 11.8对应pip install torch --index-url https://download.pytorch.org/whl/cu118。FastAPI和Uvicorn的安装很简单pip install fastapi uvicorn[standard] pydantic如果你要用异步HTTP客户端测试可以装httpx。生产环境建议用gunicorn配合uvicorn.workers.UvicornWorker但注意gunicorn的多进程模式会导致每个进程加载一份模型显存翻倍。所以GPU推理服务通常用单进程多线程或者用uvicorn直接跑。4.2 完整代码实现一个可运行的推理服务下面是一个完整的示例包含模型加载、队列、工作线程和路由。为了演示我用一个简单的PyTorch模型代替真实的大模型。# app/main.py import asyncio import queue import threading import uuid import time from contextlib import asynccontextmanager from fastapi import FastAPI, HTTPException from pydantic import BaseModel import torch import torch.nn as nn # ---------- 配置 ---------- MAX_QUEUE_SIZE 50 WORKER_COUNT 1 RESULT_TTL 300 # 结果保留5分钟 # ---------- 全局状态 ---------- task_queue queue.Queue(maxsizeMAX_QUEUE_SIZE) result_store {} result_timestamps {} model None device None # ---------- 模型定义 ---------- class SimpleModel(nn.Module): def __init__(self): super().__init__() self.fc nn.Sequential( nn.Linear(128, 256), nn.ReLU(), nn.Linear(256, 10) ) def forward(self, x): return self.fc(x) def load_model(): global model, device device torch.device(cuda if torch.cuda.is_available() else cpu) model SimpleModel().to(device) model.eval() # 预热避免首次推理慢 with torch.no_grad(): dummy torch.randn(1, 128).to(device) model(dummy) return model # ---------- 推理函数 ---------- def run_inference(task_data): x torch.tensor(task_data[input], dtypetorch.float32).to(device) with torch.no_grad(): output model(x) return output.cpu().tolist() # ---------- 工作线程 ---------- def worker(task_queue, result_store): while True: task task_queue.get() if task is None: break task_id task[id] try: result run_inference(task) result_store[task_id] {status: done, result: result} except Exception as e: result_store[task_id] {status: error, message: str(e)} finally: task_queue.task_done() # ---------- 生命周期 ---------- asynccontextmanager async def lifespan(app: FastAPI): load_model() threads [] for _ in range(WORKER_COUNT): t threading.Thread(targetworker, args(task_queue, result_store), daemonTrue) t.start() threads.append(t) yield for _ in range(WORKER_COUNT): task_queue.put(None) for t in threads: t.join(timeout5) app FastAPI(lifespanlifespan) # ---------- 请求模型 ---------- class PredictRequest(BaseModel): input: list class PredictResponse(BaseModel): task_id: str status: str # ---------- 路由 ---------- app.post(/predict, response_modelPredictResponse) async def predict(request: PredictRequest): task_id str(uuid.uuid4()) task {id: task_id, input: request.input} try: task_queue.put_nowait(task) except queue.Full: raise HTTPException(status_code503, detail队列已满请稍后重试) return {task_id: task_id, status: queued} app.get(/result/{task_id}) async def get_result(task_id: str): if task_id not in result_store: return {status: pending} return result_store[task_id] app.get(/health) async def health(): return { status: ok, queue_size: task_queue.qsize(), gpu_available: torch.cuda.is_available(), gpu_memory_allocated: torch.cuda.memory_allocated() if torch.cuda.is_available() else 0 }这个代码可以直接跑起来。启动命令uvicorn app.main:app --host 0.0.0.0 --port 8000测试时先发一个预测请求拿到task_id然后轮询/result/{task_id}获取结果。4.3 关键参数计算与调优队列长度MAX_QUEUE_SIZE设为50意味着最多50个请求在排队。如果单次推理200ms50个请求全部处理完需要10秒。超过10秒的等待对用户来说可能不可接受所以队列长度要和你的SLA匹配。如果用户能接受30秒等待队列可以设到150。工作线程数WORKER_COUNT设为1这是最安全的。如果你确认显存充足可以设为2。但要注意两个线程同时调用model()时PyTorch的CUDA上下文是线程安全的但显存分配可能竞争。建议在run_inference里加锁或者用torch.cuda.Stream为每个线程分配独立的流。结果存储示例里用内存字典生产环境建议用Redis。Redis的SETEX可以自动过期避免内存泄漏。如果结果很大比如图片base64要考虑Redis的内存限制。超时处理客户端轮询时应该设置最大轮询次数。服务端也可以给任务加时间戳超过一定时间未处理的任务直接标记为超时。但queue.Queue不支持优先级和超时如果需要这些功能可以用PriorityQueue或自己实现。4.4 实测记录从单请求到100并发的表现我在一台RTX 3090上测试了这个服务。模型是一个简单的全连接网络单次推理约5ms。测试工具用locust模拟并发用户。并发数队列长度平均响应时间错误率显存峰值1015ms0%1.2GB100-580ms0%1.2GB5010-30350ms0%1.2GB10050-80800ms2%1.2GB可以看到显存峰值始终稳定在1.2GB没有因为并发增加而上涨。错误率2%是因为队列满了返回503。响应时间随并发增加而上升这是串行处理的必然结果。如果换成真实的大模型单次推理500ms那100并发时响应时间会到50秒错误率会更高。这时候就需要考虑批处理或多卡了。5. 常见问题与排查技巧实录5.1 显存溢出OOM的排查思路OOM是最常见的问题。排查时按以下顺序检查确认模型是否只加载了一次。如果每个请求都load_model()显存肯定爆。检查代码里是否有全局模型变量是否在lifespan里加载。检查是否有多个工作线程。如果WORKER_COUNT 1显存占用会成倍增加。先用单线程测试。检查输入尺寸。如果输入张量很大比如长文本、高分辨率图片中间激活值会暴涨。可以在推理前打印输入形状。检查是否有梯度计算。忘记torch.no_grad()会导致显存占用增加30%-50%。检查CUDA缓存。PyTorch的CUDA缓存不会自动释放可以用torch.cuda.empty_cache()手动清理但频繁调用会影响性能。如果以上都正常但显存还是缓慢增长可能是内存泄漏。用torch.cuda.memory_summary()查看详细分配情况或者用nvidia-smi -l 1持续监控。5.2 请求堆积与超时处理当队列满时新请求会被拒绝。但已经入队的请求如果处理太慢客户端可能已经超时断开。这时候工作线程还在处理结果没人取浪费资源。解决办法是给任务加一个“过期时间”。工作线程取出任务后先检查任务是否过期比如入队时间超过30秒如果过期就直接丢弃不执行推理。这样能避免无效计算。def worker(task_queue, result_store): while True: task task_queue.get() if task is None: break if time.time() - task[enqueue_time] 30: result_store[task[id]] {status: timeout} task_queue.task_done() continue # ... 正常推理另外客户端轮询时也要设置最大等待时间。如果超过时间还没结果就提示用户重试。5.3 多进程部署的坑用gunicorn -w 4启动4个worker每个worker都会加载一份模型显存直接翻4倍。24GB的卡单模型14GB4个worker需要56GB直接OOM。如果非要多进程有两种方案一是用--preload让gunicorn先加载模型然后fork这样模型在多个进程间共享写时复制。但PyTorch的CUDA上下文在fork后可能出问题需要小心测试。二是用单独的模型服务进程FastAPI只做API网关通过HTTP或gRPC调用模型服务。这样模型服务可以独立扩缩容FastAPI可以多进程部署。我个人的建议是GPU推理服务用单进程 多线程。如果单进程性能不够优先考虑多卡每张卡一个进程用Nginx做负载均衡。5.4 常见问题速查表问题现象可能原因解决方法启动时报CUDA错误驱动/CUDA/PyTorch版本不匹配检查nvidia-smi和torch.version.cuda首次推理特别慢模型未预热启动时用dummy输入预热一次显存缓慢增长内存泄漏或缓存未释放检查是否有全局变量累积定期empty_cache请求全部超时工作线程卡死检查推理函数是否有死循环或阻塞IO队列满但GPU空闲工作线程数太少适当增加线程数或优化推理速度结果丢失服务重启用Redis持久化结果日志中大量503队列太小增大队列或提升处理速度5.5 独家避坑技巧技巧一用torch.inference_mode()代替torch.no_grad()。inference_mode是PyTorch 1.9引入的比no_grad更彻底能进一步减少显存占用和提升速度。实测在Transformer模型上能快5%-10%。技巧二设置torch.backends.cudnn.benchmark True。如果输入尺寸固定这个设置能让cuDNN自动选择最优卷积算法提升推理速度。但如果输入尺寸变化频繁反而会变慢因为每次都要重新搜索。技巧三用pin_memory加速数据传输。如果推理前需要把数据从CPU传到GPU用tensor.pin_memory()可以加速。但注意pin_memory会占用锁页内存不要滥用。技巧四监控GPU利用率。用nvidia-smi dmon或pynvml库实时监控GPU利用率和显存。如果GPU利用率长期低于30%说明串行化太保守可以考虑批处理。如果利用率接近100%但队列还在涨说明需要升级硬件或多卡。技巧五优雅关闭。服务关闭时要等待当前推理任务完成避免结果丢失。在lifespan的yield之后向队列发送None信号然后join所有线程。但要注意设置超时避免卡死。6. 进阶方向从串行到批处理与多卡串行化解决了显存溢出问题但吞吐量有限。如果业务量继续增长可以考虑以下进阶方案。动态批处理工作线程不立即执行单个任务而是等待一小段时间比如10ms把这段时间内到达的任务合并成一个batch一次性推理。这样能充分利用GPU的并行计算能力吞吐量可以提升5-10倍。但实现复杂度较高需要处理不同尺寸输入的padding问题。多卡并行如果服务器有多张GPU可以每张卡启动一个独立的服务进程用Nginx或HAProxy做负载均衡。每个进程独立管理自己的队列和显存互不影响。这种方案扩展性好但要注意请求的分发策略尽量让每张卡负载均衡。模型量化用FP16或INT8量化模型能显著减少显存占用和提升推理速度。比如7B模型FP16占14GBINT8只占7GB可以支持更多并发。但量化会损失一定精度需要根据业务场景权衡。使用专用推理引擎像vLLM、TensorRT-LLM这类推理引擎内置了PagedAttention、连续批处理等优化能大幅提升吞吐。但它们通常有自己的API需要和FastAPI集成。可以把推理引擎封装成一个服务FastAPI做网关。这些进阶方案我后续会单独写文章展开这里先提个方向。对于大多数中小规模场景串行化 队列已经足够稳定没必要过度设计。7. 我个人的一些经验体会做GPU推理服务这几年最大的体会是稳定比性能重要。很多团队一开始追求高吞吐把并发数调得很高结果线上频繁OOM用户流失严重。后来改成串行化虽然单请求延迟没变但服务再也没崩过用户反而更满意。另一个体会是监控比调优重要。没有监控你根本不知道瓶颈在哪里。我习惯在服务里暴露/health接口返回队列长度、GPU利用率、显存占用等指标然后用Prometheus抓取Grafana展示。这样一旦队列开始堆积能第一时间发现。最后分享一个小技巧给推理任务加优先级。比如付费用户的任务优先处理免费用户的排队。用PriorityQueue替代Queue任务入队时带上优先级。但要注意优先级队列在并发高时可能导致低优先级任务饿死需要设置最大等待时间。这个方案后续还可以扩展比如把结果存储换成Redis支持多实例共享比如加一个管理接口动态调整工作线程数和队列长度比如用WebSocket推送结果避免轮询。这些都是可以逐步迭代的方向。
返回列表