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

如何在Keras LSTM中获取每次predict调用时的内部遗忘门值?

获取Keras LSTM遗忘门值的可行方案

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.15 08:01:57