如何为不同输入形状的手动剪枝RNN维护隐藏状态
问题描述
手动实现了一个由带剪枝连接的多层线性层构成的RNN,通过next_hidden_states保存t时刻的隐藏状态供t+1时刻复用,该变量尺寸为(batch_size, N)。需要模型同时支持两种输入场景:
- 带batch的输入(用于训练智能体)
- 无batch的输入(用于在环境中运行单episode)
希望复刻PyTorch原生RNN模块的隐式batch处理逻辑,不想将隐藏状态作为输入输出传递,认为这种方式不够优雅。
补充说明
以下是原始极简代码:
import numpy as np import torch import torch.nn.utils.prune as prune import torch.nn as nn class BrainRNN(nn.Module): def __init__(self, activation=torch.sigmoid, batch_size=8): super(BrainRNN, self).__init__() self.n_neurons = 3*4 self.activation = activation self.batch_size = batch_size self.reset_hidden_states() # Create the input layer self.input_layer = nn.Linear(4, 4) # Create forward hidden layers self.hidden_layers = nn.ModuleList([]) new_layer = nn.Linear(4,4) mask = np.ones((4,4))-np.eye(4) prune.custom_from_mask(new_layer, name='weight', mask=torch.tensor(mask.T)) # delete fictive connections self.hidden_layers.append(new_layer) # Create the backward weights self.recurrent_layers = nn.ModuleList([]) # recurrent_layers[i](hidden_states) = layer j>i to i new_layer = nn.Linear(self.n_neurons, 4, bias=False) # no bias for backward connection mask = np.zeros((12,4)) mask[1,0] = 1 prune.custom_from_mask(new_layer, name='weight', mask=torch.tensor(mask.T)) # delete fictive connections self.recurrent_layers.append(new_layer) # Create the output layer self.output_layer = nn.Linear(4,4) def forward(self, x): next_hidden_states = torch.empty(x.shape[0], self.n_neurons) if x.dim() > 1 else torch.empty(self.n_neurons) skips = [] # list of current states for skip connections # Input layer x = self.activation(self.input_layer(x) + self.recurrent_layers[0](self.hidden_states)) next_hidden_states[...,[0,1,2,3]] = x # Hidden layers x = self.hidden_layers[0](x) x = self.activation(x) next_hidden_states[...,[4,5,6,7]] = x # Output layer x = self.output_layer(x) # no activation nor recurrent/skip connection for the last one self.hidden_states = next_hidden_states return x def reset_hidden_states(self, hidden_states=None): if self.batch_size > 0: self.hidden_states = nn.init.normal_(torch.empty(self.n_neurons), std=1).repeat(self.batch_size,1) # same hidden states for all batches else: self.hidden_states = nn.init.normal_(torch.empty(self.n_neurons), std=1) nn = BrainRNN() nn(torch.zeros(8,4)) # works well nn(torch.zeros(4)) # shape issue at next_hidden_states[...,[0,1,2,3]] = x
该RNN包含3层,每层4个节点,隐藏层与输入层间存在循环连接,且部分连接已被剪枝。
解决方案
核心是统一输入与隐藏状态的维度处理逻辑,让模型自动适配有无batch的场景,修改后的代码如下:
import numpy as np import torch import torch.nn.utils.prune as prune import torch.nn as nn class BrainRNN(nn.Module): def __init__(self, activation=torch.sigmoid): super(BrainRNN, self).__init__() self.n_neurons = 3 * 4 self.activation = activation self.hidden_states = None # 初始化不绑定固定batch size # 输入层 self.input_layer = nn.Linear(4, 4) # 前向隐藏层 self.hidden_layers = nn.ModuleList([]) new_layer = nn.Linear(4, 4) mask = np.ones((4, 4)) - np.eye(4) prune.custom_from_mask(new_layer, name='weight', mask=torch.tensor(mask.T)) self.hidden_layers.append(new_layer) # 循环连接层 self.recurrent_layers = nn.ModuleList([]) new_layer = nn.Linear(self.n_neurons, 4, bias=False) mask = np.zeros((12, 4)) mask[1, 0] = 1 prune.custom_from_mask(new_layer, name='weight', mask=torch.tensor(mask.T)) self.recurrent_layers.append(new_layer) # 输出层 self.output_layer = nn.Linear(4, 4) def forward(self, x): # 记录原始输入维度,用于后续输出调整 original_dim = x.dim() # 统一输入为(batch_size, feature_dim)格式 if original_dim == 1: x = x.unsqueeze(0) batch_size = x.shape[0] # 初始化或适配隐藏状态维度:如果未初始化或batch尺寸不匹配则重置 if self.hidden_states is None or self.hidden_states.shape[0] != batch_size: self.reset_hidden_states(batch_size=batch_size) # 创建与输入同设备、同类型的隐藏状态张量 next_hidden_states = torch.empty(batch_size, self.n_neurons, device=x.device, dtype=x.dtype) # 输入层计算 x = self.activation(self.input_layer(x) + self.recurrent_layers[0](self.hidden_states)) next_hidden_states[:, [0,1,2,3]] = x # 隐藏层计算 x = self.hidden_layers[0](x) x = self.activation(x) next_hidden_states[:, [4,5,6,7]] = x # 输出层计算 x = self.output_layer(x) # 更新隐藏状态 self.hidden_states = next_hidden_states # 如果原始输入无batch,压缩输出维度 if original_dim == 1: x = x.squeeze(0) return x def reset_hidden_states(self, hidden_states=None, batch_size=1): if hidden_states is not None: # 确保传入的隐藏状态维度符合要求 if hidden_states.dim() == 1: hidden_states = hidden_states.unsqueeze(0) self.hidden_states = hidden_states else: # 动态生成对应batch size的隐藏状态 self.hidden_states = nn.init.normal_(torch.empty(batch_size, self.n_neurons), std=1) # 测试验证 nn = BrainRNN() # 带batch输入测试 output_batch = nn(torch.zeros(8,4)) print(f"带batch输出维度: {output_batch.shape}") # 输出: torch.Size([8, 4]) # 无batch输入测试 output_single = nn(torch.zeros(4)) print(f"无batch输出维度: {output_single.shape}") # 输出: torch.Size([4])
关键修改点
- 移除固定batch_size参数:初始化时不再绑定固定batch尺寸,改为根据每次输入动态适配隐藏状态的batch大小。
- 输入维度统一:通过
unsqueeze(0)将无batch的输入转为(1, feature_dim)格式,内部计算完全基于带batch的张量,最后再用squeeze(0)还原输出维度。 - 隐藏状态动态适配:在forward中检查当前隐藏状态的batch尺寸是否与输入匹配,不匹配则自动重置,避免形状不兼容问题。
- 张量属性对齐:创建
next_hidden_states时指定device和dtype,保证与输入张量的设备、数据类型一致,避免运行时错误。
内容的提问来源于stack exchange,提问作者samje
相关产品推荐
相关产品推荐

