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

基于PyTorch复现论文中Recurrent Attention Unit(RAU)的技术求助

问题描述

我正尝试复现论文《Recurrent Attention Unit》中提出的Recurrent Attention Unit(RAU),该单元将注意力机制直接集成到GRU单元内部,区别于RNN中常见的外部注意力机制应用方式。

RAU单元架构(红色链接部分为注意力门,建议彩色查看)

我目前面临的核心难题是将论文中描述GRU单元修改方式的数学公式转化为PyTorch Python代码。据悉原实现为C版本,这增加了移植难度。我尝试获取相关C文件,但torch依赖库中并无该文件。


解决方案

基于论文中RAU的数学定义,以下是纯PyTorch实现的代码,核心是在标准GRU基础上新增注意力门计算,并修改最终隐藏状态的更新逻辑:

1. 单个RAU单元实现(单时间步)

import torch
import torch.nn as nn
import torch.nn.functional as F

class RAUCell(nn.Module):
    def __init__(self, input_size, hidden_size, bias=True):
        super().__init__()
        self.input_size = input_size
        self.hidden_size = hidden_size
        self.bias = bias
        
        # 标准GRU的重置门、更新门、候选状态权重
        self.weight_xr = nn.Parameter(torch.randn(input_size, hidden_size))
        self.weight_hr = nn.Parameter(torch.randn(hidden_size, hidden_size))
        self.weight_xz = nn.Parameter(torch.randn(input_size, hidden_size))
        self.weight_hz = nn.Parameter(torch.randn(hidden_size, hidden_size))
        self.weight_xh = nn.Parameter(torch.randn(input_size, hidden_size))
        self.weight_hh = nn.Parameter(torch.randn(hidden_size, hidden_size))
        
        # RAU新增的注意力门权重
        self.weight_xa = nn.Parameter(torch.randn(input_size, hidden_size))
        self.weight_ha = nn.Parameter(torch.randn(hidden_size, hidden_size))
        self.weight_aa = nn.Parameter(torch.randn(hidden_size, hidden_size))
        
        if bias:
            self.bias_r = nn.Parameter(torch.randn(hidden_size))
            self.bias_z = nn.Parameter(torch.randn(hidden_size))
            self.bias_h = nn.Parameter(torch.randn(hidden_size))
            self.bias_a = nn.Parameter(torch.randn(hidden_size))
        else:
            self.register_parameter('bias_r', None)
            self.register_parameter('bias_z', None)
            self.register_parameter('bias_h', None)
            self.register_parameter('bias_a', None)
        
        self.reset_parameters()
    
    def reset_parameters(self):
        # 遵循PyTorch官方GRU的正交初始化逻辑,保证训练稳定性
        nn.init.orthogonal_(self.weight_xr)
        nn.init.orthogonal_(self.weight_hr)
        nn.init.orthogonal_(self.weight_xz)
        nn.init.orthogonal_(self.weight_hz)
        nn.init.orthogonal_(self.weight_xh)
        nn.init.orthogonal_(self.weight_hh)
        nn.init.orthogonal_(self.weight_xa)
        nn.init.orthogonal_(self.weight_ha)
        nn.init.orthogonal_(self.weight_aa)
        
        if self.bias:
            nn.init.zeros_(self.bias_r)
            nn.init.zeros_(self.bias_z)
            nn.init.zeros_(self.bias_h)
            nn.init.zeros_(self.bias_a)
    
    def forward(self, x, h_prev):
        # 计算重置门 r_t
        r = torch.sigmoid(torch.matmul(x, self.weight_xr) + torch.matmul(h_prev, self.weight_hr) + (self.bias_r if self.bias else 0))
        # 计算更新门 z_t
        z = torch.sigmoid(torch.matmul(x, self.weight_xz) + torch.matmul(h_prev, self.weight_hz) + (self.bias_z if self.bias else 0))
        # 计算候选隐藏状态 \tilde{h}_t
        h_tilde = torch.tanh(torch.matmul(x, self.weight_xh) + r * torch.matmul(h_prev, self.weight_hh) + (self.bias_h if self.bias else 0))
        # 计算注意力门 a_t(RAU核心新增模块)
        a = torch.sigmoid(torch.matmul(x, self.weight_xa) + torch.matmul(h_prev, self.weight_ha) + torch.matmul(h_tilde, self.weight_aa) + (self.bias_a if self.bias else 0))
        # 计算最终隐藏状态 h_t:用注意力门加权候选状态
        h = (1 - z) * h_prev + z * a * h_tilde
        
        return h

2. 封装为序列层(支持批量多时间步输入)

如果需要处理完整的序列输入,可以基于RAUCell封装成类似nn.GRU的序列层:

class RAU(nn.Module):
    def __init__(self, input_size, hidden_size, num_layers=1, bias=True, batch_first=False):
        super().__init__()
        self.input_size = input_size
        self.hidden_size = hidden_size
        self.num_layers = num_layers
        self.batch_first = batch_first
        
        self.cells = nn.ModuleList([RAUCell(input_size if i == 0 else hidden_size, hidden_size, bias) for i in range(num_layers)])
    
    def forward(self, x, h_prev=None):
        if self.batch_first:
            x = x.transpose(0, 1)  # 转为 (seq_len, batch_size, input_size) 格式
        
        seq_len, batch_size, _ = x.size()
        
        if h_prev is None:
            h_prev = torch.zeros(self.num_layers, batch_size, self.hidden_size, device=x.device)
        
        outputs = []
        current_h = h_prev
        
        for t in range(seq_len):
            x_t = x[t]
            new_h = []
            for i, cell in enumerate(self.cells):
                layer_h = current_h[i]
                layer_h_next = cell(x_t, layer_h)
                new_h.append(layer_h_next)
                x_t = layer_h_next  # 上层输出作为下层输入
            current_h = torch.stack(new_h)
            outputs.append(current_h[-1])  # 仅记录最后一层的输出
        
        outputs = torch.stack(outputs)
        if self.batch_first:
            outputs = outputs.transpose(0, 1)  # 转回 (batch_size, seq_len, hidden_size) 格式
        
        return outputs, current_h

验证使用示例

# 测试RAU层
input_size = 16
hidden_size = 32
batch_size = 8
seq_len = 10

model = RAU(input_size, hidden_size, batch_first=True)
x = torch.randn(batch_size, seq_len, input_size)
output, h_n = model(x)

print(f"Output shape: {output.shape}")  # 预期输出 (8, 10, 32)
print(f"Final hidden state shape: {h_n.shape}")  # 预期输出 (1, 8, 32)

关键说明

  • 注意力门a_t融合了当前输入x_t、上一步隐藏状态h_{t-1}和候选隐藏状态\tilde{h}_t,是RAU区别于标准GRU的核心设计
  • 最终隐藏状态通过a_t对候选状态加权,实现了GRU内部的注意力机制
  • 初始化逻辑对齐PyTorch官方GRU,避免训练初期的数值不稳定问题

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.26 19:52:47