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

如何为不同输入形状的手动剪枝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])

关键修改点

  1. 移除固定batch_size参数:初始化时不再绑定固定batch尺寸,改为根据每次输入动态适配隐藏状态的batch大小。
  2. 输入维度统一:通过unsqueeze(0)将无batch的输入转为(1, feature_dim)格式,内部计算完全基于带batch的张量,最后再用squeeze(0)还原输出维度。
  3. 隐藏状态动态适配:在forward中检查当前隐藏状态的batch尺寸是否与输入匹配,不匹配则自动重置,避免形状不兼容问题。
  4. 张量属性对齐:创建next_hidden_states时指定device和dtype,保证与输入张量的设备、数据类型一致,避免运行时错误。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.26 01:52:34