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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.15 11:22:07