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

基于PyTorch的RNN模型能否通过输出结果反推输入?

Can we infer the input of a PyTorch RNN from its output?

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:

  1. Extract RNN parameters: Get ( W_{ih}, b_{ih}, W_{hh}, b_{hh} ) from the model.
  2. 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).
  3. 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} ]
  4. Solve the underdetermined system: Your ( W_{ih} ) is a 2x4 matrix — 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's torch.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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.28 15:52:45