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

如何在Keras LSTM层中修改门控(遗忘门、输入门等)?

How to Modify Gating Mechanisms (Forget Gate, Input Gate, etc.) in Keras LSTM Layers

Great question! Being able to tweak LSTM gates is a game-changer when you need to tailor a recurrent model to your specific task—whether it’s adjusting activation functions, adding custom connections like peepholes, or even redefining how gates compute their outputs. Let’s break down the two main approaches you can take in Keras.

1. Quick Adjustments to Default LSTM Gates

If you just need to modify shared parameters across all gates (like activation functions or regularizers), you don’t need to rebuild the entire layer. The built-in keras.layers.LSTM class has arguments that let you tweak these easily:

  • Change the recurrent activation (used by forget, input, and output gates)
    By default, all gates use the sigmoid activation. You can swap this out for another built-in function (like hard_sigmoid) or even a custom one:

    from tensorflow import keras
    from keras.layers import LSTM
    
    lstm_layer = LSTM(
        units=64,
        recurrent_activation='hard_sigmoid',  # Applies to all gates
        activation='tanh',  # Activation for cell state updates
        recurrent_regularizer=keras.regularizers.l2(0.01)  # Add L2 regularization to gate weights
    )
    
  • Customize weight initializers
    Use kernel_initializer (for input-to-gate weights) and recurrent_initializer (for hidden-state-to-gate weights) to control how gate parameters are initialized:

    lstm_layer = LSTM(
        units=64,
        kernel_initializer='he_normal',
        recurrent_initializer='glorot_uniform'
    )
    

2. Build a Fully Custom LSTM Layer for Gate-Specific Modifications

If you need to modify individual gates (e.g., change only the forget gate’s activation, add peephole connections, or feed extra inputs to a gate), you’ll need to create a custom LSTM cell and wrap it in an RNN layer. This gives you full control over every part of the gate logic.

Here’s a concrete example where we set a custom activation for the forget gate (ReLU instead of sigmoid) while keeping the other gates as default:

import tensorflow as tf
from keras.layers import RNN, Layer

class CustomLSTMCell(Layer):
    def __init__(self, units, forget_gate_activation='relu', **kwargs):
        self.units = units
        # Convert activation string to actual function
        self.forget_activation = tf.keras.activations.get(forget_gate_activation)
        super().__init__(**kwargs)

    def build(self, input_shape):
        # Weights for input -> all 4 gates (forget, input, cell, output)
        self.kernel = self.add_weight(
            shape=(input_shape[-1], self.units * 4),
            initializer='glorot_uniform',
            name='input_kernel'
        )
        # Weights for hidden state -> all 4 gates
        self.recurrent_kernel = self.add_weight(
            shape=(self.units, self.units * 4),
            initializer='orthogonal',
            name='recurrent_kernel'
        )
        # Biases for each gate (split into 4 parts later)
        self.bias = self.add_weight(
            shape=(self.units * 4,),
            initializer='zeros',
            name='gate_biases'
        )
        # Optional: Custom bias init for forget gate (e.g., start with higher values to reduce forgetting)
        # self.bias.assign(tf.concat([tf.ones(self.units), tf.zeros(self.units*3)], axis=0))
        super().build(input_shape)

    def call(self, inputs, states):
        h_prev = states[0]  # Previous hidden state
        c_prev = states[1]  # Previous cell state

        # Split weights/biases into individual gate components
        kernel_f, kernel_i, kernel_c, kernel_o = tf.split(self.kernel, 4, axis=1)
        rec_kernel_f, rec_kernel_i, rec_kernel_c, rec_kernel_o = tf.split(self.recurrent_kernel, 4, axis=1)
        bias_f, bias_i, bias_c, bias_o = tf.split(self.bias, 4, axis=0)

        # Compute each gate with custom logic
        # Forget gate (using our custom activation)
        f = self.forget_activation(
            tf.matmul(inputs, kernel_f) + tf.matmul(h_prev, rec_kernel_f) + bias_f
        )
        # Input gate (default sigmoid)
        i = tf.keras.activations.sigmoid(
            tf.matmul(inputs, kernel_i) + tf.matmul(h_prev, rec_kernel_i) + bias_i
        )
        # Cell update (default tanh)
        c_update = tf.keras.activations.tanh(
            tf.matmul(inputs, kernel_c) + tf.matmul(h_prev, rec_kernel_c) + bias_c
        )
        # Output gate (default sigmoid)
        o = tf.keras.activations.sigmoid(
            tf.matmul(inputs, kernel_o) + tf.matmul(h_prev, rec_kernel_o) + bias_o
        )

        # Update cell and hidden states
        c_current = f * c_prev + i * c_update
        h_current = o * tf.keras.activations.tanh(c_current)

        return h_current, [h_current, c_current]

    def get_config(self):
        # Ensure the layer can be saved/loaded
        config = super().get_config()
        config.update({
            'units': self.units,
            'forget_gate_activation': tf.keras.activations.serialize(self.forget_activation)
        })
        return config

# Use the custom cell in an RNN layer
custom_lstm = RNN(CustomLSTMCell(units=64, forget_gate_activation='relu'), return_sequences=True)

Extending This Further

If you want to make more advanced modifications, here are some ideas:

  • Add peephole connections: Let gates depend on the previous cell state by adding peephole weights in the build method and including them in the gate calculations.
  • Feed extra inputs to a gate: Pass additional features to specific gates by modifying the call method to accept extra inputs and adding corresponding weights.
  • Redefine gate logic: For example, make the forget gate depend on global context (like an attention vector) instead of just the current input and previous hidden state.

Key Notes

  • Always implement the get_config method if you plan to save and reload your custom layer.
  • Test your modifications carefully—changing gate behavior can drastically affect model performance, so start with small tweaks and iterate.

内容的提问来源于stack exchange,提问作者Nazrin Taba-Tabai

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.27 19:07:25