百川 AI 对话服务、微调任务、两种负载的显存账要分开算

文章导读
同一张卡上既要跑百川对话服务,又要跑微调任务,最容易踩的坑是把两种负载的显存按同一套账去估:推理时看着还剩不少,一开训练就 OOM;或者训练能起来,一接线上请求就掉。推理和微调的显存构成几乎没有重叠项,推理的大头是键值缓存,微调的大头是优化器状态和激活值。先把两份账分别列出来,再决定是共存、限流还是排队。
📋 目录
  1. 拆出推理侧的显存构成
  2. 拆出微调侧的显存构成
  3. 给出参数量与精度的估算骨架
  4. 用命令实测两种负载的峰值占用
  5. 决定两种负载的共存或排队方式
A A

同一张卡上既要跑百川对话服务,又要跑微调任务,最容易踩的坑是把两种负载的显存按同一套账去估:推理时看着还剩不少,一开训练就 OOM;或者训练能起来,一接线上请求就掉。推理和微调的显存构成几乎没有重叠项,推理的大头是键值缓存,微调的大头是优化器状态和激活值。先把两份账分别列出来,再决定是共存、限流还是排队。

先按推理和微调各拆一份显存预算:推理按权重、KV 缓存、并发请求三块估,微调按权重、梯度、优化器状态、激活值四块估。适用场景是单卡、显存紧张、两种负载都要保留。操作动作是先用参数量乘每参数字节数算骨架,再加余量系数,最后用按秒采样的显存命令核对两种负载的峰值时段。验证方式是看两条显存曲线是否在同一时刻叠加。风险边界在于显存碎片和框架开销会让估算偏低,估算值只能当参考下限,实际以采样峰值为准。

拆出推理侧的显存构成

百川对话服务在推理时的显存通常由三部分组成:模型权重、键值缓存(KV cache)、以及框架运行时开销(CUDA context、临时缓冲区、并行通信缓冲)。权重是固定项,加载完基本不动;KV 缓存是变动项,随序列长度和并发请求数增长;运行时开销随框架和并行策略变化,需要单独留一块。三者的作用是:权重决定模型能不能装进去,KV 缓存决定能同时接多少请求、能放多长上下文。

KV 缓存可以直接按公式估:

每 token 的 KV 字节 ≈ 2 × 层数 × KV 头数 × head_dim × 每元素字节数
KV 缓存总量 ≈ 每 token 字节 × 平均序列长度 × 并发请求数

其中 2 表示 key 和 value 各一份;
每元素字节数按精度取,fp16/bf16 通常为 2,int8 量化为 1。

从式子里能看出方向:序列长度和并发数都是线性放大因子,最大上下文从 4k 提到 8k,KV 缓存大致翻倍;并发从 4 提到 16,也大致是四倍。给对话服务做预算时,建议用“最长上下文 × 预期最大并发”算一个上限,再用“平均上下文 × 日常并发”算一个常态值,两个数都要心里有底。启动参数里通常会有最大序列长度、最大并发或最大运行请求数一类的开关,不同推理框架命名不同,需要按实际框架确认后再调。

拆出微调侧的显存构成

微调侧的显存由四部分组成:权重、梯度、优化器状态、激活值。前三项与参数量和精度直接相关,第四项与批大小、序列长度、是否开启梯度重计算(gradient checkpointing)相关。

百川 AI 对话服务、微调任务、两种负载的显存账要分开算
  • 权重:加载的模型参数,全参微调时通常还要留一份 fp32 主权重副本。
  • 梯度:与可训练参数量同量级,冻结的参数不产生梯度。
  • 优化器状态:用 Adam/AdamW 时,每个可训练参数要存一阶动量和二阶动量,通常按 fp32 存,优化器状态往往是权重本身的好几倍。
  • 激活值:前向过程中为反向传播保留的中间结果,随批大小和序列长度增长。

批大小加大以后,最先吃紧的通常是激活值,因为激活值与批大小近似成正比,而权重、梯度、优化器状态在同一轮训练里基本不变。做法上可以先加大批大小并配合梯度累积观察峰值,或者开启梯度重计算用时间换显存;如果做的是 LoRA 这类参数高效微调,梯度和优化器状态按可训练参数量算,规模小很多,此时瓶颈更容易落在权重和激活值上。

给出参数量与精度的估算骨架

换模型或换精度时,先用这条通用骨架重估,不求精确,只求量级不出错:

权重显存 ≈ 参数量 × 每参数字节数
每参数字节数参考:fp32 = 4,fp16/bf16 = 2,int8 ≈ 1,int4 ≈ 0.5

全参微调(Adam 类优化器 + 混合精度)粗估:
权重 + 梯度 + 优化器状态 ≈ 参数量 × (12 ~ 18) 字节
落点取决于是否有 fp32 主权重、梯度按什么精度存。

LoRA 等参数高效微调:
权重部分按原模型算,梯度与优化器状态按可训练参数量算。

推理侧:
权重量 + KV 缓存上限 + 运行时开销

估算结果只能当下限参考,建议再乘一个余量系数,通常取 1.1 到 1.3,用来吸收显存碎片、CUDA context、通信缓冲和临时张量。余量取多少需要结合环境确认:并行度越高、序列越长、框架越复杂,留得越宽。如果按估算加余量已经贴近卡的上限,就不要指望靠调批大小硬挤,换方案更省事。

用命令实测两种负载的峰值占用

估算是为了判断能不能做,实测是为了验证能不能稳定跑。最直接的方式是按秒采样显存占用并打时间戳,便于和训练日志、服务日志对齐:

百川 AI 对话服务、微调任务、两种负载的显存账要分开算
# 按秒记录显存占用,Ctrl+C 结束
while true; do
  echo "$(date +%H:%M:%S) $(nvidia-smi `--query-gpu`=memory.used `--format`=csv,noheader,nounits) MiB"
  sleep 1
done

# 或者用 nvidia-smi 自带的采样模式
nvidia-smi dmon -s mu -d 1

采样时机要覆盖几个关键点:模型加载完成但未接请求时(得到权重加运行时开销的基线)、首批请求带长上下文 prefill 时(推理侧峰值通常出现在这里)、训练 step 0、训练若干步之后、以及保存 checkpoint 时。训练侧还可以在代码里读 PyTorch 的峰值统计,reserved 通常比 allocated 更接近真实占用:

torch.cuda.reset_peak_memory_stats()
# ... 跑一个完整的 step ...
print(torch.cuda.max_memory_allocated() / 1024**3, 'GiB')
print(torch.cuda.max_memory_reserved() / 1024**3, 'GiB')

判断峰值出现在哪一步,看采样曲线和日志时间戳是否对齐:对话服务的峰值一般在长序列请求的 prefill 阶段或并发最高的时刻;微调的峰值一般在反向传播到优化器 step 之间,保存检查点时如果额外做格式转换或权重合并,也可能出现一个次高峰。把两条曲线叠在同一时间轴上,就能判断共存时会不会撞在一起。

决定两种负载的共存或排队方式

实测完峰值,通常有三种处理方式,按代价从低到高排:

  1. 错峰运行:训练放在对话低峰期,服务保持在线。适用条件是两者峰值之和仍有余量,且训练可以随时暂停。验证方式是把两条显存曲线叠起来看是否有交集,并在训练进行中发请求确认服务响应正常。
  2. 限制并发:把服务的最大并发请求数、最大上下文长度调小,给训练腾出固定空间。适用条件是业务能接受排队或超时。验证方式是压到预期并发上限后看显存是否稳定,以及队列等待时间是否落在可接受范围。
  3. 切换设备:把训练和服务拆到不同卡,或者在同一张卡上分时启停。适用条件是手上有第二张卡,或训练任务可以离线跑。验证方式是确认服务进程和训练进程的可见设备互不重叠,切换前后各做一次峰值采样。

如果三种方案都试过仍然顶到上限,就不要用扩大虚拟内存之类的止血手段硬扛,那类做法会拖慢训练速度,属于临时措施而不是性能方案。更实际的选择是缩小模型规模、缩短训练序列、或者把其中一种负载挪到别的机器上。