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

Keras视频帧预测:4帧输入转1帧输出的网络构建求助

解决Keras DNN视频帧预测:4帧输入转1帧输出的方案

Got it, let's fix this frame prediction problem for you. The core issue you're facing is adapting a network that outputs the same number of frames as input to instead collapse 4 consecutive frames into a single predicted next frame. Here's a straightforward approach with concrete code examples:

第一步:调整输入数据维度

First, let's get your input shape aligned with Keras' preferred format. Your current input is (number_samples, 4, 60, 60) — we need to reshape this to play nicely with convolutional layers:

  • If you're working with grayscale frames:
    • For 2D convolution-based networks: Convert to (number_samples, 60, 60, 4) (treat the 4 frames as 4 separate "channels" per pixel).
    • For 3D convolution-based networks (better for explicit spatiotemporal features): Convert to (number_samples, 4, 60, 60, 1) (add a single channel dimension for grayscale).

You can do this with NumPy or TensorFlow:

# 2D卷积输入转换
import numpy as np
input_data = np.moveaxis(input_data, 1, -1)  # 从 (samples,4,60,60) 转到 (samples,60,60,4)

# 3D卷积输入转换
input_data = np.expand_dims(input_data, axis=-1)  # 从 (samples,4,60,60) 转到 (samples,4,60,60,1)

第二步:构建适配的网络结构

The key here is replacing any "per-frame" layers (like TimeDistributed) with layers that fuse information across all 4 input frames. Below are two proven architectures:

方案1:基于2D卷积的通道融合网络

This treats the 4 frames as extra channels, using 2D convolutions to learn spatial patterns combined with temporal information from the 4 channels:

import tensorflow as tf
from tensorflow.keras import layers, Model

# 输入形状:(60,60,4)
input_shape = (60, 60, 4)
inputs = layers.Input(shape=input_shape)

# 特征提取
x = layers.Conv2D(32, (3,3), activation='relu', padding='same')(inputs)
x = layers.MaxPooling2D((2,2), padding='same')(x)
x = layers.Conv2D(64, (3,3), activation='relu', padding='same')(x)
x = layers.MaxPooling2D((2,2), padding='same')(x)
x = layers.Conv2D(128, (3,3), activation='relu', padding='same')(x)

# 上采样恢复原尺寸
x = layers.UpSampling2D((2,2))(x)
x = layers.Conv2D(64, (3,3), activation='relu', padding='same')(x)
x = layers.UpSampling2D((2,2))(x)
x = layers.Conv2D(32, (3,3), activation='relu', padding='same')(x)

# 输出单帧:先得到 (60,60,1),再调整为你需要的 (1,60,60) 格式
outputs = layers.Conv2D(1, (3,3), activation='sigmoid', padding='same')(x)
outputs = layers.Reshape((1, 60, 60))(outputs)  # 匹配输出维度 (number_samples,1,60,60)

# 编译模型
model = Model(inputs=inputs, outputs=outputs)
model.compile(optimizer='adam', loss='mean_squared_error')
model.summary()

方案2:基于3D卷积的时空特征网络

3D convolutions explicitly learn patterns across time (the 4 frames) and space, which is often more effective for video tasks:

# 输入形状:(4,60,60,1)
input_shape = (4, 60, 60, 1)
inputs = layers.Input(shape=input_shape)

# 时空特征提取
x = layers.Conv3D(32, (2,3,3), activation='relu', padding='same')(inputs)
x = layers.MaxPooling3D((1,2,2), padding='same')(x)  # 只在空间维度下采样,保留时间维度
x = layers.Conv3D(64, (2,3,3), activation='relu', padding='same')(x)
x = layers.MaxPooling3D((1,2,2), padding='same')(x)

# 上采样恢复空间尺寸
x = layers.UpSampling3D((1,2,2))(x)
x = layers.Conv3D(64, (2,3,3), activation='relu', padding='same')(x)
x = layers.UpSampling3D((1,2,2))(x)
x = layers.Conv3D(32, (2,3,3), activation='relu', padding='same')(x)

# 输出时间维度为1的帧,调整为目标格式
outputs = layers.Conv3D(1, (1,3,3), activation='sigmoid', padding='same')(x)
outputs = layers.Reshape((1, 60, 60))(outputs)

model = Model(inputs=inputs, outputs=outputs)
model.compile(optimizer='adam', loss='mean_squared_error')
model.summary()

关键注意事项

  • 数据配对: Make sure your training data is correctly paired: each input of 4 frames (t, t+1, t+2, t+3) should map to the output frame t+4.
  • 激活函数: Use sigmoid if your pixel values are normalized to [0,1], or skip activation (use linear) if you're working with raw [0,255] values.
  • 损失函数: Mean Squared Error (MSE) is standard for pixel-wise regression tasks like frame prediction; you can also try Mean Absolute Error (MAE) if you want to reduce sensitivity to outliers.

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.25 07:10:59