【Python深度学习系列】模型训练与推理过程中内存、CPU、显存各自承担的关键任务和瓶颈点总结
这是我的第453篇原创文章。
一、引言
本文总结一下模型训练与推理过程中内存、CPU、显存各自承担的关键任务和瓶颈点。
二、实现过程
2.1 数据加载与预处理(CPU + RAM 为主)
吃内存(RAM):
- 数据集大(如图像、视频、音频)会加载到内存或映射到内存(
mmap)。 - 多进程数据加载(
num_workers > 0)会复制多份数据,内存飙升。 - 数据增强(如 Albumentations、OpenCV)会生成临时图像,内存峰值高。
吃CPU:
- 图像解码(JPEG/PNG)、resize、normalize、数据增强(旋转、裁剪、颜色变换)全是CPU活。
num_workers设置越高,CPU占用越高(建议设为 CPU核心数 × 0.5 ~ 0.8)。
吃显存:
- 几乎不吃,数据还没进GPU。
✅ 优化建议:
-
用
pin_memory=True加速数据从RAM到GPU的传输。 -
用
persistent_workers=True避免重复进程开销。 -
对大数据集用
memory_map=True或流式加载(如 WebDataset)。
2.2 前向传播(Forward):显存大头(激活值)
显存占用主要来源:
- 激活值(activations):每层输出的特征图(feature maps)。
- 激活值显存 ∝ Batch Size × 序列长度 × 隐藏层维度 × 层数。
-
举例:
-
ViT-B/16 输入 224×224,Batch=32,中间激活值可能占 数GB。
-
LLM 中,序列长度 4096,Batch=4,激活值显存可能 >20GB。
-
✅ 优化建议:
减小
batch_size(最直接)。用 梯度检查点(Gradient Checkpointing):以时间换空间,激活值不保存,反向时重计算。
Hugging Face:
model.gradient_checkpointing_enable()PyTorch:
torch.utils.checkpoint.checkpoint()用 FlashAttention(LLM 必备):减少注意力机制显存和计算量。
2.3 反向传播(Backward)
显存爆炸(梯度 + 优化器状态)
显存占用来源:
-
Adam:保存 动量 + 方差,×2 倍参数显存。
-
AdamW + FP16:显存 ≈ 参数 × 2(模型) + 参数 × 2(梯度) + 参数 × 4(Adam状态) = 6倍参数显存!
- SGD:无额外状态。
-
梯度(gradients):与参数同量级。
-
优化器状态:
-
-
举例:
-
1B 参数模型,FP16 + AdamW,训练时显存 ≈ 6GB × 3 = 18GB(粗略)。
-
✅ 优化建议:
-
用 AdamW + FP16 + DeepSpeed ZeRO Stage 2/3:把优化器状态分片到多卡,显存线性减少。
-
用 LoRA / QLoRA:只训练低秩适配器,显存骤降(QLoRA 可训练 70B 模型在 48GB 卡上)。
-
用 8-bit 优化器(如
bitsandbytes.optim.AdamW8bit):Adam 状态压缩到 1 byte,显存减半。
2.4 模型初始化
模型初始化(显存为主,把权重从磁盘读进内存/显存,创建网络对象)
显存占用公式(粗略估算):显存≈参数量×精度字节数
|
精度类型 |
字节数 |
示例:1B参数模型 |
|---|---|---|
|
FP32 |
4 bytes |
≈ 4 GB |
|
FP16 |
2 bytes |
≈ 2 GB |
|
INT8 |
1 byte |
≈ 1 GB |
-
例如:LLaMA 7B 模型,FP16 精度,初始化就占 ~14GB 显存(7B × 2 bytes)。
✅ 优化建议:
-
用
torch.cuda.empty_cache()清理未用显存。 -
用
model.half()或bfloat16降低精度。 -
用
device_map='auto'(Hugging Face)自动分配多卡。
2.5 推理
Inference,把真实音频/图像/text 送进网络,做一次前向传播,拿到输出):显存低于训练,但仍需关注
显存占用:
- 只有 模型参数 + 输入数据 + 激活值(无梯度、无优化器)。
- 可粗略估算为:2 × 模型参数显存(FP16 时)。
内存(RAM):
- 加载模型文件(
.bin、.safetensors)到内存。 - 多线程推理时,CPU 占用可能高(如 LLM 解码)。
✅ 优化建议:
-
用 模型编译(torch.compile())加速推理。
-
用 KV Cache(LLM 必备):缓存键值对,避免重复计算。
-
用 INT8/INT4 量化(GPTQ、AWQ、BitsAndBytes):显存减半甚至更多。
作者简介:
读研期间发表6篇SCI数据挖掘相关论文,现在某研究院从事数据算法相关科研工作,结合自身科研实践经历不定期分享关于Python、机器学习、深度学习、人工智能系列基础知识与应用案例。致力于只做原创,以最简单的方式理解和学习,关注我一起交流成长。需要数据集和源码的小伙伴可以关注底部公众号添加作者微信。
更多推荐



所有评论(0)