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

寻求无框架依赖的LSTM/GRU原生实现以理解其内部机制

Understanding LSTMs via a Framework-Free Implementation

Hey there! I totally feel your frustration—when you're trying to dig into the actual mechanics of LSTMs/GRUs, framework tutorials and dense math papers can feel like they're skipping the most critical part: how all those equations translate to code that runs step-by-step. Let's fix that with a vanilla LSTM implementation from scratch in Python (easy to adapt to Java/C#/VB later) that includes backpropagation through time (BPTT) for supervised learning.

First, Why These Implementations Are Hard to Find

Most resources focus on frameworks because:

  • Hand-coding LSTMs is verbose and error-prone (no optimized ops under the hood)
  • Production code never uses raw implementations—frameworks are way faster
  • But for learning? Raw code is gold. Your keyword issue was probably using terms like "LSTM tutorial" which pulls up framework guides. Try searching for vanilla LSTM from scratch or LSTMs backpropagation through time manual code next time!

Framework-Free LSTM with Supervised Learning Demo

Below is a minimal LSTM implementation that learns to predict the next value in a simple sequence (e.g., given [1,2,3,4], predict 5). We'll include all core components: input/forget/output gates, cell state updates, and BPTT for weight updates.

import numpy as np

# Sigmoid activation (for gates)
def sigmoid(x):
    return 1 / (1 + np.exp(-x))

# Derivative of sigmoid
def sigmoid_deriv(x):
    return sigmoid(x) * (1 - sigmoid(x))

# Tanh activation (for cell state)
def tanh(x):
    return np.tanh(x)

# Derivative of tanh
def tanh_deriv(x):
    return 1 - np.tanh(x)**2

class VanillaLSTM:
    def __init__(self, input_size, hidden_size, output_size):
        # Initialize weights and biases for all gates
        self.input_size = input_size
        self.hidden_size = hidden_size
        self.output_size = output_size

        # Input gate weights: W_ii (input->input), W_hi (hidden->input), b_i
        self.W_ii = np.random.uniform(-0.1, 0.1, (hidden_size, input_size))
        self.W_hi = np.random.uniform(-0.1, 0.1, (hidden_size, hidden_size))
        self.b_i = np.zeros((hidden_size, 1))

        # Forget gate weights: W_if (input->forget), W_hf (hidden->forget), b_f
        self.W_if = np.random.uniform(-0.1, 0.1, (hidden_size, input_size))
        self.W_hf = np.random.uniform(-0.1, 0.1, (hidden_size, hidden_size))
        self.b_f = np.zeros((hidden_size, 1))

        # Cell state candidate weights: W_ig (input->candidate), W_hg (hidden->candidate), b_g
        self.W_ig = np.random.uniform(-0.1, 0.1, (hidden_size, input_size))
        self.W_hg = np.random.uniform(-0.1, 0.1, (hidden_size, hidden_size))
        self.b_g = np.zeros((hidden_size, 1))

        # Output gate weights: W_io (input->output), W_ho (hidden->output), b_o
        self.W_io = np.random.uniform(-0.1, 0.1, (hidden_size, input_size))
        self.W_ho = np.random.uniform(-0.1, 0.1, (hidden_size, hidden_size))
        self.b_o = np.zeros((hidden_size, 1))

        # Output layer weights: W_y (hidden->output), b_y
        self.W_y = np.random.uniform(-0.1, 0.1, (output_size, hidden_size))
        self.b_y = np.zeros((output_size, 1))

        # For backprop: store hidden states, cell states, and gate values
        self.hidden_states = []
        self.cell_states = []
        self.input_gates = []
        self.forget_gates = []
        self.candidate_cells = []
        self.output_gates = []

    def forward(self, x_seq):
        # Initialize hidden and cell state to zeros
        h_prev = np.zeros((self.hidden_size, 1))
        c_prev = np.zeros((self.hidden_size, 1))

        self.hidden_states = [h_prev]
        self.cell_states = [c_prev]

        for x in x_seq:
            x = x.reshape(-1, 1)  # Reshape to (input_size, 1)

            # Calculate input gate: i_t = sigmoid(W_ii*x_t + W_hi*h_{t-1} + b_i)
            i_t = sigmoid(np.dot(self.W_ii, x) + np.dot(self.W_hi, h_prev) + self.b_i)
            self.input_gates.append(i_t)

            # Calculate forget gate: f_t = sigmoid(W_if*x_t + W_hf*h_{t-1} + b_f)
            f_t = sigmoid(np.dot(self.W_if, x) + np.dot(self.W_hf, h_prev) + self.b_f)
            self.forget_gates.append(f_t)

            # Calculate cell state candidate: g_t = tanh(W_ig*x_t + W_hg*h_{t-1} + b_g)
            g_t = tanh(np.dot(self.W_ig, x) + np.dot(self.W_hg, h_prev) + self.b_g)
            self.candidate_cells.append(g_t)

            # Update cell state: c_t = f_t * c_{t-1} + i_t * g_t
            c_t = f_t * c_prev + i_t * g_t
            self.cell_states.append(c_t)

            # Calculate output gate: o_t = sigmoid(W_io*x_t + W_ho*h_{t-1} + b_o)
            o_t = sigmoid(np.dot(self.W_io, x) + np.dot(self.W_ho, h_prev) + self.b_o)
            self.output_gates.append(o_t)

            # Update hidden state: h_t = o_t * tanh(c_t)
            h_t = o_t * tanh(c_t)
            self.hidden_states.append(h_t)

            h_prev = h_t
            c_prev = c_t

        # Final output: y = W_y*h_t + b_y
        y_pred = np.dot(self.W_y, h_t) + self.b_y
        return y_pred

    def backward(self, x_seq, y_true, learning_rate=0.01):
        # Initialize gradients to zero
        dW_ii, dW_hi = np.zeros_like(self.W_ii), np.zeros_like(self.W_hi)
        dW_if, dW_hf = np.zeros_like(self.W_if), np.zeros_like(self.W_hf)
        dW_ig, dW_hg = np.zeros_like(self.W_ig), np.zeros_like(self.W_hg)
        dW_io, dW_ho = np.zeros_like(self.W_io), np.zeros_like(self.W_ho)
        dW_y = np.zeros_like(self.W_y)

        db_i, db_f, db_g, db_o, db_y = np.zeros_like(self.b_i), np.zeros_like(self.b_f), np.zeros_like(self.b_g), np.zeros_like(self.b_o), np.zeros_like(self.b_y)

        # Initial error: dL/dy = y_pred - y_true
        y_pred = self.forward(x_seq)
        dy = y_pred - y_true.reshape(-1, 1)

        # Gradient for output layer
        dW_y += np.dot(dy, self.hidden_states[-1].T)
        db_y += dy

        # Initialize hidden state gradient (dh_t)
        dh = np.dot(self.W_y.T, dy)

        # Iterate backwards through time steps
        for t in reversed(range(len(x_seq))):
            x = x_seq[t].reshape(-1, 1)
            h_prev = self.hidden_states[t]
            h_t = self.hidden_states[t+1]
            c_prev = self.cell_states[t]
            c_t = self.cell_states[t+1]

            i_t = self.input_gates[t]
            f_t = self.forget_gates[t]
            g_t = self.candidate_cells[t]
            o_t = self.output_gates[t]

            # Calculate gradients for cell state and output gate
            dc = dh * o_t * tanh_deriv(c_t)
            do = dh * tanh(c_t) * sigmoid_deriv(o_t)
            di = dc * g_t * sigmoid_deriv(i_t)
            df = dc * c_prev * sigmoid_deriv(f_t)
            dg = dc * i_t * tanh_deriv(g_t)

            # Update weight gradients
            dW_ii += np.dot(di, x.T)
            dW_hi += np.dot(di, h_prev.T)
            db_i += di

            dW_if += np.dot(df, x.T)
            dW_hf += np.dot(df, h_prev.T)
            db_f += df

            dW_ig += np.dot(dg, x.T)
            dW_hg += np.dot(dg, h_prev.T)
            db_g += dg

            dW_io += np.dot(do, x.T)
            dW_ho += np.dot(do, h_prev.T)
            db_o += do

            # Update hidden gradient for previous time step
            dh = np.dot(self.W_hi.T, di) + np.dot(self.W_hf.T, df) + np.dot(self.W_hg.T, dg) + np.dot(self.W_ho.T, do)

        # Clip gradients to prevent exploding gradients
        for grad in [dW_ii, dW_hi, dW_if, dW_hf, dW_ig, dW_hg, dW_io, dW_ho, dW_y, db_i, db_f, db_g, db_o, db_y]:
            np.clip(grad, -1, 1, out=grad)

        # Update weights and biases
        self.W_ii -= learning_rate * dW_ii
        self.W_hi -= learning_rate * dW_hi
        self.b_i -= learning_rate * db_i

        self.W_if -= learning_rate * dW_if
        self.W_hf -= learning_rate * dW_hf
        self.b_f -= learning_rate * db_f

        self.W_ig -= learning_rate * dW_ig
        self.W_hg -= learning_rate * dW_hg
        self.b_g -= learning_rate * db_g

        self.W_io -= learning_rate * dW_io
        self.W_ho -= learning_rate * dW_ho
        self.b_o -= learning_rate * db_o

        self.W_y -= learning_rate * dW_y
        self.b_y -= learning_rate * db_y

        # Return loss (MSE)
        return np.mean(dy**2)

# ------------------------------
# Supervised Learning Demo
# ------------------------------
if __name__ == "__main__":
    # Simple task: predict next number in sequence [1,2,3,4] -> 5, [2,3,4,5] ->6, etc.
    X = [np.array([1,2,3,4]), np.array([2,3,4,5]), np.array([3,4,5,6]), np.array([4,5,6,7])]
    y = [np.array([5]), np.array([6]), np.array([7]), np.array([8])]

    # Initialize LSTM: input_size=1 (each step is a single number), hidden_size=4, output_size=1
    lstm = VanillaLSTM(input_size=1, hidden_size=4, output_size=1)

    # Train for 1000 epochs
    epochs = 1000
    for epoch in range(epochs):
        total_loss = 0
        for x_seq, y_true in zip(X, y):
            # Reshape sequence to list of single-element arrays (each step is one input)
            x_seq = np.split(x_seq, len(x_seq))
            loss = lstm.backward(x_seq, y_true, learning_rate=0.005)
            total_loss += loss
        if (epoch + 1) % 100 == 0:
            print(f"Epoch {epoch+1}, Average Loss: {total_loss/len(X):.4f}")

    # Test the trained model
    test_seq = np.array([5,6,7,8])
    test_seq_split = np.split(test_seq, len(test_seq))
    prediction = lstm.forward(test_seq_split)
    print(f"\nTest sequence: {test_seq}, Predicted next value: {prediction[0][0]:.2f}")

Key Notes for Adapting to Other Languages

If you prefer Java/C#/VB, the core logic translates directly:

  • Replace NumPy arrays with native arrays or matrix libraries (e.g., MathNet.Numerics for C#, Apache Commons Math for Java)
  • Implement the activation functions and their derivatives as helper methods
  • Track all intermediate states (hidden/cell states, gates) in lists/arrays during forward pass
  • The BPTT loop runs backwards through time steps, calculating gradients for each weight matrix

Better Search Keywords for Future Reference

To find more raw implementations, use these terms:

  • vanilla LSTM implementation from scratch
  • LSTM backpropagation through time manual code
  • GRU raw code without ML framework
  • RNN weight update step-by-step code

These will filter out most framework-focused results and point you to code that exposes the inner workings.

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.21 04:03:42