LSTM序列图像生成模型输出固定问题排查求助
我用LSTM模型生成电影帧序列的下一张图像(因需求需要完整图像输入到下一轮序列迭代,所以未用CNN)。数据集为电影帧,序列构建规则:若一个场景含n张图像、序列长度为s,则输入为image_1~image_s,输出为image_s+1;下一组输入为image_2~image_s+1,输出为image_s+2,以此类推。
模型代码:
class LSTM(nn.Module): def __init__(self, input_len, hidden_size, num_layers): super(LSTM, self).__init__() self.hidden_size = hidden_size self.num_layers = num_layers self.lstm = nn.LSTM(input_len, hidden_size, num_layers, batch_first=True) self.output_layer = nn.Linear(hidden_size, input_len) self.dropout = nn.Dropout(.2) def forward(self, X): hidden_states = torch.zeros(self.num_layers, X.size(0), self.hidden_size, device=device) cell_states = torch.zeros(self.num_layers, X.size(0), self.hidden_size, device=device) out, _ = self.lstm(X, (hidden_states, cell_states)) out = self.dropout(out) out = self.output_layer(out[:, -1, :]) return out
训练代码:
def train(num_epochs, model, loss_func, optimizer): total_steps = loader.getSizeWithBatch() for epoch in range(num_epochs): loader.reset() for item in range(total_steps-1): element = loader.next()[0] x_images,y_image = element x_images = x_images.reshape(-1, sequence_len, input_len) output = model(x_images) y_image = y_image.reshape(-1,input_len) loss = loss_func(output, y_image) optimizer.zero_grad() loss.backward() optimizer.step() if (item + 1) % 1 == 0: print(f'Epoch: {epoch + 1}; Batch: {item + 1} / {total_steps};Loss: {loss.item():>4f}') if (epoch + 1) % int(config['SAVE']['model_save_interval']) == 0: if (epoch + 1) % int(config['SAVE']['clean_save_interval']) == 0: torch.save(model.state_dict(), os.path.join(config['PATH']['model_path'], config['PATH']['model_name'] + str(epoch+1))) else: torch.save(model.state_dict(), os.path.join(config['PATH']['model_path'], config['PATH']['model_name']))
Loader以预序列化张量加载图像节省内存,采用MSE损失与Adam优化器。
问题:训练时损失降到目标值0.003(图像已归一化到0-1),但预测时生成的是模糊的融合场景图,且无论输入来自哪个场景,输出完全相同(像素值差为0),类似数据集所有图像的叠加效果。尝试过加Dropout、增大hidden_size(当前128)、增加层数、调整学习率(从0.001降到0.0001),均无效。
1. LSTM状态初始化错误(最关键)
当前模型在每次forward时都将hidden和cell状态重置为0,导致LSTM无法捕捉序列的时序依赖关系——模型根本没学到帧之间的前后关联,只能学到数据集的全局均值,自然输出所有输入的平均图像。
修改方案:允许LSTM状态在序列迭代中传递,而非每次重置:
def forward(self, X, hidden=None): batch_size = X.size(0) # 仅在首次输入时初始化状态 if hidden is None: hidden_states = torch.zeros(self.num_layers, batch_size, self.hidden_size, device=device) cell_states = torch.zeros(self.num_layers, batch_size, self.hidden_size, device=device) else: hidden_states, cell_states = hidden # 保存更新后的状态,用于下一轮序列输入 out, (h_n, c_n) = self.lstm(X, (hidden_states, cell_states)) out = self.dropout(out) out = self.output_layer(out[:, -1, :]) return out, (h_n, c_n)
训练时,针对同一个场景的连续序列,需要传递状态;场景切换时重置状态:
def train(num_epochs, model, loss_func, optimizer): for epoch in range(num_epochs): loader.reset() hidden_state = None while loader.has_more_scenes(): # 每个场景单独处理,切换场景时重置状态 scene_loader = loader.get_next_scene_loader() hidden_state = None for x_images, y_image in scene_loader: x_images = x_images.reshape(-1, sequence_len, input_len) y_image = y_image.reshape(-1, input_len) # 传递状态时分离计算图,避免显存泄漏 if hidden_state is not None: hidden_state = (hidden_state[0].detach(), hidden_state[1].detach()) output, hidden_state = model(x_images, hidden_state) loss = loss_func(output, y_image) optimizer.zero_grad() loss.backward() optimizer.step() print(f'Epoch: {epoch + 1}; Loss: {loss.item():>4f}') # 模型保存逻辑不变
2. 输入维度过大导致模型坍缩
直接将整张图像展平成一维向量输入LSTM,维度极高(比如224×224×3的图像展平后是150528维),LSTM难以处理这么高维度的输入,很容易陷入全局最优解——输出所有样本的均值。
修改方案:添加CNN编码器将图像压缩为低维特征,再输入LSTM,最后用解码器还原为图像,既保留完整图像信息,又降低LSTM输入维度:
import torchvision class LSTMImageGenerator(nn.Module): def __init__(self, img_channels=3, img_size=224, hidden_size=256, num_layers=2): super().__init__() self.img_size = img_size # CNN编码器:将图像压缩为低维特征 self.encoder = nn.Sequential( nn.Conv2d(img_channels, 32, kernel_size=3, stride=2, padding=1), nn.ReLU(), nn.Conv2d(32, 64, kernel_size=3, stride=2, padding=1), nn.ReLU(), nn.Conv2d(64, 128, kernel_size=3, stride=2, padding=1), nn.ReLU(), nn.Flatten() ) # 预计算编码器输出特征长度 with torch.no_grad(): dummy = torch.randn(1, img_channels, img_size, img_size) self.feature_len = self.encoder(dummy).shape[1] self.hidden_size = hidden_size self.num_layers = num_layers self.lstm = nn.LSTM(self.feature_len, hidden_size, num_layers, batch_first=True) self.dropout = nn.Dropout(0.2) # CNN解码器:将LSTM输出还原为图像 self.decoder = nn.Sequential( nn.Linear(hidden_size, 128 * (img_size//8) * (img_size//8)), nn.ReLU(), nn.Unflatten(1, (128, img_size//8, img_size//8)), nn.ConvTranspose2d(128, 64, kernel_size=4, stride=2, padding=1), nn.ReLU(), nn.ConvTranspose2d(64, 32, kernel_size=4, stride=2, padding=1), nn.ReLU(), nn.ConvTranspose2d(32, img_channels, kernel_size=4, stride=2, padding=1), nn.Sigmoid() # 匹配图像归一化到0-1的范围 ) def forward(self, X, hidden=None): batch_size, seq_len, channels, h, w = X.shape # 对序列中的每帧图像编码 encoded_frames = [] for i in range(seq_len): frame = X[:, i, :, :, :] feat = self.encoder(frame) encoded_frames.append(feat) encoded_seq = torch.stack(encoded_frames, dim=1) # LSTM处理序列特征 if hidden is None: hidden_states = torch.zeros(self.num_layers, batch_size, self.hidden_size, device=device) cell_states = torch.zeros(self.num_layers, batch_size, self.hidden_size, device=device) else: hidden_states, cell_states = hidden out, (h_n, c_n) = self.lstm(encoded_seq, (hidden_states, cell_states)) out = self.dropout(out[:, -1, :]) # 解码为图像 generated_frame = self.decoder(out) return generated_frame, (h_n, c_n)
3. MSE损失的局限性
MSE损失会最小化预测与真实图像的像素级误差,天然倾向于生成所有样本的均值,导致图像模糊。
优化方案:替换为MSE+感知损失的组合,利用预训练CNN提取的特征计算损失,让模型学习到更细节的图像特征:
class PerceptualLoss(nn.Module): def __init__(self): super().__init__() # 使用预训练VGG16的前16层提取特征 vgg = torchvision.models.vgg16(pretrained=True).features[:16].eval() for param in vgg.parameters(): param.requires_grad = False self.vgg = vgg self.mse = nn.MSELoss() def forward(self, pred, target): # 计算像素级MSE损失 pixel_loss = self.mse(pred, target) # 计算特征级MSE损失 pred_feat = self.vgg(pred) target_feat = self.vgg(target) feat_loss = self.mse(pred_feat, target_feat) # 加权组合损失 return pixel_loss + 0.1 * feat_loss
4. 训练数据的序列连续性问题
检查Loader是否正确将同一个场景的帧分组为连续序列,如果存在跨场景的序列,模型无法学习到场景内的时序规律,只能输出全局均值。确保每个训练batch的序列都来自同一个场景,场景之间独立处理。
内容的提问来源于stack exchange,提问作者Tamás Csepely

