自定义RNN Cell的Keras模型调用正常但保存时触发维度不匹配错误
你遇到的这个保存模型时的维度不匹配问题,根源在于两个核心点:
自定义RNNCell未声明constants的大小:Keras的RNNCell类需要通过
constants_size属性告知框架constants的形状信息。你的MinimalRNNCell只设置了state_size,却没有指定constants_size,导致模型序列化(保存)时,Keras无法正确推断constants的形状,错误地将其与输入x的形状对齐(变成三维的[?, ?, 5]),而不是保持你原本定义的二维形状[?, 1]。广播依赖动态形状,静态推断失败:在
predict时,你传入的是固定时间步长(1)的输入,TensorFlow的动态广播机制可以临时兼容维度差异;但模型保存时需要做静态形状推断,此时matmul(inputs, self.kernel)的形状是[?, ?, 32],而被错误推断形状的constants是[?, ?, 5],两者维度完全不匹配,自然触发ValueError。
要解决这个问题,需要做两个关键修改:
1. 为自定义RNNCell添加constants_size属性
在MinimalRNNCell的__init__方法中,明确声明constants的大小。因为你的z是单个标量输入,这里设置为1:
def __init__(self, units, **kwargs): self.units = units self.state_size = units self.constants_size = 1 # 新增这一行,告诉Keras constants的大小 super(MinimalRNNCell, self).__init__(**kwargs)
2. 在call方法中显式对齐constants的形状
确保constants的形状能和matmul结果正确广播。我们可以通过扩展维度,把原本的[batch, 1]转换成[batch, 1, 1],这样TensorFlow就能自动广播到[batch, timesteps, 32]的形状:
def call(self, inputs, states=None, constants=None, *args, **kwargs): prev_output = states[0] const = constants[0] # 扩展两个维度,匹配timesteps和units维度 const = tf.expand_dims(tf.expand_dims(const, axis=1), axis=2) h = matmul(inputs, self.kernel) + const output = h + matmul(prev_output, self.recurrent_kernel) return output, [output]
如果你的z本来应该是和RNN单元数(32)同维度的向量(比如作为偏置项),那更合理的做法是调整输入形状并适配:
- 构建模型时把
z的输入改为(32,):z = tfk.Input((32,), name='z') - call方法中只扩展时间步维度:
const = tf.expand_dims(constants[0], axis=1)
import tensorflow as tf from tensorflow.linalg import matmul import tensorflow.keras as tfk import tensorflow.keras.backend as K import numpy as np class MinimalRNNCell(tfk.layers.Layer): def __init__(self, units, **kwargs): self.units = units self.state_size = units self.constants_size = 1 # 新增:声明constants大小 super(MinimalRNNCell, self).__init__(**kwargs) def build(self, input_shape): self.kernel = self.add_weight(shape=(input_shape[-1], self.units), initializer='uniform', name='kernel') self.recurrent_kernel = self.add_weight( shape=(self.units, self.units), initializer='uniform', name='recurrent_kernel') self.built = True def call(self, inputs, states=None, constants=None, *args, **kwargs): prev_output = states[0] const = constants[0] # 扩展维度以匹配广播要求 const = tf.expand_dims(tf.expand_dims(const, axis=1), axis=2) h = matmul(inputs, self.kernel) + const output = h + matmul(prev_output, self.recurrent_kernel) return output, [output] def get_config(self): return dict(super().get_config(), **{'units': self.units}) cell = MinimalRNNCell(32) x = tfk.Input((None, 5), name='x') z = tfk.Input((1,), name='z') layer = tfk.layers.RNN(cell, name='rnn') y = layer(x, constants=[z]) model = tfk.Model(inputs=[x, z], outputs=[y]) model.compile(optimizer='adam', loss='mse') model.predict([np.array([[[0,0,0,0,0]]]), np.array([[0]])]) model.save('tmp.model')
运行这段修改后的代码,模型就能正常保存了。
内容的提问来源于stack exchange,提问作者Itamar Turner-Trauring

