如何构建支持未知输入形状的自定义Keras ConvLSTM层
Keras自定义ConvLSTM层报错解决方案
核心问题
你对报错原因的判断正确:Keras模型构建阶段输入张量的batch维度为动态维度,静态值为None,无法直接作为参数传入tf.zeros等需要固定尺寸的接口,需要调整为获取运行时动态维度,同时修复代码中其他逻辑和语法问题。
具体修复点
- 状态初始化时使用
tf.shape(inputs)获取运行时动态维度,替代静态的inputs.shape属性 - 子层(Padding2D、Conv2D)统一在
__init__方法中实例化,保证权重被正确注册到模型 - 修复类内部参数调用错误,所有
__init__中定义的参数调用时补充self前缀 - 修复
tf.zeros参数格式错误,将通道数放入尺寸列表中 - 补充自定义层必需的
**kwargs参数,符合Keras层开发规范 - 未导入的算子统一使用
tf或tf.keras.layers前缀调用
修复后完整可运行代码
import tensorflow as tf from tensorflow import keras from keras.layers import InputSpec, Layer, Conv2D class Padding2D(Layer): def __init__(self, padding = (1,1), **kwargs): self.padding = tuple(padding) self.input_spec = [InputSpec(ndim = 4)] super(Padding2D,self).__init__(**kwargs) def compute_output_shape(self, s): return (s[0], s[1] + 2*self.padding[0], s[2] + 2*self.padding[1], s[3]) def call(self, x): w_pad, h_pad = self.padding return tf.pad(x, [[0,0], [h_pad,h_pad],[w_pad,w_pad],[0,0]]) class ConvLSTM(Layer): def __init__(self, out_channels, kernel_size=5, forget_bias=1.0, padding=0, **kwargs): super(ConvLSTM, self).__init__(**kwargs) self.out_channels = out_channels self.forget_bias = forget_bias self.padding = padding self.kernel_size = kernel_size # 子层在init中实例化,保证权重被正确注册 self.pad_layer = Padding2D(padding = (padding,padding)) self.conv_layer = Conv2D(4 * self.out_channels, kernel_size, strides=1) def build(self, input_shape): self.states = None super().build(input_shape) def call(self, inputs): if self.states is None: # 用tf.shape获取运行时动态维度 batch_size = tf.shape(inputs)[0] h = tf.shape(inputs)[1] w = tf.shape(inputs)[2] # 初始化两个状态张量 self.states = ( tf.zeros([batch_size, h, w, self.out_channels]), tf.zeros([batch_size, h, w, self.out_channels]) ) c, h_state = self.states if not (len(c.shape) == 4 and len(h_state.shape) == 4 and len(inputs.shape) == 4): raise TypeError("Incorrect shapes") inputs_h = tf.concat((inputs, h_state), axis=3) padded_inputs_h = self.pad_layer(inputs_h) i_j_f_o = self.conv_layer(padded_inputs_h) i = i_j_f_o[:,:,:,: self.out_channels] j = i_j_f_o[:,:,:,self.out_channels : 2*self.out_channels] f = i_j_f_o[:,:,:, 2*self.out_channels : 3*self.out_channels] o = i_j_f_o[:,:,:, 3*self.out_channels :] new_c = c * tf.sigmoid(f + self.forget_bias) + tf.sigmoid(i) * tf.tanh(j) new_h = tf.tanh(new_c) * tf.sigmoid(o) self.states = (new_c, new_h) return new_h input0 = tf.keras.Input(shape= (2,2,1)) x = ConvLSTM(out_channels=1)(input0) model = tf.keras.Model(input0,x) print(model(tf.ones((1,2,2,1))))
额外说明
如果你的ConvLSTM需要处理时序序列(即输入维度为[batch, timestep, height, width, channel]),可以在call方法中遍历时间维度逐帧计算状态即可,当前实现为单帧输入的ConvLSTM逻辑。
内容的提问来源于stack exchange,提问作者danix
相关产品推荐
相关产品推荐

