如何在Keras LSTM层中修改门控(遗忘门、输入门等)?
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 (likehard_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
Usekernel_initializer(for input-to-gate weights) andrecurrent_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
buildmethod and including them in the gate calculations. - Feed extra inputs to a gate: Pass additional features to specific gates by modifying the
callmethod 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_configmethod 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

