基于Numpy实现LSTM时的输入矩阵维度理解疑问
Hey there! Let's walk through this step by step, starting with what you already know about RNNs to build up to LSTMs. It's totally normal to get stuck on the matrix mechanics when first implementing LSTMs from scratch—let's clear it up.
First, Recap RNN Dimensions (To Align Our Baseline)
You mentioned you understand RNNs where x_t (shape (b, d): b = batch size, d = input feature dimension) interacts with weight matrices. Let's formalize that quickly to make sure we're on the same page:
- For a vanilla RNN, the hidden state update is:
h_t = np.tanh(x_t @ W_xh + h_prev @ W_hh + b_h)W_xh: Weight matrix mapping input to hidden state → shape(d, h)(so(b,d) @ (d,h) = (b,h))W_hh: Weight matrix mapping previous hidden state to current → shape(h, h)(so(b,h) @ (h,h) = (b,h))h_prev: Previous hidden state → shape(b, h)b_h: Bias term → shape(h,)(broadcasts to(b,h)automatically in NumPy)
LSTMs build on this core idea but add four separate "gates" (forget, input, candidate cell state, output) that each use x_t and h_prev—so we'll extend this dimension logic to each gate.
LSTM Gate-by-Gate Dimension Breakdown
Each gate in an LSTM follows the same input-to-output dimension pattern as the RNN's hidden state calculation, but with separate weights for each gate. Let's define h as the LSTM's hidden state dimension (same as the cell state dimension c).
1. Individual Gate Calculations (No Weight Merging)
For each gate, we compute a transformation of x_t and h_prev, then apply an activation function:
Forget Gate (
f_t): Controls what to discard from the cell statef_t = sigmoid(x_t @ W_xf + h_prev @ W_hf + b_f)W_xf: Input-to-forget gate weights → shape(d, h)W_hf: Hidden-to-forget gate weights → shape(h, h)b_f: Forget gate bias → shape(h,)- Output
f_t: Shape(b, h)(element-wise sigmoid squashes values to 0-1)
Input Gate (
i_t): Controls what new information to store in the cell statei_t = sigmoid(x_t @ W_xi + h_prev @ W_hi + b_i)- Dimensions match the forget gate:
W_xi=(d,h),W_hi=(h,h),b_i=(h,), outputi_t=(b,h)
- Dimensions match the forget gate:
Candidate Cell State (
c_tilde): Generates new candidate values to add to the cell statec_tilde = np.tanh(x_t @ W_xc + h_prev @ W_hc + b_c)- Dimensions:
W_xc=(d,h),W_hc=(h,h),b_c=(h,), outputc_tilde=(b,h)
- Dimensions:
Output Gate (
o_t): Controls what to output as the hidden stateo_t = sigmoid(x_t @ W_xo + h_prev @ W_ho + b_o)- Dimensions:
W_xo=(d,h),W_ho=(h,h),b_o=(h,), outputo_t=(b,h)
- Dimensions:
2. Cell State & Hidden State Updates
Once we have all four gates, we update the cell state and hidden state using element-wise multiplication (since we're applying gate "masks"):
# Update cell state: forget old values, add new candidate values c_t = f_t * c_prev + i_t * c_tilde # All terms shape (b, h) # Update hidden state: output a filtered version of the cell state h_t = o_t * np.tanh(c_t) # Both terms shape (b, h)
c_prev: Previous cell state → shape(b, h)(same as hidden state dimension)
Optimized Weight Merging (How Real Implementations Work)
Doing four separate matrix multiplications for x_t and h_prev is inefficient. Instead, we merge all four gate weights into single large matrices:
- Combine input weights:
W_x = np.concatenate([W_xf, W_xi, W_xc, W_xo], axis=1)→ shape(d, 4h) - Combine hidden weights:
W_h = np.concatenate([W_hf, W_hi, W_hc, W_ho], axis=1)→ shape(h, 4h) - Combine biases:
b = np.concatenate([b_f, b_i, b_c, b_o])→ shape(4h,)
Then we do one matrix multiplication for each input type, split the result into four parts, and apply activations:
# Single projection for all gates x_proj = x_t @ W_x # Shape (b, 4h) h_proj = h_prev @ W_h # Shape (b, 4h) combined = x_proj + h_proj + b # Shape (b, 4h) # Split into individual gate outputs f_t = sigmoid(combined[:, :h]) i_t = sigmoid(combined[:, h:2*h]) c_tilde = np.tanh(combined[:, 2*h:3*h]) o_t = sigmoid(combined[:, 3*h:])
This is the approach you'll usually see in NumPy implementations (and frameworks like TensorFlow/PyTorch under the hood) because it's faster.
Key Dimension Rules to Remember
- All gate outputs, cell states, and hidden states have shape
(batch_size, hidden_size) - Input weight matrices (per gate) are always
(input_size, hidden_size); merged is(input_size, 4*hidden_size) - Hidden weight matrices (per gate) are always
(hidden_size, hidden_size); merged is(hidden_size, 4*hidden_size) - Biases match the output dimension of their respective gates: per gate
(hidden_size,), merged(4*hidden_size,) - NumPy uses row-major ordering (each row is a sample), so we don't need to transpose matrices for batch operations—this keeps the math straightforward.
内容的提问来源于stack exchange,提问作者James

