公司动态

小红书 算法一面 九

📅 2026/8/20 22:41:08
小红书 算法一面 九
显存占用和哪些因素有关? 算一下一个7B的模型推理需要多少显存训练呢?显存占用是深度学习模型训练和推理的核心瓶颈其大小由**模型参数、激活值、KV缓存、优化器状态、梯度**等多个核心因素共同决定且**训练阶段的显存需求远高于推理阶段**训练需额外存储梯度、优化器状态等。下面先拆解显存占用的关键影响因素再以7B70亿参数大模型为例分推理和训练场景计算理论显存占用。## 一、显存占用的核心影响因素显存占用可分为 **静态开销**与输入无关和 **动态开销**与输入、batch size、序列长度强相关推理和训练的开销构成差异显著。| 开销类型 | 核心组成 | 推理阶段是否需要 | 训练阶段是否需要 | 关键影响因子 ||----------|----------|------------------|------------------|--------------|| **静态开销** | 模型参数权重、偏置 | ✅ 必须 | ✅ 必须 | 参数数量、数据精度FP32/FP16/INT8等 || **动态开销-推理特有** | KV缓存注意力机制的Key/Value缓存 | ✅ 必须 | ❌ 无需训练时KV随计算生成不缓存 | batch size、序列长度、KV头数、头维度 || **动态开销-训练特有** | 梯度 | ❌ 无需 | ✅ 必须 | 与参数同大小、数据精度 || **动态开销-训练特有** | 优化器状态 | ❌ 无需 | ✅ 必须 | 优化器类型AdamW/SGD、数据精度 || **动态开销-通用** | 激活值各层前向计算的中间结果 | ✅ 少量 | ✅ 大量 | batch size、序列长度、模型层数、隐藏维度 || **其他开销** | 框架临时变量、数据存储、显存碎片 | ✅ 少量 | ✅ 少量 | 深度学习框架PyTorch/TensorFlow、硬件架构 |### 关键因子详解1. **数据精度**是影响所有静态/动态开销的核心因子不同精度的单参数字节数如下| 数据精度 | 单参数字节数 | 典型适用场景 ||----------|--------------|--------------|| FP32单精度浮点数 | 4 | 模型训练基线、高精度需求场景 || FP16半精度浮点数 | 2 | 混合精度训练/推理、主流选择 || BF16脑浮点数 | 2 | 大模型训练、对溢出更鲁棒 || INT88位整数 | 1 | 量化推理、低显存部署 || INT44位整数 | 0.5 | 极致量化推理、精度略有损失 |2. **KV缓存**仅推理阶段需要用于加速自回归生成避免重复计算已生成token的K/V大小公式为$$KV_{显存} batch\_size \times seq\_len \times n_{kv\_heads} \times head\_dim \times 2 \times dtype\_size$$- $n_{kv\_heads}$KV头数MQA1GQA分组头数MHA与Q头数相同- 系数2代表K和V两个向量- 推理时KV缓存随序列长度线性增长是长上下文推理的主要显存开销。3. **优化器状态**训练阶段特有不同优化器的状态大小差异极大- **SGD**无额外状态仅需存储参数梯度2倍参数大小- **AdamW**需存储一阶矩m、二阶矩v共2倍参数大小总开销为参数梯度优化器状态**4倍参数大小**大模型训练的主要静态开销4. **激活值**训练阶段的核心动态开销大小与模型结构强相关公式简化为$$激活值_{显存} \approx batch\_size \times seq\_len \times d_{model} \times num\_layers \times dtype\_size$$- 可通过**激活重计算Activation Checkpointing** 减少50%以上的激活显存以牺牲计算速度为代价## 二、7B模型的显存计算理论值以**LLaMA-7B结构**为参考基准模型核心参数如下- 参数总量$N 7 \times 10^9$70亿- 隐藏维度$d_{model}4096$- 注意力头数$n_{heads}32$KV头数GQA$n_{kv\_heads}8$- 头维度$head\_dim d_{model}/n_{heads}128$- 模型层数$num\_layers32$### 2.1 推理阶段显存计算推理显存 **模型参数显存 KV缓存显存**其中模型参数显存占主导地位。#### 步骤1计算模型参数显存参数显存公式$Param_{显存} N \times dtype\_size$| 数据精度 | 单参数字节 | 理论参数显存 | 备注 ||----------|------------|--------------|------|| FP32 | 4 | $7e9 \times 4 28,000 \ MB 28 \ GB$ | 极少使用精度过剩 || FP16/BF16 | 2 | $7e9 \times 2 14,000 \ MB 14 \ GB$ | 主流推理精度 || INT8 | 1 | $7e9 \times 1 7,000 \ MB 7 \ GB$ | 量化推理精度损失小 || INT4 | 0.5 | $7e9 \times 0.5 3,500 \ MB 3.5 \ GB$ | 极致量化适合边缘设备 |#### 步骤2计算KV缓存显存以**典型推理配置**为例$batch\_size1$$seq\_len2048$$dtypeFP16$2字节代入公式$$\begin{align*}KV_{显存} 1 \times 2048 \times 8 \times 128 \times 2 \times 2 \\ 1 \times 2048 \times 8 \times 128 \times 4 \\ 8,388,608 \ Bytes \\ 8.192 \ MB\end{align*}$$#### 推理总显存理论值| 数据精度 | 参数显存 | KV缓存显存seq_len2048 | 总理论显存 | 实际显存含框架开销 ||----------|----------|-----------------------------|------------|------------------------|| FP16 | 14 GB | 8.2 MB | ~14.01 GB | ~15-16 GB || INT8 | 7 GB | 4.1 MB | ~7.004 GB | ~8-9 GB | 注意长上下文场景下KV缓存占比会提升例如 $seq\_len128K$ 时FP16的KV缓存为 $1 \times 128000 \times 8 \times 128 \times 4 524 \ MB$总显存仍以参数为主。### 2.2 训练阶段显存计算训练显存 **参数显存 梯度显存 优化器状态显存 激活值显存**其中**优化器状态和激活值**是主要开销。我们以**主流训练配置**为例优化器AdamW精度FP16$batch\_size4$$seq\_len2048$不使用激活重计算。#### 步骤1计算静态训练显存参数梯度优化器状态AdamW优化器下静态显存 $4 \times Param_{显存}$参数1倍 梯度1倍 优化器状态2倍FP16精度下$$静态训练显存 4 \times 14 \ GB 56 \ GB$$#### 步骤2计算激活值显存代入激活值简化公式$$\begin{align*}激活值_{显存} 4 \times 2048 \times 4096 \times 32 \times 2 \\ 4 \times 2048 \times 4096 \times 64 \\ 214,748,364,800 \ Bytes \\ 204,800 \ MB 200 \ GB\end{align*}$$#### 训练总显存理论值FP16精度、AdamW优化器、batch_size4、seq_len2048下$$总训练显存 56 \ GB 200 \ GB 256 \ GB$$### 训练显存优化技术的影响上述理论值是**无优化的极端情况**实际训练时通过以下技术可大幅降低显存占用| 优化技术 | 显存降低比例 | 核心原理 ||----------|--------------|----------|| 激活重计算Checkpointing | 50%-70% | 舍弃部分激活值反向传播时重新计算 || 梯度累积 | 按累积步数等比例降低 | 拆分batch模拟大batch训练降低单步激活显存 || 模型并行 | 按层数拆分模型到多GPU | 每个GPU仅存储部分层的参数/激活 || ZeRO优化ZeRO-1/2/3 | 30%-90% | 拆分优化器状态/梯度/参数到多GPU避免冗余存储 |以**ZeRO-2 激活重计算**为例7B模型训练显存可降至 **20-30 GB/单GPU**满足消费级旗舰GPU如RTX 4090 24GB的训练需求。## 三、核心结论1. **推理显存**以模型参数显存为主KV缓存占比极低长上下文除外量化是降低推理显存的最优方案INT8可减半显存。2. **训练显存**远高于推理核心开销是激活值和AdamW优化器状态需依赖激活重计算、ZeRO、模型并行等技术压缩显存。3. **7B模型典型显存需求**- 推理FP16~15 GB/单GPU- 训练FP16ZeRO-2Checkpointing~25 GB/单GPU遇到过灾难性遗忘吗?怎么缓解的# 灾难性遗忘Catastrophic Forgetting及缓解方法**灾难性遗忘**是深度学习模型在**连续学习新任务**时的核心痛点模型在学习新任务的过程中会快速遗忘之前已掌握的旧任务知识导致旧任务的性能急剧下降。这种现象的本质是**参数共享冲突**——深度学习模型的参数是共享的学习新任务时的梯度更新会覆盖旧任务的关键参数使模型参数从旧任务的最优解向新任务的最优解偏移最终丢失旧任务的知识。该问题在持续学习终身学习、增量学习、大模型多任务微调等场景中尤为突出下面先分析成因再详细介绍主流的缓解方法。## 一、灾难性遗忘的核心成因1. **参数共享机制**深度学习模型如Transformer、CNN的参数是全局共享的旧任务和新任务依赖同一组参数完成预测。学习新任务时梯度下降会驱动参数朝着新任务的损失最小化方向更新从而覆盖旧任务的关键参数配置。2. **数据分布偏移**新旧任务的数据分布差异越大参数更新的冲突越剧烈。例如用预训练大模型先学习“文本分类”再学习“机器翻译”两类任务的目标函数和数据分布差异极大翻译任务的参数更新会严重破坏分类任务的知识。3. **梯度更新的偏向性**新任务的训练数据通常是“当前可用”的模型在训练时会优先拟合新数据梯度更新的方向完全由新任务损失主导没有机制约束参数远离旧任务的最优解。## 二、缓解灾难性遗忘的主流方法缓解灾难性遗忘的核心思路分为三类**约束参数更新**、**隔离新旧任务参数**、**保留旧任务知识**下面分别介绍各类方法的原理、典型算法和适用场景。### 1. 基于正则化的方法约束关键参数的更新这类方法的核心是**识别旧任务的关键参数**并在学习新任务时对这些参数的更新施加正则化约束防止其被大幅修改。#### 1.1 弹性权重整合EWC, Elastic Weight ConsolidationEWC是最经典的正则化方法核心步骤如下- **步骤1评估参数重要性**在完成旧任务训练后计算每个参数对旧任务损失的贡献度常用**Fisher信息矩阵**衡量——Fisher矩阵的对角线元素越大说明该参数对旧任务越关键。- **步骤2添加参数更新正则**学习新任务时在损失函数中加入正则项限制关键参数的更新幅度。新任务的总损失为$$\mathcal{L}_{total} \mathcal{L}_{new}(\theta) \lambda \sum_i \frac{F_i}{2} (\theta_i - \theta_i^*)^2$$其中- $\theta_i^*$ 是旧任务训练完成后的参数值- $F_i$ 是Fisher矩阵的对角线元素代表参数 $i$ 的重要性- $\lambda$ 是正则强度系数越大对参数约束越强。**优势**无需存储旧任务数据仅需存储参数重要性和旧参数值**局限**仅适用于少量任务的连续学习任务过多时Fisher矩阵存储和计算成本会急剧上升。#### 1.2 突触智能SI, Synaptic IntelligenceSI是EWC的改进版核心是**跟踪参数在旧任务中的“贡献轨迹”**而非依赖Fisher矩阵。它通过计算参数在旧任务训练过程中的累积变化量衡量参数的重要性然后在新任务训练时施加正则约束。**优势**无需在旧任务结束后单独计算Fisher矩阵可在线跟踪参数重要性**局限**正则强度的调参难度较高对任务顺序敏感。### 2. 基于参数隔离的方法避免新旧任务参数冲突这类方法的核心是**为不同任务分配独立的参数空间**让新旧任务的参数互不干扰从根本上解决参数共享冲突。该方法是**大模型多任务微调的主流方案**。#### 2.1 参数高效微调PEFT冻结主干训练轻量适配器这是大模型缓解灾难性遗忘的**最优实践**核心思路是**冻结预训练模型的主干参数**旧任务知识的载体仅训练插入到模型中的轻量级适配器模块新任务的知识完全存储在适配器中不会修改主干参数。典型的PEFT方法包括- **LoRALow-Rank Adaptation**在Transformer的注意力层中插入低秩矩阵新任务仅训练低秩矩阵的参数主干参数冻结。低秩矩阵的参数量仅为模型总参数量的0.1%~1%几乎不增加显存开销。- **Adapter**在Transformer的FFN层和注意力层之间插入小型瓶颈网络如Down-Projection Non-Linear Up-Projection仅训练Adapter的参数主干参数冻结。- **Prefix Tuning**在输入序列前添加可训练的前缀向量新任务仅优化前缀向量主干模型参数冻结。**优势**完全避免主干参数被修改旧任务知识零遗忘参数量小训练和部署成本低**适用场景**大模型多任务微调、增量任务学习是工业界的首选方案。#### 2.2 增量网络Incremental Network为每个新任务**新增独立的网络层或参数子集**不修改旧任务的参数。例如- 对于CNN学习新任务时新增卷积层和全连接层旧层参数冻结- 对于Transformer学习新任务时新增注意力头或FFN层旧层参数冻结。**优势**旧任务性能无损失**局限**任务过多时模型参数量会线性增长导致“模型膨胀”不利于部署。### 3. 基于数据回放的方法保留旧任务数据混合训练这类方法的核心是**存储旧任务的部分数据**在学习新任务时将旧任务数据与新任务数据混合训练让模型在学习新知识的同时“复习”旧知识从而保留旧任务的性能。#### 3.1 核心集回放Core Set Replay- **核心思路**不存储旧任务的全部数据避免存储成本过高而是通过聚类、代表性采样等方法选择旧任务的**核心样本**如覆盖数据分布的关键样本存储。学习新任务时用“新数据 旧核心样本”混合训练。- **典型算法**Herding、K-Center Greedy等用于高效选择代表性样本。**优势**效果稳定对新旧任务的性能都有保障**局限**需要存储旧任务数据存在隐私和存储成本问题例如医疗、金融数据不能随意存储。#### 3.2 生成式回放Generative Replay- **核心思路**不存储旧任务的真实数据而是训练一个**生成模型**如GAN、VAE、扩散模型学习旧任务的数据分布。学习新任务时用生成模型生成旧任务的“伪数据”与新数据混合训练。- **进阶方案**用旧任务训练好的模型作为**教师模型**新任务训练时让当前模型的输出与教师模型的输出对齐知识蒸馏同时拟合新任务数据。例如经典的**LwFLearning without Forgetting** 算法总损失为$$\mathcal{L}_{total} \mathcal{L}_{new}(\theta) \alpha \cdot KL(q_{old}(y|x) \parallel q_{new}(y|x))$$其中 $q_{old}$ 是旧任务模型的输出分布$q_{new}$ 是当前模型的输出分布KL散度用于约束当前模型保留旧任务的知识。**优势**无需存储旧任务真实数据解决隐私和存储问题**局限**生成模型的质量决定回放效果若生成的伪数据与真实数据差异大缓解遗忘的效果会大打折扣。### 4. 其他方法动态架构与知识蒸馏- **动态架构方法**例如**神经网络模块化**将模型拆分为多个独立的模块每个任务激活不同的模块组合避免参数冲突- **跨任务知识蒸馏**训练一个“通用模型”将旧任务的知识蒸馏到通用模型中再用通用模型初始化新任务的训练实现知识的迁移和保留。## 三、大模型场景下的最优缓解策略在大模型如7B、13B、70B参数的持续学习和多任务微调中**参数高效微调PEFT 少量数据回放**是工业界的主流方案原因如下1. **PEFT保证旧知识不遗忘**冻结主干参数仅训练适配器完全避免旧任务关键参数被覆盖2. **少量数据回放提升新任务性能**混合少量预训练语料或旧任务核心样本可防止模型“过拟合”新任务同时增强新旧任务的知识融合3. **成本可控**适配器的参数量极小回放数据量仅需总数据量的5%~10%训练和显存成本远低于全参数微调。例如用Llama-7B先微调“文本摘要”任务再微调“情感分类”任务时可采用**LoRA 摘要任务核心样本回放**- 冻结Llama-7B的主干参数仅训练注意力层的LoRA矩阵- 训练时混合80%的情感分类数据 20%的摘要任务核心样本- 最终模型既能保持摘要任务的性能又能达到情感分类的最优精度。## 四、总结缓解灾难性遗忘的方法需根据**任务场景、数据隐私、算力成本**选择- 数据隐私受限、无法存储旧数据 → 选择**正则化方法EWC/SI或参数隔离PEFT**- 数据隐私无限制、追求最优性能 → 选择**数据回放 参数隔离**的组合方案- 大模型多任务微调 → **PEFTLoRA/Adapter是首选**兼顾性能和效率。