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

我目前面临的核心难题是将论文中描述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
相关产品推荐
相关产品推荐

