You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

MNIST目标定位TensorFlow转PyTorch训练损失不下降问题排查

问题解答

1. 网络结构问题

卷积层、池化层、全连接层的通道数、核大小、连接逻辑和原TensorFlow代码一致,展平层的输入维度3136计算正确(输入7575单通道图,三次卷积+池化后特征图尺寸为77,7764=3136),但存在两处致命错误:

  • nn.Softmax(dim=0)维度设置错误:分类输出形状为[batch_size, 10],dim=0是批次维度,相当于在同一个批次的不同样本之间做softmax计算,而非在单个样本的10个类别维度上计算,这直接导致你看到的所有类别预测值都是1/批次大小(你批次大小设为64,1/64≈0.0156,和你贴的输出完全吻合)。
  • 分类头额外添加Softmax激活和后续损失函数不匹配:PyTorch的nn.CrossEntropyLoss内部已经集成了LogSoftmax和NLLLoss计算,要求输入是未经过激活的原始logits,你在网络输出端提前加Softmax会导致梯度计算异常,出现梯度消失问题。

2. 损失函数配置问题

损失函数的选型思路正确(分类用交叉熵、回归用MSE),但存在两处配置错误:

  • 输入给nn.CrossEntropyLoss的张量错误:你传入的是经过Softmax激活的概率值,而非原始logits,导致交叉熵计算逻辑完全失效。
  • 标签格式不匹配:原TensorFlow代码用的categorical_crossentropy接收one-hot格式标签,而nn.CrossEntropyLoss默认接收0-9的类别索引标签(形状为[batch_size]),如果你的数据集标签是one-hot格式,直接传入会得到错误的损失值。
  • 学习率设置严重错误:原TensorFlow代码用的Adam优化器默认学习率为0.001,你设置的lr=0.1是默认值的100倍,会导致参数更新步长过大,无法收敛。

3. 预测结果始终不变的原因

是上述几个错误叠加导致的:

  • Softmax维度设为dim=0时,模型初始化阶段全连接层权重接近0,所有样本的分类logit值差异极小,在批次维度做softmax后每个类别的输出值恒等于1/批次大小,不会随训练变化。
  • 提前加Softmax导致传入交叉熵损失的是概率值,损失计算出的梯度极小,加上学习率虽然数值大但有效梯度几乎为0,参数长期无法更新,输出自然固定不变。

4. 自定义训练循环和TF的compile/fit等价性

你写的训练循环(梯度清空、前向传播、损失计算、反向传播、参数更新)的核心逻辑和TF的compile/fit是等价的,但因为上述的网络结构、损失输入、学习率、标签匹配的错误,导致实际训练效果完全不一致。修正所有错误后可以达到和原TF代码几乎一致的收敛效果。


修正后的关键代码
class ConvNetwork(nn.Module):
    def __init__(self):
        super(ConvNetwork, self).__init__()
        self.conv2d_1 = nn.Conv2d(in_channels=1, out_channels=16, kernel_size=3)
        self.conv2d_2 = nn.Conv2d(in_channels=16, out_channels=32, kernel_size=3)
        self.conv2d_3 = nn.Conv2d(in_channels=32, out_channels=64, kernel_size=3)
        self.avgPooling2D = nn.AvgPool2d((2,2))
        self.dense_1 = nn.Linear(in_features=3136, out_features=128)
        
        self.dense_classifier = nn.Linear(in_features=128, out_features=10)
        # 删除维度设置错误的Softmax层,预测阶段再在类别维度计算概率
        self.dense_regression = nn.Linear(in_features=128, out_features=4)


    def forward(self, input):
        x = self.avgPooling2D(F.relu(self.conv2d_1(input)))
        x = self.avgPooling2D(F.relu(self.conv2d_2(x)))
        x = self.avgPooling2D(F.relu(self.conv2d_3(x)))
        x = nn.Flatten()(x)
        x = F.relu(self.dense_1(x))

        # 分类头输出原始logits给损失计算,预测概率单独运算
        logits_classifier = self.dense_classifier(x)
        output_regression = self.dense_regression(x)
        return [logits_classifier, output_regression]

######################################################

learning_rate = 0.001 # 修正为Adam默认学习率,和TensorFlow实现对齐
EPOCHS = 10 # 原TF代码默认训练10轮,1轮训练不足以观察到收敛
BATCH_SIZE = 64

model = ConvNetwork()
model = model.to(device)
optimizer = torch.optim.Adam(params=model.parameters(), lr=learning_rate)
classification_loss = nn.CrossEntropyLoss()
regression_loss = nn.MSELoss()

######################################################

begin_time = time.time()
for epoch in range(EPOCHS) : 
    tot_loss = 0
    train_start = time.time()
    training_losses = []
    
    print("-"*20)
    print(" "*5 + f"EPOCH {epoch+1}/{EPOCHS}")
    print("-"*20)

    model.train()
    for batch, (digits, labels, bbox_coords) in enumerate(training_dataset):
        digits, labels, bbox_coords = digits.to(device), labels.to(device), bbox_coords.to(device)
        # 如果标签是one-hot格式,取消下一行注释转换为类别索引
        # labels = labels.argmax(dim=1)
        optimizer.zero_grad()
        
        [label_logits, bbox_coords_preds] = model(digits)
        # 单独计算分类概率用于预测,不参与损失运算
        label_preds = F.softmax(label_logits, dim=1)
        
        class_loss = classification_loss(label_logits, labels)
        box_loss = regression_loss(bbox_coords_preds, bbox_coords)

        training_loss = class_loss + box_loss
        training_loss.backward()
        
        optimizer.step()
        
        ######### 打印逻辑 #######################
        training_losses.append(training_loss.item())
        if batch+1 <= len_training_ds//BATCH_SIZE:
            current_training_sample = (batch+1)*BATCH_SIZE
        else:
            current_training_sample = (batch)*BATCH_SIZE + len_training_ds%BATCH_SIZE
        
        if (batch+1) == 1 or (batch+1)%100 == 0 or (batch+1) == len_training_ds//BATCH_SIZE +1:
            print(f"Elapsed time : {(time.time()-train_start)/60:.3f}",\
                  f" --- Digit : {current_training_sample}/{len_training_ds}",\
                  f" : loss = {training_loss:.5f}")
            if batch+1 == (len_training_ds//BATCH_SIZE)+1:
                print(f"Total elapsed time for training : {(time.time()-begin_time)/60:.3f}")

补充注意:确保输入图像像素值归一化到0-1区间,和原TF代码的预处理逻辑一致,否则也会拖慢收敛速度。


内容的提问来源于stack exchange,提问作者ThOpaque

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.09.03 08:09:23