基于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

