PyTorch逻辑回归与交叉熵损失函数实战指南
1. 逻辑回归与交叉熵损失函数基础
在深度学习中,逻辑回归是最基础的分类模型之一,而交叉熵损失函数则是分类任务中最常用的优化目标。理解这两者的配合使用对于掌握PyTorch中的分类模型构建至关重要。
1.1 逻辑回归的本质
逻辑回归虽然名字中有"回归",但实际上是一个分类算法。它的核心思想是通过线性变换加上sigmoid激活函数,将输入特征映射到0-1之间的概率值。数学表达式为:
P(y=1|x) = σ(w·x + b) = 1 / (1 + e^-(w·x + b))
其中σ表示sigmoid函数,w是权重,b是偏置项。这个概率值可以直观理解为样本属于正类的可能性。
在PyTorch中,我们通常使用 torch.nn.Linear 实现线性变换部分,然后通过 torch.sigmoid 函数得到概率输出。这种组合既简单又高效,特别适合二分类问题。
1.2 为什么选择交叉熵损失
交叉熵损失(Cross-Entropy Loss)衡量的是模型预测概率分布与真实概率分布之间的差异。对于二分类问题,我们使用二元交叉熵(BCELoss),其数学定义为:
Loss = -[y·log(p) + (1-y)·log(1-p)]
与均方误差(MSE)相比,交叉熵有几个显著优势:
- 当预测值与真实值差距较大时,交叉熵能提供更大的梯度,加快模型收敛
- 它直接优化分类准确率,而不是连续值的误差
- 对分类边界附近的样本更加敏感
在实际应用中,我们观察到使用MSE损失时模型准确率可能只有57%,而切换到交叉熵后可以达到接近100%的准确率。这种差异在小批量训练时尤为明显。
2. PyTorch实现细节解析
2.1 数据准备与模型构建
在PyTorch中,我们首先需要自定义数据集类。一个好的实践是继承 torch.utils.data.Dataset 并实现三个关键方法:
class Data(Dataset):
def __init__(self):
self.x = torch.arange(-2, 2, 0.1).view(-1, 1)
self.y = torch.zeros(self.x.shape[0], 1)
self.y[self.x[:, 0] > 0.2] = 1
self.len = self.x.shape[0]
def __getitem__(self, idx):
return self.x[idx], self.y[idx]
def __len__(self):
return self.len
这个数据集创建了一个简单的线性可分数据,其中x>0.2的样本被标记为1,其余为0。这种设计让我们可以清晰地观察模型的学习效果。
模型构建方面,我们创建一个继承自 nn.Module 的逻辑回归类:
class LogisticRegression(nn.Module):
def __init__(self, n_inputs):
super().__init__()
self.linear = nn.Linear(n_inputs, 1)
def forward(self, x):
y_pred = torch.sigmoid(self.linear(x))
return y_pred
关键点在于forward方法中,我们先进行线性变换,然后通过sigmoid函数将输出压缩到0-1之间。这种结构虽然简单,但包含了逻辑回归的所有核心要素。
2.2 训练过程优化技巧
训练逻辑回归模型时,有几个关键参数需要特别注意:
-
学习率(lr):对于这个简单模型,我们可以使用较大的学习率(如2.0),因为数据是线性可分的。但在更复杂场景中,通常需要更小的值。
-
批量大小(batch_size):这里设置为2,主要是为了演示小批量训练的效果。实际应用中,根据数据规模和内存情况,可能需要调整。
-
优化器选择:随机梯度下降(SGD)对于这种简单模型已经足够。对于更复杂网络,可以考虑Adam等自适应优化器。
训练循环的标准结构如下:
optimizer = torch.optim.SGD(model.parameters(), lr=2)
criterion = nn.BCELoss()
for epoch in range(epochs):
for x, y in train_loader:
optimizer.zero_grad()
y_pred = model(x)
loss = criterion(y_pred, y)
loss.backward()
optimizer.step()
特别要注意的是 zero_grad() 的调用位置,它必须在每次反向传播前清空梯度,否则梯度会累积导致训练不稳定。
3. 模型评估与结果分析
3.1 训练过程监控
在训练过程中,我们记录了每个batch的损失值,这可以帮助我们观察模型的学习动态。理想情况下,损失曲线应该呈现稳定的下降趋势:
plt.plot(Loss)
plt.xlabel("Iterations")
plt.ylabel("Loss")
plt.show()
如果出现损失震荡或无法下降的情况,可能需要调整学习率或检查数据质量。在这个例子中,我们可以看到损失快速收敛到接近0的值,表明模型学习效果良好。
3.2 测试集性能评估
模型评估阶段,我们不仅关心损失值,更关注分类准确率。计算准确率的关键步骤是:
y_pred = model(data_set.x)
predictions = y_pred > 0.5 # 以0.5为阈值
accuracy = (predictions == data_set.y.byte()).float().mean()
这里0.5是一个常用的分类阈值,但对于类别不平衡的数据集,可能需要调整这个值。在我们的简单例子中,模型可以达到100%的准确率,这得益于:
- 数据的线性可分性
- 合适的损失函数选择
- 充分的训练迭代
3.3 与MSE损失的对比实验
为了展示交叉熵的优势,我们可以进行对比实验。当使用MSE损失时,同样的模型架构和训练参数,准确率可能只有57%左右。这是因为:
- MSE对概率输出的惩罚不够直接
- 在分类边界附近梯度信号较弱
- 容易陷入局部最优
这个对比清晰地展示了损失函数选择对模型性能的重大影响。
4. 实战技巧与常见问题
4.1 权重初始化策略
虽然逻辑回归对初始化不太敏感,但好的初始化可以加速收敛。PyTorch的Linear层默认使用均匀初始化,对于这个简单任务已经足够。对于更复杂网络,可以考虑:
torch.nn.init.xavier_uniform_(self.linear.weight)
4.2 学习率调整策略
固定学习率虽然简单,但动态调整往往能获得更好效果。PyTorch提供了多种学习率调度器:
scheduler = torch.optim.lr_scheduler.StepLR(optimizer, step_size=10, gamma=0.1)
在每个epoch后调用 scheduler.step() 即可实现学习率衰减。
4.3 常见问题排查
-
损失不下降 :
- 检查学习率是否合适
- 确认梯度是否正常传播(打印梯度值)
- 验证数据是否正常加载
-
准确率波动大 :
- 增大批量大小
- 添加正则化项
- 检查数据是否有噪声
-
过拟合 :
- 增加L2正则化
- 使用更简单的模型
- 获取更多训练数据
4.4 扩展应用
虽然我们演示的是简单的一维逻辑回归,但同样的原理可以扩展到:
- 多维特征的分类问题
- 多分类任务(使用Softmax和NLLLoss)
- 更复杂的深度网络作为分类器
在实际项目中,逻辑回归常常作为基线模型,或者作为更复杂网络的最后一层。掌握这些基础知识对理解深度学习的分类任务至关重要。
更多推荐


所有评论(0)