SFT 硬件资源怎么配?
GPU 卡数、机器数需求分析,影响资源的关键因素
原题:在进行大模型监督微调(SFT)时,典型的硬件资源配置是怎样的?例如需要多少张GPU卡、多少台机器,影响资源需求的主要因素有哪些?
模型训练 · 阿里真题
回答与解析
资源估算必须先固定训练方案
先明确模型参数量、全量微调还是 PEFT、权重/梯度/优化器精度、序列长度、micro-batch、梯度累积、激活检查点、并行策略和目标吞吐。脱离这些条件回答“需要几张卡”没有意义。
主要显存项
- 权重、梯度、优化器状态及可能的 FP32 master weights;
- 激活与 attention 中间量,标准稠密 attention 对序列长度含二次项;
- 通信 buffer、临时 workspace 和框架碎片。
全量 Adam 训练通常由训练状态和激活主导,可用 ZeRO/FSDP 分片参数、梯度和优化器状态,并用 activation checkpointing、FlashAttention 等减少激活。普通 LoRA 冻结基座但仍需以所选精度加载它;QLoRA 才是把冻结基座以 4-bit 存储并训练 LoRA adapter。QLoRA 论文的单卡 48GB、65B 结果不能外推为普通 LoRA 或所有 70B 配置。
多机时,tensor/pipeline/data parallel 的选择取决于单层是否放得下、网络拓扑与 batch。Pipeline parallelism 存在 bubble,调度和 micro-batch 只能减少而不能保证消除。最终应先在单机小规模 profile 峰值显存和 tokens/s,再按通信效率扩展,并预留 checkpoint、评测和故障恢复资源。
口语版讲法(约3分钟)
- 先固定训练配置再算卡数
- 拆分训练状态和激活显存
- 区分全量、LoRA 与 QLoRA
- 选择分片并行并实测扩展效率
这道题如果直接报“几张 A100”其实是不合格的,因为资源取决于一组配置。我会先问清模型参数量、全量微调还是 PEFT、权重和优化器精度、序列长度、micro-batch、梯度累积、激活检查点以及目标训练时长。只有这些条件固定以后,卡数才可估算。
显存要拆开算。第一类是模型训练状态,包括权重、梯度、Adam 的一阶和二阶矩,以及某些混合精度方案里的 FP32 master weights。第二类是激活,受层数、hidden size、batch 和序列长度影响;标准稠密 attention 还有序列长度二次项,不能笼统说近似线性。第三类是通信 buffer、kernel workspace 和碎片。
训练方式差异也很大。全量微调要更新全部参数,通常需要 ZeRO 或 FSDP 把参数、梯度和优化器状态分片,再结合 activation checkpointing 与 FlashAttention 控制激活。普通 LoRA 虽然只更新低秩 adapter,但冻结基座仍要按所选精度装入显存,反向也要穿过基座。QLoRA 才是把冻结基座以 4-bit 形式加载,再训练 LoRA。论文里单张 48GB 卡微调 65B 是特定 QLoRA 配置,不能改写成普通 LoRA 单卡能跑所有 70B 模型。
如果单层都放不下一张卡,需要 tensor parallel;层能放下但整模放不下,可以考虑 pipeline parallel;数据并行用来扩大吞吐。Pipeline 会有 bubble,增加 micro-batch 和改调度只能减少,不能保证消失,多机还要看互联带宽。
我的落地流程是先用目标序列长度和 micro-batch 在一台机器 profile 峰值显存、tokens/s 和通信占比,再外推不同 ZeRO/FSDP stage 和节点数,最后留出 checkpoint、在线评测与故障恢复空间。这样给出的硬件表会明确前提和误差范围,而不会虚构某家公司用了多少卡。
估算时我会先做一张逐项预算表,而不是只算参数乘字节数。表里分别列出常驻权重、可训练参数对应的梯度和优化器状态、峰值激活、临时通信与算子工作区,再标明哪些会被分片、量化或重计算。随后用一个短跑实验校准框架额外开销和碎片,因为理论和真实峰值之间往往还有差距。
卡数决定后还要验证吞吐是否值得。某个方案显存能放下,不代表扩到多机就高效;如果通信占比持续上升,增加设备可能只缩短很少时间。我的交付结果会同时给最小可运行配置、满足目标工期的推荐配置,以及序列长度或 batch 改变后的敏感性。这样面试官看到的是可复算的容量规划,而不是脱离前提的硬件数字。
容量表还要给突发峰值留余量,例如保存 checkpoint、切换评测或通信重叠时的临时占用。这个余量应由实测峰值和故障恢复流程决定,不能简单写成统一比例。若配置只在理想情况下刚好放下,长序列样本或负载波动就可能频繁溢出。
关键一句:参数状态、激活和通信的主导项会随精度、序列长度及并行策略改变,卡数必须由实测成本模型反推。
核验来源
面试官还可能这样问
- 问法 1 · 场景切入
假设你要给电商客服微调一个70B模型做长文本回复,现有8张A100 80G,你觉得够用吗?如果不够,你会怎么调整?
- 问法 2 · 层层追问
你做过SFT吗?通常需要多少资源?……比如7B和70B模型,配置差多少?……影响显存的主要因素有哪些?
- 问法 3 · 直球架构
给定一个70B模型,做全量SFT,序列长度4k,batch size 32,用FP16。请估算所需GPU数量和节点数,并说明主要影响因素。