基于PyTorch的RNN模型能否通过输出结果反推输入?
Great question! Let's break this down specifically for your nn.RNN example:
myrnn = nn.RNN(4, 2, 1, batch_first=True) expected_out, hidden = myrnn(input) # expected_out shape: (5, 1, 2)
Short Answer
Yes, but only if you have full access to the RNN's parameters and initial hidden state — and crucially, you won't get a unique input. There will be infinitely many possible inputs that produce the given output.
Detailed Explanation
To understand why, let's start with the core math of the single-layer RNN you're using:
For each time step ( t ), the hidden state (which is your output when return_sequence=True) is calculated as:
[ h_t = \tanh(W_{ih} x_t + b_{ih} + W_{hh} h_{t-1} + b_{hh}) ]
Where:
- ( W_{ih} ): Input-to-hidden weight matrix (shape:
(hidden_size, input_size)=(2,4)in your case) - ( b_{ih} ): Input-to-hidden bias (shape:
(2,)) - ( W_{hh} ): Hidden-to-hidden weight matrix (shape:
(2,2)) - ( b_{hh} ): Hidden-to-hidden bias (shape:
(2,)) - ( h_{t-1} ): Hidden state from the previous time step (or initial hidden state ( h_0 ) for the first step)
Key Observations for Your Example
Your expected_out has shape (5,1,2), meaning:
- Batch size = 5
- Sequence length = 1 (each input sample is a single time step)
- Hidden size = 2
Assuming you didn't manually set ( h_0 ), PyTorch defaults it to a tensor of zeros with shape (num_layers, batch_size, hidden_size) = (1,5,2).
Steps to Infer Possible Inputs
If you have access to all RNN parameters (via myrnn.state_dict()), here's how you can find valid inputs:
- Extract RNN parameters: Get ( W_{ih}, b_{ih}, W_{hh}, b_{hh} ) from the model.
- Compute the pre-tanh value: Since ( h_t = \tanh(\text{pre_tanh}) ), we can reverse this with ( \text{pre_tanh} = \text{arctanh}(h_t) ). Note: This only works if ( h_t ) is strictly within
(-1, 1)(which it almost always is in practice, since tanh never actually reaches ±1 numerically). - Set up the linear equation: Rearrange the RNN formula to solve for ( x_t ):
[ W_{ih} x_t = \text{pre_tanh} - b_{ih} - W_{hh} h_0 - b_{hh} ] - Solve the underdetermined system: Your ( W_{ih} ) is a
2x4matrix — this means we have 2 equations but 4 variables (the input features). Such systems have infinitely many solutions. You can use a least-squares solver to find one valid input, e.g., with PyTorch'storch.linalg.lstsq().
Example Code Snippet
import torch import torch.nn as nn # Your RNN setup myrnn = nn.RNN(4, 2, 1, batch_first=True) expected_out = torch.tensor([[[-0.7773, -0.2031]], [[-0.4129, -0.1802]], [[ 0.0599, -0.0151]], [[-0.9273, 0.2683]], [[ 0.6161, 0.5412]]]) # Get parameters W_ih = myrnn.weight_ih_l0 b_ih = myrnn.bias_ih_l0 W_hh = myrnn.weight_hh_l0 b_hh = myrnn.bias_hh_l0 # Default initial hidden state (all zeros) h0 = torch.zeros(1, expected_out.shape[0], 2) # Process each batch sample pre_tanh = torch.arctanh(expected_out.squeeze(1)) # Shape: (5,2) # Compute the right-hand side of the equation rhs = pre_tanh - b_ih - (W_hh @ h0.squeeze(0).T).T - b_hh # Solve for x (least-squares solution) x_solution = torch.linalg.lstsq(W_ih.T, rhs.T).solution.T # Shape: (5,4) # Verify: Pass the solution back through the RNN output, _ = myrnn(x_solution.unsqueeze(1)) print(torch.allclose(output, expected_out, atol=1e-4)) # Should be True (within numerical error)
Critical Limitations
- No unique input: Since the input dimension (4) is larger than the hidden size (2), there are infinitely many inputs that will produce the same output. The least-squares solution is just one of them.
- Requires full parameter access: If you don't know the RNN's weights and biases, you can't reverse-engineer the input — different parameter sets can map completely different inputs to the same output.
- Numerical stability: If the output ( h_t ) is very close to ±1,
arctanh()will produce extremely large values, leading to unstable calculations.
内容的提问来源于stack exchange,提问作者BING

