本节摘要:训练前先算账——"这个模型要用什么显卡、训多久、花多少钱"。本节讲清显存需求的估算方法、训练时长的量级、CPU/GPU/多卡的取舍,以及"先算账再开跑"的工程习惯。
阅读完本节,你应当能够:
"跑 124M 模型要什么显卡?"——不回答这个问题就开跑,大概率中途 OOM。训练硬件的账其实有规律可循:显存看模型大小与批次,时长看算力与数据量。先算清这两笔账,选择自然清晰。
这套"算账"能力可以迁移到任何模型:看到参数量和训练数据量,就能大致估出需要的显卡和天数,避免盲目开跑。
粗略估算:模型参数 × 12-16 字节(训练时)——124M 参数约需 1.5-2GB,训练批次的激活另算。
# 估算训练显存(不含激活) def estimate_memory_gb(num_params, bytes_per_param=16): # fp32 训练:权重 4 + 梯度 4 + Adam 状态 8 = 16 字节/参数 return num_params * bytes_per_param / 1024**3 # 各规模模型的训练显存估算 for params in [10e6, 124e6, 1e9]: print(f"{params/1e6:.0f}M 参数:约 {estimate_memory_gb(params):.1f} GB") # 输出:10M → 0.2 GB;124M → 1.9 GB;1B → 16 GB
这个 16 字节/参数的系数是"经验账本":权重、梯度各 4 字节,Adam 优化器再存一阶二阶动量各 4 字节。使用混合精度(bf16)可以把权重和梯度减半,总需求降到约 8-12 字节/参数——这也是 Modded-NanoGPT 提速的关键之一。
| 场景 | 量级 |
|---|---|
| CPU 小模型(10M 级) | 分钟到小时 |
| 单卡 GPU 小模型 | 分钟到小时 |
| 单卡 GPU 124M | 数天 |
| 多卡 124M | 天级(并行加速) |
| 模型规模 | 权重显存 | 训练估算 |
|---|---|---|
| 10M | 40MB | 几百 MB |
| 100M | 400MB | 2-4GB |
| 124M | 500MB | 2-4GB+ |
| 1B | 4GB | 16GB+ |
💡 关键直觉:显存不够的通用解法:降批次、降精度、降模型——三个"降"按顺序试,多数场景能解决。
| 预算 | 方案 |
|---|---|
| 无 GPU | CPU 跑小模型验证流程 |
| 消费级 GPU | 小模型 + 低精度 |
| 多卡 | 数据并行训练 |
| 云 GPU | 按需租用 |
数据并行:每卡一份模型副本,各训一批 梯度同步:每步交换梯度,参数保持一致 效果:近似"批次变大",训练加速
NanoGPT 的 train.py 内置了 PyTorch 的 DDP 支持,多卡训练只需在启动命令里加 torchrun 参数:
# 4 卡数据并行训练 torchrun --standalone --nproc_per_node=4 train.py config/train_gpt2.py
DDP 的核心机制是"梯度同步":每张卡算完自己的梯度后,AllReduce 汇总平均,再同步更新——模型参数在各卡始终一致,等价于把批次扩大了 4 倍。
显存:模型 + 批次能否放下 时长:按算力估算天数 预算:云 GPU 时薪 × 时长 价值:这次训练值不值这个成本
| 手段 | 效果 | 代价 |
|---|---|---|
| 混合精度 bf16 | 显存减半、速度提升 | 需新 GPU 支持 |
| 梯度累积 | 等效大批次 | 无额外显存开销 |
| Flash Attention | 注意力加速省显存 | 需对应实现 |
| 数据并行 | 多卡线性提速 | 需多卡环境 |
⚠️ 常见坑:不估算就开跑。跑到一半 OOM 或发现要跑两周——先花五分钟算账,能避免大量浪费。
除了显存,另一笔账是"每秒能处理多少 token"。训练时长 ≈ 总 token 数 ÷ 每秒吞吐。粗略量级:
| 设备 | 10M 模型吞吐 | 124M 模型吞吐 |
|---|---|---|
| CPU(多核) | 几千 token/秒 | 几百 token/秒 |
| 消费级 GPU | 几万 token/秒 | 几千 token/秒 |
| A100 | 十万+ token/秒 | 数万 token/秒 |
用莎士比亚数据(约 100 万 token)估算:消费级 GPU 训 10M 模型只要几百秒,训 124M 模型则要数小时。这个数量级提前知道,安排实验就不慌。
# 估算一次训练的总时长 def estimate_hours(tokens_total, tokens_per_sec): return tokens_total / tokens_per_sec / 3600 # 示例:124M 模型、OpenWebText 规模约 100 亿 token # 单张 A100 约 3 万 token/秒 hours = estimate_hours(10e9, 30000) print(f"单卡估算 {hours:.0f} 小时") # 约 93 小时,接近 4 天 # 8 卡并行理论上缩短到约 12 小时(实际受通信等影响约 1-2 天)
这个计算和官方"8×A100 约 4 天"互相印证,说明数量级估算是可靠的。你不需要精确建模,能算出"天级还是小时级"就够了。
| 方案 | 优点 | 注意点 |
|---|---|---|
| 按小时租 A100 | 强算力、灵活 | 贵,要盯时间 |
| 租消费级卡 | 便宜、够用小模型 | 显存有限 |
| 包月整机 | 适合长期学习 | 空闲也计费 |
选择原则:先估"这次实验要跑多久",再算"哪档价格最划算"。跑一次 2 小时的小实验租 A100 纯属浪费;跑 124M 复现则值得上多卡。
成本算清了,下一节学怎么提速——优化技术与策略。