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

基于PyTorch的LSTM自编码器分类任务联合训练问题排查

问题描述

编辑:尴尬的是,我的错误是只打乱了数据却没有打乱标签。

我用PyTorch实现了LSTM自编码器来重构MNIST图像,现在要修改网络,让它同时完成图像重构和分类任务,核心要求是基于两种损失联合训练。

我的自编码器实现如下:

def __init__(self, input_size, hidden_size, num_layers, output_size, epochs, optimizer, learning_rate, grad_clip, batch_size):
    super(AE, self).__init__()
    self.encoder = Encoder(input_size, hidden_size, num_layers)
    self.decoder = Decoder(input_size, hidden_size, num_layers, output_size)
    self.epochs = epochs
    self.optimizer = optimizer
    self.learning_rate = learning_rate
    self.grad_clip = grad_clip
    self.batch_size = batch_size
    self.criterion = nn.MSELoss()
    self.losses = []

该自编码器的forward和train方法运行正常,在MNIST数据集上重构图像效果良好,MSE损失平均约为1e-6。

我通过继承自编码器类添加了分类模块:

class AeWithClassifier(AE):
    def __init__(self, input_size, hidden_size, num_layers, output_size, epochs, optimizer, learning_rate, grad_clip, batch_size, num_classes):
        super(AeWithClassifier, self).__init__(input_size, hidden_size, num_layers, output_size, epochs, optimizer, learning_rate, grad_clip, batch_size)

        self.classifier = nn.Sequential(
            nn.Linear(output_size*output_size, num_classes))
        self.classifier_criterion = nn.CrossEntropyLoss()

相关方法如下:

def forward(self, x):
    predictions = super().forward(x)
    classifier_predictions = self.classifier(predictions.reshape(-1, 28*28))
    return predictions, classifier_predictions
def train(self, x, y):
        losses = []
        optimizer = self.optimizer(self.parameters(), lr=self.learning_rate)

        for epoch in range(self.epochs):
            cur_loss = 0
            batch_idx = 0
            for batch_idx, x_batch in enumerate(x):
                x_batch = x_batch.to(device)
                y_batch = y[batch_idx*self.batch_size:(batch_idx+1)*self.batch_size]
                optimizer.zero_grad()
                predictions, classifier_predictions = self.forward(x_batch)

                recon_loss = self.criterion(predictions, x_batch)
                class_loss = self.classifier_criterion(classifier_predictions, y_batch)
                cur_loss = loss = recon_loss + class_loss

                loss.backward()
                nn.utils.clip_grad_norm_(self.parameters(), self.grad_clip)
                optimizer.step()
            losses.append(cur_loss.item())
            print(f'Epoch: {epoch+1}/{self.epochs}, Loss: {cur_loss.item()}')
        self.losses = losses

我将重构损失与分类损失相加作为总损失进行梯度计算和优化,但网络虽能正常重构图像,分类效果极差,交叉熵损失无法降至2.3以下。调整损失权重也无改善,请问是网络结构还是训练过程存在问题?


解决思路

1. 数据匹配问题(已发现标签未打乱)

只打乱数据不打乱标签会导致训练时图像和标签完全不匹配,分类器根本学不到有效特征。正确做法是用TensorDataset和DataLoader打包数据与标签,自动处理打乱和分批:

from torch.utils.data import TensorDataset, DataLoader

dataset = TensorDataset(x, y)
dataloader = DataLoader(dataset, batch_size=self.batch_size, shuffle=True)

训练时直接遍历dataloader即可拿到对应x_batch和y_batch,避免手动切片的匹配错误。

2. 分类器输入选择不合理

当前用重构后的图像作为分类器输入完全没必要——重构图像是编码器+解码器的输出,属于原图像的近似,包含冗余信息。分类应该直接使用编码器的瓶颈特征(即自编码器的压缩表征),这才是图像的核心语义信息,能大幅提升分类效率。

需要先修改原AE类的forward方法,让它同时返回编码器的特征和重构预测,再调整子类的forward:

# 原AE类的forward示例修改
def forward(self, x):
    encoder_feature = self.encoder(x)
    predictions = self.decoder(encoder_feature)
    return encoder_feature, predictions

# AeWithClassifier的forward修改
def forward(self, x):
    encoder_feature, predictions = super().forward(x)
    classifier_predictions = self.classifier(encoder_feature.reshape(-1, hidden_size))
    return predictions, classifier_predictions

3. 损失尺度不匹配

MSE损失(约1e-6)和交叉熵损失(2.3左右)尺度差6个数量级,直接相加会导致优化器几乎只关注重构损失,分类损失的梯度被淹没。需统一两者尺度:

  • 给重构损失乘以放大系数,比如recon_loss = self.criterion(predictions, x_batch) * 1e6,让数值范围与交叉熵接近;
  • 或使用动态权重,根据当前损失值自动调整比例,确保两者对总损失的贡献均衡。

4. 分类器结构过于简单

仅用一层线性层拟合MNIST分类在联合训练场景下能力不足,建议添加非线性层提升拟合能力:

self.classifier = nn.Sequential(
    nn.Linear(hidden_size, 128),
    nn.ReLU(),
    nn.Linear(128, num_classes)
)

5. 训练过程细节修正

  • train方法中cur_loss每次循环被覆盖,最后仅保存最后一个batch的损失,应累加所有batch损失后取平均:
cur_loss = 0
total_batches = len(dataloader)
for batch_idx, (x_batch, y_batch) in enumerate(dataloader):
    # ... 训练步骤 ...
    cur_loss += loss.item()
cur_loss /= total_batches
losses.append(cur_loss)
  • 确保y_batch移到对应设备上,添加y_batch = y_batch.to(device),避免CPU与GPU数据不匹配的错误。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.29 06:27:20