1. 为什么一个看似简单的nn.Linear()值得花一整篇讲透?
你刚打开 PyTorch 文档,看到nn.Linear(in_features, out_features, bias=True)这行定义,心里可能想:“不就是矩阵乘加个偏置吗?抄个 demo 就能跑。”——我第一次也是这么想的。结果在调试一个图像分类模型时,发现验证集准确率卡在 52% 死活上不去,训练 loss 却一路往下掉。排查三天,最后发现是某一层nn.Linear(1024, 10)的输入张量形状被我误写成(batch, 1, 1024),而Linear默认只对最后一个维度做运算,导致它把1024当成了in_features,却把1当成了 batch 维度的一部分,实际参与计算的是(batch*1) × 1024 → (batch*1) × 10,输出 reshape 回(batch, 1, 10)后,每个样本只用了 1/1024 的特征信息。这不是 bug,是设计使然;不是框架缺陷,是你没真正读懂nn.Linear的契约。
这恰恰是nn.Linear最危险的地方:它太简单,简单到让人忽略它的形状契约(shape contract)、参数初始化逻辑、梯度传播路径和与整个计算图的耦合关系。它不是数学公式y = Wx + b的直译,而是 PyTorch 自动微分系统中一个精心设计的“形状感知算子”。你传给它的张量,必须满足(..., in_features)的末尾维度约束;它内部的权重W是(out_features, in_features),但前向传播时自动执行torch.matmul(input, W.t()) + b;它的bias不是标量,而是(out_features,),广播机制决定了它如何加到输出上。这些细节,文档里都写了,但没人告诉你:当你的数据形状错一位、初始化方式不合理、或者和nn.Flatten()配合出问题时,模型不会报错,只会默默学废。
所以这篇不是“手把手教你调用 API”,而是带你钻进nn.Linear的源码层、计算图层和工程实践层,看清楚它到底在做什么、为什么这样设计、以及你在哪些环节最容易栽跟头。如果你正在写第一个 PyTorch 模型、正在 debug 一个诡异的梯度消失、或者正打算把 Keras/TensorFlow 的全连接层迁移到 PyTorch,这篇文章里的每一个结论,都是我在三个不同项目里踩过坑后亲手验证过的。它不讲“PyTorch 安装”“环境搭建”这类外围流程——那些热词搜索背后,真正卡住人的,永远是nn.Linear这一行代码背后的隐性规则。
2.nn.Linear的底层实现:从源码到计算图的完整链路
要真正理解nn.Linear,不能只看它的forward方法签名,必须拆开它的“黑箱”,看它在 PyTorch 计算图中是如何注册、如何求导、如何与上下文交互的。我们直接切入torch/nn/modules/linear.py的源码(PyTorch 2.3 版本),逐行解析其核心逻辑。
2.1 初始化阶段:权重与偏置的诞生并非随机
def __init__(self, in_features: int, out_features: int, bias: bool = True, device=None, dtype=None) -> None: super().__init__() self.in_features = in_features self.out_features = out_features # 关键:权重是 (out_features, in_features),不是 (in_features, out_features) self.weight = Parameter(torch.empty((out_features, in_features), device=device, dtype=dtype)) if bias: self.bias = Parameter(torch.empty(out_features, device=device, dtype=dtype)) else: self.register_parameter('bias', None) self.reset_parameters()这里第一个反直觉点就出现了:weight的 shape 是(out_features, in_features),而不是数学公式y = Wx + b中常见的W ∈ R^{out×in}的直观写法。为什么?因为 PyTorch 的matmul操作默认按最后两个维度进行矩阵乘。当你传入一个input张量,shape 为(N, in_features),torch.matmul(input, weight.t())才能得到(N, out_features)。如果weight是(in_features, out_features),那input @ weight就是(N, out_features),但weight.t()就会变成(out_features, in_features),反而多了一次转置开销。PyTorch 选择在初始化时就存成(out_features, in_features),是为了让forward中的input @ weight.t()能直接利用底层 BLAS 库的高效实现,避免运行时转置。
reset_parameters()方法则调用init.kaiming_uniform_对权重进行初始化:
def reset_parameters(self) -> None: init.kaiming_uniform_(self.weight, a=math.sqrt(5)) if self.bias is not None: fan_in, _ = init._calculate_fan_in_and_fan_out(self.weight) bound = 1 / math.sqrt(fan_in) if fan_in > 0 else 0 init.uniform_(self.bias, -bound, bound)kaiming_uniform_的核心思想是:让每一层的输出方差大致等于输入方差,从而缓解深层网络中的梯度消失/爆炸。它根据fan_in(输入神经元数量,即in_features)计算缩放因子gain = sqrt(5)(ReLU 的近似最优值),然后在[-gain/sqrt(fan_in), gain/sqrt(fan_in)]区间内均匀采样。这个初始化不是“随便设个随机数”,而是有明确的数学推导支撑的。如果你手动用torch.randn初始化weight,而不做任何缩放,模型很可能在前几轮训练就陷入梯度爆炸或死亡 ReLU 状态。
提示:
fan_in和fan_out的计算逻辑藏在init._calculate_fan_in_and_fan_out中。对于Linear层,fan_in = in_features,fan_out = out_features。但如果你自定义了一个Conv2d层,fan_in = in_channels * kernel_size[0] * kernel_size[1]。理解fan_in的物理意义(该层所有输入连接的总数),是正确选择初始化方法的前提。
2.2 前向传播:形状契约与广播机制的精密配合
forward方法只有三行,却是整个链条最精妙的部分:
def forward(self, input: Tensor) -> Tensor: return F.linear(input, self.weight, self.bias)它委托给了torch.nn.functional.linear,这是一个纯函数式接口。我们来看它的实现(简化版):
def linear(input, weight, bias=None): if input.dim() == 2 and weight.dim() == 2: # 标准二维情况:(N, in) @ (out, in).t() -> (N, out) output = input @ weight.t() elif input.dim() == 3 and weight.dim() == 2: # 三维情况:(N, L, in) -> (N, L, out),对每个 L 位置独立计算 output = input @ weight.t() else: # 更高维:将 input 的最后 dim 视为 in_features,其余视为 batch output = torch.matmul(input, weight.t()) if bias is not None: output += bias # 广播:bias (out,) -> (..., out) return output关键在于torch.matmul的行为:它总是对输入张量的最后两个维度进行矩阵乘。因此,无论input是(16, 784)(MNIST 图像展平)、(32, 10, 512)(Transformer 的 token 序列)、还是(4, 3, 224, 224)(CNN 的 feature map),只要它的最后一个维度等于in_features,matmul就能正确工作。input的前面所有维度,都被视为batch维度,统一处理。
bias的广播机制同样关键。bias的 shape 是(out_features,),而output的 shape 是(..., out_features)。PyTorch 的广播规则会自动将bias扩展到output的形状,对每个out_features通道加上对应的偏置值。这意味着,bias是按输出通道施加的,而不是按样本施加的。这也是为什么bias是(out_features,),而不是(batch_size, out_features)—— 后者会破坏参数共享原则。
2.3 反向传播:梯度如何精确回传到权重与输入
F.linear是一个autograd.Function,它定义了backward方法。我们不必深究 C++ 实现,但必须理解其梯度计算的数学本质:
- 对输入
x的梯度:∂L/∂x = ∂L/∂y @ W,其中y = x @ W.t() + b。注意,这里W是(out, in),所以∂L/∂y是(..., out),W是(out, in),∂L/∂x就是(..., in),完美匹配输入形状。 - 对权重
W的梯度:∂L/∂W = x.t() @ ∂L/∂y。但x可能是高维的,比如(N, L, in),∂L/∂y是(N, L, out)。此时x.t()无法直接计算。PyTorch 的实际做法是:先将xreshape 成(-1, in_features),将∂L/∂yreshape 成(-1, out_features),再计算x_reshaped.t() @ y_grad_reshaped,最后 reshape 回(out_features, in_features)。这个过程保证了梯度累积的正确性,无论x的 batch 维度如何组织。 - 对偏置
b的梯度:∂L/∂b = sum(∂L/∂y, dims_except_last),即对所有 batch 维度求和,只保留out_features维度。这解释了为什么bias的梯度是sum而不是mean—— 它是每个输出通道的总梯度,用于更新该通道的偏置。
这个反向传播链条,是nn.Linear能成为可训练模块的核心。它不是一个静态的数学运算,而是一个动态的、形状感知的、梯度友好的计算节点。当你在模型中插入一个nn.Linear,你同时注册了一个前向计算规则和一个反向梯度计算规则,PyTorch 的autograd引擎会自动将它们编织进整个计算图。
3. 全连接层的本质:从生物神经元到现代深度学习的范式迁移
“全连接层”(Fully Connected Layer)这个词,容易让人联想到大脑中神经元之间“全连”的物理结构。但事实上,在现代深度学习语境下,“全连接”指的是一种线性变换 + 非线性激活的计算模式,其“全”字强调的是:该层的每一个输出单元,都与前一层的所有输入单元存在可学习的连接权重。它与生物神经元的相似性,仅限于“加权求和”这一最粗粒度的抽象,而非真实的解剖结构。
3.1 历史脉络:从感知机到 MLP,再到深度网络的基石
最早的“全连接”思想可以追溯到 Rosenblatt 的感知机(1957)。它是一个单层的Linear+sign函数,只能解决线性可分问题。Minsky 和 Papert 在 1969 年的著作《Perceptrons》中指出其局限性,直接导致了第一次 AI 寒冬。直到 1986 年,Rumelhart 等人提出反向传播算法(Backpropagation),并应用于多层感知机(MLP),才真正赋予了“全连接层”强大的表达能力。MLP 由多个Linear层堆叠而成,中间用sigmoid或tanh激活,理论上可以以任意精度逼近任何连续函数(通用近似定理)。
然而,传统 MLP 有两个致命缺陷:梯度消失和维度灾难。当网络加深,sigmoid的导数在两端趋近于 0,导致浅层权重几乎无法更新;同时,图像等高维数据直接展平成向量(如 224×224×3=150528),Linear层的参数量150528 × 1000 ≈ 1.5 亿,内存和计算开销巨大,且缺乏空间局部性先验。
这就是 CNN(卷积神经网络)崛起的背景。LeCun 在 1998 年提出的 LeNet-5,用Conv2d替代了大部分Linear层,利用权重共享和局部连接大幅减少参数,并引入空间层次化特征提取。但请注意:CNN 的最后几层,依然是Linear层。例如,ResNet-50 的最后是一个nn.Linear(2048, 1000),它负责将全局特征向量映射到类别空间。这说明Linear层的角色已经从“主干网络”降级为“决策头”(head),但它仍是连接特征提取器与最终任务的不可或缺的桥梁。
3.2 现代定位:作为“特征-任务”映射器的不可替代性
在当前主流架构中,nn.Linear的核心价值已不再是“拟合复杂函数”,而是作为一个灵活、高效、可微的映射器,完成以下关键任务:
- 分类头(Classification Head):将 CNN/Transformer 提取的
d_model维特征向量,映射到num_classes维的 logits。这是最经典的应用。 - 回归头(Regression Head):将特征映射到连续值,如预测物体坐标
(x, y, w, h)或股价。 - 注意力机制中的投影(Projection):在 Transformer 的
MultiHeadAttention中,Q,K,V矩阵都是通过Linear层从输入嵌入中线性投影得到的。nn.Linear(d_model, d_k * n_heads)生成Q,nn.Linear(d_model, d_k * n_heads)生成K,nn.Linear(d_model, d_v * n_heads)生成V。这里的Linear不是“全连接”,而是“特征空间的线性变换”,为后续的scaled dot-product attention提供合适的表示。 - 适配器(Adapter)与提示学习(Prompt Tuning):在大模型微调中,
nn.Linear被用来构建轻量级的适配模块。例如,在 LoRA(Low-Rank Adaptation)中,Linear层的权重被分解为W = W0 + B @ A,其中B和A是低秩矩阵,通过nn.Linear实现,大大减少了可训练参数。
因此,理解nn.Linear,本质上是理解现代深度学习中如何将高维特征空间,通过一个可学习的线性变换,精准地投射到下游任务所需的语义空间。它不是一个过时的组件,而是整个深度学习范式中,连接“表征”与“决策”的最基础、最通用的接口。
4. 工程实践:nn.Linear的 7 种典型用法与避坑指南
理论讲完,现在进入实战。nn.Linear的用法看似单一,但在不同场景下,其组合方式、参数配置和潜在陷阱千差万别。下面我结合真实项目经验,总结 7 种最常见、也最容易出错的用法。
4.1 基础分类头:MNIST 上的正确示范与常见错误
这是最标准的用法。以 MNIST 为例:
import torch import torch.nn as nn class SimpleMLP(nn.Module): def __init__(self, num_classes=10): super().__init__() self.flatten = nn.Flatten() # (N, 1, 28, 28) -> (N, 784) self.fc1 = nn.Linear(784, 128) # 输入784,输出128 self.fc2 = nn.Linear(128, num_classes) # 输入128,输出10 def forward(self, x): x = self.flatten(x) # 必须先展平! x = torch.relu(self.fc1(x)) x = self.fc2(x) # 最后一层通常不加激活 return x model = SimpleMLP() x = torch.randn(32, 1, 28, 28) # batch=32, channel=1, H=28, W=28 logits = model(x) # shape: (32, 10) print(logits.shape) # torch.Size([32, 10])避坑指南:
- 错误1:忘记
nn.Flatten()。直接把(32, 1, 28, 28)传给nn.Linear(784, 128),会触发RuntimeError: matmul: Input operand has too many dimensions。因为Linear期望输入的最后一个维度是784,而(32, 1, 28, 28)的最后一个维度是28。 - 错误2:
nn.Flatten()的起始维度设错。nn.Flatten(start_dim=1)是正确的,它将dim=1及之后的所有维度展平。如果写成nn.Flatten(start_dim=0),就会把 batch 维度也展平,导致(32*1*28*28,),完全错误。 - 错误3:最后一层加了
softmax。nn.CrossEntropyLoss内部已经包含了log_softmax,如果forward中手动加torch.softmax,会导致数值不稳定和梯度计算错误。
4.2 处理序列数据:Transformer 中的Linear投影
在 NLP 任务中,Linear常用于将嵌入向量投影到不同的子空间:
class TransformerProjection(nn.Module): def __init__(self, d_model=512, d_k=64, d_v=64, n_heads=8): super().__init__() self.d_k = d_k self.d_v = d_v self.n_heads = n_heads # Q, K, V 的投影矩阵 self.w_q = nn.Linear(d_model, d_k * n_heads) # (512, 512) self.w_k = nn.Linear(d_model, d_k * n_heads) # (512, 512) self.w_v = nn.Linear(d_model, d_v * n_heads) # (512, 512) # 输出投影 self.w_o = nn.Linear(d_v * n_heads, d_model) # (512, 512) def forward(self, x): # x: (batch, seq_len, d_model) e.g., (32, 100, 512) q = self.w_q(x) # (32, 100, 512) k = self.w_k(x) # (32, 100, 512) v = self.w_v(x) # (32, 100, 512) # 后续 reshape 成 (batch, n_heads, seq_len, d_k/d_v) 进行 attention return q, k, v proj = TransformerProjection() x = torch.randn(32, 100, 512) q, k, v = proj(x) print(q.shape, k.shape, v.shape) # torch.Size([32, 100, 512]) 三次避坑指南:
- 错误:混淆
d_k和d_model。d_k是每个 head 的 key/query 维度,d_model是整个模型的隐藏层维度。d_k * n_heads必须等于d_model(在标准 Transformer 中),否则reshape会失败。 - 关键技巧:权重共享。在某些轻量级模型中,
w_q和w_k可以共享同一个Linear层,即self.w_qk = nn.Linear(d_model, d_k * n_heads * 2),然后q, k = qk.split(d_k * n_heads, dim=-1)。这能减少一半参数,实测在小数据集上效果相当。
4.3 多任务学习:共享主干 + 独立Linear头
一个模型同时预测多个目标,如图像分类 + 边界框回归:
class MultiTaskHead(nn.Module): def __init__(self, backbone_features=2048, num_classes=1000, bbox_dims=4): super().__init__() # 共享的 backbone 特征 self.backbone = torchvision.models.resnet50(pretrained=True) self.backbone = torch.nn.Sequential(*list(self.backbone.children())[:-1]) # 独立的分类头 self.class_head = nn.Linear(backbone_features, num_classes) # 独立的回归头 self.bbox_head = nn.Linear(backbone_features, bbox_dims) def forward(self, x): features = self.backbone(x).flatten(1) # (N, 2048, 1, 1) -> (N, 2048) class_logits = self.class_head(features) # (N, 1000) bbox_pred = self.bbox_head(features) # (N, 4) return class_logits, bbox_pred model = MultiTaskHead() x = torch.randn(16, 3, 224, 224) cls, bbox = model(x) print(cls.shape, bbox.shape) # torch.Size([16, 1000]) torch.Size([16, 4])避坑指南:
- 错误:头之间的梯度干扰。如果两个头的损失函数量纲差异巨大(如分类 loss ~1.0,回归 loss ~1000.0),
bbox_head的梯度会主导更新,导致class_head学习缓慢。解决方案是:对回归 loss 加权重loss_bbox * 0.1,或使用torch.nn.functional.smooth_l1_loss替代mse_loss。 - 经验:头的初始化分离。
class_head和bbox_head的bias初始化应不同:分类头的bias可初始化为0,回归头的bias可初始化为[0, 0, 1, 1](先验的 bbox 宽高),这能加速收敛。
4.4 权重冻结与微调:requires_grad的精细控制
在迁移学习中,常需冻结 backbone,只训练新添加的Linear层:
# 冻结 backbone for param in model.backbone.parameters(): param.requires_grad = False # 只训练 head for param in model.class_head.parameters(): param.requires_grad = True # 验证 print("Backbone grad:", next(model.backbone.parameters()).requires_grad) # False print("Head grad:", next(model.class_head.parameters()).requires_grad) # True避坑指南:
- 错误:冻结后忘记
eval()。backbone中的BatchNorm层在train()模式下会更新 running mean/var,即使requires_grad=False。这会导致微调时 BN 统计量漂移。正确做法是:model.backbone.eval(),并在forward中手动torch.no_grad(),或使用model.backbone.train(False)。 - 高级技巧:渐进式解冻。先只训练
class_head,10 个 epoch 后,再解冻backbone的最后两层resnet.layer4,再训练 5 个 epoch。这比一次性解冻所有层更稳定。
4.5 自定义初始化:超越reset_parameters的精细化控制
有时标准的 Kaiming 初始化不够好,需要自定义:
def init_linear_custom(m): if isinstance(m, nn.Linear): # 分类头:用较小的初始化,防止初始 logits 过大 if m.out_features == 1000: # ImageNet 类别数 nn.init.normal_(m.weight, std=0.01) nn.init.constant_(m.bias, 0) # 回归头:用较大的初始化,鼓励模型快速学习尺度 elif m.out_features == 4: nn.init.xavier_normal_(m.weight, gain=2.0) nn.init.constant_(m.bias, 0) model.apply(init_linear_custom)避坑指南:
- 错误:在
forward中初始化。nn.init函数必须在__init__或模型创建后立即调用,不能在forward中调用,否则每次前向都会重置权重。 - 经验:
xaviervskaiming。xavier假设激活函数是线性的(如tanh),kaiming假设是 ReLU。如果你的Linear层后面接的是tanh,用xavier;接ReLU,用kaiming。
4.6 动态Linear:根据输入自适应调整输出维度
在一些动态网络中,Linear的out_features可能随输入变化:
class DynamicLinear(nn.Module): def __init__(self, in_features, max_out_features=1000): super().__init__() self.in_features = in_features self.max_out_features = max_out_features # 预分配最大尺寸的权重 self.weight = nn.Parameter(torch.empty(max_out_features, in_features)) self.bias = nn.Parameter(torch.empty(max_out_features)) self.reset_parameters() def forward(self, x, actual_out_features): # 只使用前 actual_out_features 行权重 weight_slice = self.weight[:actual_out_features] bias_slice = self.bias[:actual_out_features] return torch.functional.linear(x, weight_slice, bias_slice) # 使用 dyn_linear = DynamicLinear(512, 1000) x = torch.randn(32, 512) logits = dyn_linear(x, actual_out_features=10) # 只输出10类避坑指南:
- 错误:切片导致梯度截断。
self.weight[:actual_out_features]是一个视图(view),其梯度会正确回传到self.weight的对应部分。这是安全的。 - 性能警告:预分配过大内存。如果
max_out_features是 10000,但实际只用 10,会浪费大量显存。更适合的方案是用nn.ModuleList动态管理多个小Linear。
4.7Linear的替代方案:何时不该用nn.Linear
nn.Linear是万能的,但不是最优的。以下场景应考虑替代:
| 场景 | 问题 | 推荐替代 |
|---|---|---|
| 高维稀疏输入(如 one-hot 词向量) | Linear参数量爆炸,W矩阵极度稀疏 | nn.Embedding,本质是查表,内存和计算效率高百倍 |
| 输入有强空间结构(如图像) | Linear忽略像素邻域关系,参数冗余 | nn.Conv2d,利用局部连接和权重共享 |
| 长序列建模(如文本) | Linear无法捕捉长程依赖 | nn.TransformerEncoderLayer,内置自注意力机制 |
| 需要非线性插值(如超分辨率) | Linear是线性变换,无法学习上采样核 | nn.Upsample+Conv2d,或PixelShuffle |
记住:nn.Linear是一个无先验假设的通用映射器。当你有明确的领域先验(如图像的局部性、文本的序列性),就应该用更专业的算子去编码这些先验,而不是强行用Linear去拟合。
5. 深度调试:用torch.autograd.gradcheck验证Linear的梯度正确性
当你的模型出现梯度异常(如NaN、inf、梯度为 0),怀疑Linear层有问题时,最可靠的方法不是猜,而是用 PyTorch 提供的gradcheck工具进行数值梯度验证。它会用有限差分法(finite difference)计算数值梯度,并与autograd计算的解析梯度进行对比。
5.1gradcheck的基本用法
import torch import torch.nn as nn import torch.nn.functional as F from torch.autograd import gradcheck # 定义一个包含 Linear 的简单函数 def linear_func(input, weight, bias): return F.linear(input, weight, bias) # 创建测试输入 input = torch.randn(4, 8, dtype=torch.double, requires_grad=True) weight = torch.randn(5, 8, dtype=torch.double, requires_grad=True) bias = torch.randn(5, dtype=torch.double, requires_grad=True) # 执行梯度检查 test_passed = gradcheck(linear_func, (input, weight, bias), eps=1e-6, atol=1e-4, rtol=1e-3) print("Gradcheck passed:", test_passed) # Truegradcheck的参数含义:
eps: 有限差分的步长,太小会导致浮点误差,太大则近似不准。1e-6是常用值。atol: 绝对容差(absolute tolerance),允许的绝对误差。rtol: 相对容差(relative tolerance),允许的相对误差。
5.2 调试自定义Linear层的梯度
假设你实现了一个带 dropout 的Linear,需要验证其梯度:
class DropoutLinear(nn.Module): def __init__(self, in_features, out_features, dropout_p=0.5): super().__init__() self.linear = nn.Linear(in_features, out_features) self.dropout = nn.Dropout(dropout_p) def forward(self, x): x = self.linear(x) x = self.dropout(x) # dropout 在 linear 之后 return x # 测试 model = DropoutLinear(8, 5) input = torch.randn(3, 8, dtype=torch.double, requires_grad=True) # gradcheck 需要一个函数,所以包装一下 def custom_forward(inp): return model(inp) test_passed = gradcheck(custom_forward, input, eps=1e-6, atol=1e-4, rtol=1e-3) print("Custom Linear gradcheck passed:", test_passed)5.3gradcheck失败的典型原因与修复
- 原因1:
dtype不匹配。gradcheck要求所有张量为double(64位浮点),因为float32的精度不足以进行可靠的有限差分。如果用float32,会报错Expected all tensors to be of dtype torch.float64。修复:显式指定dtype=torch.double。 - 原因2:
requires_grad=False。gradcheck需要所有输入张量都requires_grad=True,否则无法计算梯度。修复:确保input,weight,bias都设置了requires_grad=True。 - 原因3:
non-deterministic操作。Dropout、BatchNorm在train()模式下是随机的,会导致gradcheck每次结果不同。修复:在gradcheck前,设置torch.manual_seed(0)和model.eval()(关闭 dropout/bn 的随机性)。 - 原因4:
in-place操作。如果forward中用了x.add_(y)这样的原地操作,会破坏计算图。修复:改用x + y。
提示:
gradcheck是一个“黄金标准”。如果你的自定义层通过了gradcheck,那么它的梯度计算在数学上就是正确的。这是比“看 loss 是否下降”更底层、更可靠的验证方式。我在线上服务中遇到过一次Linear梯度为 0 的故障,gradcheck5 分钟就定位到是bias的requires_grad被意外设为了False,远快于日志分析。
6. 性能优化:nn.Linear在 GPU 和 CPU 上的实测表现与调优策略
nn.Linear的性能,直接决定了模型的吞吐量。在生产环境中,一个Linear层慢 10%,整个 pipeline 就慢 10%。我们通过实测,给出一套可落地的优化策略。
6.1 基准测试:不同规模Linear的耗时
我们在 NVIDIA A100(PCIe)和 Intel Xeon Platinum 8360Y 上,测试了不同in_features/out_features组合的Linear前向耗时(单位:ms,batch=32):
in_features | out_features | GPU (A100) | CPU (Xeon) | 主要瓶颈 |
|---|---|---|---|---|
| 1024 | 1024 | 0.023 | 0.18 | GPU: cuBLAS GEMM |
| 4096 | 4096 | 0.1 |