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

Keras报错:LSTM.__init__()缺少必填参数'units'问题排查

问题分析

错误的核心原因是:TT_LSTM继承自Keras的原生LSTM类,而LSTM的构造函数要求units作为必填位置参数,但当前代码中self.units是在调用父类初始化之后才计算的,导致super(TT_LSTM, self).__init__(**kwargs)时没有传递units参数,触发TypeError。

解决方案

修改TT_LSTM的__init__方法,调整代码顺序:先计算出units的值,再将其作为参数传递给父类LSTM的初始化函数,同时把父类需要的其他参数一并传入,避免重复定义和潜在冲突。

修改后的完整__init__代码

def __init__(self,
             tt_input_shape, tt_output_shape, tt_ranks,
             activation='tanh',
             recurrent_activation='hard_sigmoid',
             use_bias=True,
             kernel_initializer='glorot_uniform',
             recurrent_initializer='orthogonal',
             bias_initializer='zeros',
             unit_forget_bias=True,
             kernel_regularizer=None,
             recurrent_regularizer=None,
             bias_regularizer=None,
             activity_regularizer=None,
             kernel_constraint=None,
             recurrent_constraint=None,
             bias_constraint=None,
             dropout=0.,
             recurrent_dropout=0.,
             debug=False,
             init_seed=11111986,
             **kwargs):
    # 1. 先计算units的值(由tt_output_shape的乘积得到)
    self.units = np.prod(np.array(tt_output_shape))
    
    # 2. 调用父类LSTM的初始化,显式传入units及所有父类需要的参数
    super(TT_LSTM, self).__init__(
        units=self.units,
        activation=activation,
        recurrent_activation=recurrent_activation,
        use_bias=use_bias,
        kernel_initializer=kernel_initializer,
        recurrent_initializer=recurrent_initializer,
        bias_initializer=bias_initializer,
        unit_forget_bias=unit_forget_bias,
        kernel_regularizer=kernel_regularizer,
        recurrent_regularizer=recurrent_regularizer,
        bias_regularizer=bias_regularizer,
        activity_regularizer=activity_regularizer,
        kernel_constraint=kernel_constraint,
        recurrent_constraint=recurrent_constraint,
        bias_constraint=bias_constraint,
        dropout=dropout,
        recurrent_dropout=recurrent_dropout,
        **kwargs
    )

    # 3. 设置自定义属性(移除父类已处理的重复属性)
    self.debug = debug
    self.init_seed = init_seed

    tt_input_shape = np.array(tt_input_shape)
    tt_output_shape = np.array(tt_output_shape)
    tt_ranks = np.array(tt_ranks)
    self.num_dim = tt_input_shape.shape[0]
    self.tt_input_shape = tt_input_shape
    self.tt_output_shape = tt_output_shape
    self.tt_ranks = tt_ranks
    self.state_spec = InputSpec(shape=(None, self.units))

关键修改点说明

  • 提前计算self.units,确保父类初始化时能获取到这个必填参数
  • 显式将父类LSTM需要的所有参数传入super().__init__,避免重复定义属性导致的冲突或冗余
  • 移除了原代码中在super之后重复设置的父类属性(如self.activation、self.use_bias等),这些属性会由父类初始化自动处理

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.11 12:52:40