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

自定义RNN Cell的Keras模型调用正常但保存时触发维度不匹配错误

问题原因分析

你遇到的这个保存模型时的维度不匹配问题,根源在于两个核心点:

  1. 自定义RNNCell未声明constants的大小:Keras的RNNCell类需要通过constants_size属性告知框架constants的形状信息。你的MinimalRNNCell只设置了state_size,却没有指定constants_size,导致模型序列化(保存)时,Keras无法正确推断constants的形状,错误地将其与输入x的形状对齐(变成三维的[?, ?, 5]),而不是保持你原本定义的二维形状[?, 1]。

  2. 广播依赖动态形状,静态推断失败:在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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.29 20:24:10