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

TensorFlow自定义RNN Cell中使用add_loss方法的报错问题

解决图执行模式下自定义RNN Cell使用add_loss的InaccessibleTensorError问题

问题根源

当RNN层设置unroll=False时,TensorFlow会通过tf.while_loop迭代处理时间步。此时在自定义Cell的call方法内调用self.add_loss()添加的损失张量,属于while_loop的内部作用域。而Model.train_step()中访问self.losses时处于外部作用域,导致无法访问这些内部张量,触发InaccessibleTensorError。

可行解决方案

方案1:在Model层统一计算并添加损失

避免在Cell内部调用add_loss,将损失计算逻辑抽离到Model的call方法中,确保损失张量处于外部作用域:

# 抽离损失计算逻辑
def compute_step_loss(inputs):
    return tf.reduce_sum(inputs)

class BarCell(tf.keras.layers.Layer):
    """自定义RNN Cell,仅处理状态和输出计算"""
    def __init__(self, **kwargs):
        super().__init__(**kwargs)
        self.state_size = tf.TensorShape([1])

    def call(self, inputs, states, training=None):
        output = tf.reduce_sum(inputs, axis=1) + tf.constant(1.0)
        new_state = states[0] + 1
        return output, [new_state]

class FooModel(tf.keras.Model):
    def __init__(self, rnn=None, **kwargs):
        super().__init__(**kwargs)
        self.rnn = rnn

    def call(self, inputs, training=None):
        output = self.rnn(inputs, training=training)
        
        if training:
            # 转置输入,将时间步维度放到首位,方便逐个处理
            time_steps_inputs = tf.transpose(inputs, perm=[1, 0, 2])
            # 计算每个时间步的损失
            step_losses = tf.map_fn(
                compute_step_loss,
                time_steps_inputs,
                fn_output_signature=tf.float32
            )
            # 计算批次内的平均损失并添加到模型
            self.add_loss(tf.reduce_mean(step_losses))
        
        return output

方案2:将损失作为Cell状态的一部分传递,在Model中收集

修改Cell的状态结构,将每个时间步的损失作为状态的一部分返回,然后在Model中提取并添加损失:

class BarCell(tf.keras.layers.Layer):
    def __init__(self, **kwargs):
        super().__init__(**kwargs)
        # 状态包含两部分:实际循环状态、当前时间步损失
        self.state_size = [tf.TensorShape([1]), tf.TensorShape([])]

    def call(self, inputs, states, training=None):
        actual_state, _ = states
        output = tf.reduce_sum(inputs, axis=1) + tf.constant(1.0)
        step_loss = tf.reduce_sum(inputs)
        new_actual_state = actual_state + 1
        # 返回输出、新状态(实际状态+当前损失)
        return output, [new_actual_state, step_loss]

class FooModel(tf.keras.Model):
    def __init__(self, rnn=None, **kwargs):
        super().__init__(**kwargs)
        self.rnn = rnn

    def call(self, inputs, training=None):
        # 初始化状态:实际状态为0,初始损失无意义设为0
        batch_size = tf.shape(inputs)[0]
        initial_state = [tf.zeros((batch_size, 1)), tf.zeros(()))]
        
        # RNN返回:(时间步输出序列, 最终实际状态, 最后一步损失)
        output, final_actual_state, _ = self.rnn(inputs, initial_state=initial_state, training=training)
        
        if training:
            # 重新遍历输入计算所有时间步损失
            time_steps_inputs = tf.transpose(inputs, perm=[1, 0, 2])
            step_losses = tf.map_fn(lambda x: tf.reduce_sum(x), time_steps_inputs, fn_output_signature=tf.float32)
            self.add_loss(tf.reduce_mean(step_losses))
        
        return output

验证效果

使用原测试用例,修改后在is_eager=False且unroll=False的情况下,Model.fit()可正常运行,梯度计算能正确包含自定义损失,无报错。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.18 15:26:08