如何在Keras LSTM中获取每次predict调用时的内部遗忘门值?
Absolutely, you can get the forget gate values for every LSTM node each time you call predict() using Keras—no need to jump to another library immediately. Here are practical ways to make this happen, plus an alternative if you prefer a more intuitive workflow.
方法1:用Keras函数式API提取中间门输出
Keras doesn't expose LSTM gate outputs by default, but you can build a "monitoring model" using the functional API to tap into these internal values. Here's how:
First, let's assume you have a pre-trained LSTM model like this:
from tensorflow.keras.models import Sequential from tensorflow.keras.layers import LSTM, Dense # Your original trained model original_model = Sequential([ LSTM(64, input_shape=(10, 1)), Dense(1) ]) original_model.load_weights("your_trained_weights.h5")
Next, we'll create a new model that outputs the forget gate values from your existing LSTM layer. We can use Keras backend functions to compute the gates directly:
from tensorflow.keras.models import Model from tensorflow.keras import backend as K # Grab the LSTM layer from your original model lstm_layer = original_model.layers[0] inputs = original_model.input units = lstm_layer.units # Extract weights for the forget gate (split from LSTM's full weight set) W, U, b = lstm_layer.get_weights() W_f = W[:, :units] # Input weights for forget gate U_f = U[:, :units] # Recurrent weights for forget gate b_f = b[:units] # Bias for forget gate # Define a function to compute forget gates across all time steps def compute_forget_gates(input_data): forget_gates = [] h_prev = K.zeros((K.shape(input_data)[0], units)) # Initial hidden state for t in range(K.shape(input_data)[1]): x_t = input_data[:, t, :] # Calculate forget gate: sigmoid(W_f * x_t + U_f * h_prev + b_f) f_t = K.sigmoid(K.dot(x_t, W_f) + K.dot(h_prev, U_f) + b_f) forget_gates.append(f_t) # Update hidden state to match LSTM's sequence flow # (You can skip full LSTM state updates if you only care about forget gates) return K.stack(forget_gates, axis=1) # Build the model that returns forget gates alongside (or instead of) original output forget_gate_model = Model(inputs=inputs, outputs=[original_model.output, compute_forget_gates(inputs)])
Now, every time you call forget_gate_model.predict(your_input_data), you'll get a tuple with your original model output and the forget gate values (shape: (batch_size, sequence_length, lstm_units)).
方法2:自定义LSTM层直接输出遗忘门
If you want a cleaner setup, you can create a custom LSTM layer that returns forget gates as part of its output:
from tensorflow.keras.layers import LSTM from tensorflow.keras import backend as K class LSTMWithForgetGate(LSTM): def call(self, inputs): # Get standard LSTM outputs and states outputs, (hidden_state, cell_state) = super().call(inputs, return_state=True) # Compute forget gate values using the layer's internal weights W_f = self.kernel[:, :self.units] U_f = self.recurrent_kernel[:, :self.units] b_f = self.bias[:self.units] self.forget_gates = [] h_prev = K.zeros((K.shape(inputs)[0], self.units)) for t in range(K.shape(inputs)[1]): x_t = inputs[:, t, :] f_t = K.sigmoid(K.dot(x_t, W_f) + K.dot(h_prev, U_f) + b_f) self.forget_gates.append(f_t) h_prev = self.recurrent_activation(...) # Optional: full state update self.forget_gates = K.stack(self.forget_gates, axis=1) # Return both original output and forget gates, or just the gates return [outputs, self.forget_gates] # Build your model with this custom layer model = Sequential([ LSTMWithForgetGate(64, input_shape=(10, 1)), Dense(1) ])
Train this model as usual, and predict() will return both your target output and the forget gate values.
更便捷的替代库:PyTorch
If you find Keras's approach a bit clunky, PyTorch's dynamic computation graph makes extracting intermediate values far more intuitive. You can easily define an LSTM module that returns forget gates in its forward pass:
import torch import torch.nn as nn class LSTMWithForgetGate(nn.Module): def __init__(self, input_size, hidden_size): super().__init__() self.lstm = nn.LSTM(input_size, hidden_size) self.hidden_size = hidden_size # Extract LSTM weights self.W_ih, self.W_hh, self.b_ih, self.b_hh = self.lstm.parameters() def forward(self, x): output, (h_n, c_n) = self.lstm(x) # Split weights for the forget gate W_ih_f = self.W_ih[:self.hidden_size, :] W_hh_f = self.W_hh[:self.hidden_size, :] b_ih_f = self.b_ih[:self.hidden_size] b_hh_f = self.b_hh[:self.hidden_size] forget_gates = [] h_prev = torch.zeros(1, x.size(1), self.hidden_size).to(x.device) for t in range(x.size(0)): x_t = x[t, :, :] # Calculate forget gate f_t = torch.sigmoid(torch.matmul(x_t, W_ih_f.T) + b_ih_f + torch.matmul(h_prev, W_hh_f.T) + b_hh_f) forget_gates.append(f_t) h_prev = ... # Optional: update hidden state forget_gates = torch.stack(forget_gates, dim=0) return output, forget_gates # Usage example model = LSTMWithForgetGate(input_size=1, hidden_size=64) input_data = torch.randn(10, 32, 1) # Shape: (sequence_length, batch_size, input_size) output, forget_gates = model(input_data) # forget_gates shape: (sequence_length, batch_size, hidden_size)
With PyTorch, you don't need to build separate models—just compute and return the gates directly in the forward method, which feels much more straightforward for this kind of task.
内容的提问来源于stack exchange,提问作者humble_me

