如何将Keras自定义Attention层等效迁移至PyTorch实现
完整对齐Keras逻辑的PyTorch实现
import torch import torch.nn.functional as F class Attention_module(torch.nn.Module): def __init__(self, class_num): super().__init__() self.class_num = class_num # 占位权重,第一次前向时根据输入维度完成初始化,对齐Keras的build逻辑 self.Ws = torch.nn.Parameter(torch.empty(class_num, 0)) self._built = False def forward(self, inputs): # inputs维度: [batch_size, 序列长度, 嵌入维度] _, _, embedding_length = inputs.shape # 首次前向时初始化权重 if not self._built: self.Ws.data = torch.empty(self.class_num, embedding_length) # *Keras的glorot_uniform初始化等价于PyTorch的xavier_uniform_* torch.nn.init.xavier_uniform_(self.Ws) self._built = True sentence_trans = inputs.permute(0, 2, 1) # 带batch维度的批量矩阵乘法使用matmul,原代码的mm仅支持二维矩阵运算 at = torch.matmul(self.Ws, sentence_trans) at = torch.tanh(at) at = torch.exp(at - torch.max(at, dim=-1, keepdims=True).values) at = at / torch.sum(at, dim=-1, keepdims=True) v = torch.einsum('ijk,ikl->ijl', at, inputs) return v
核心对齐说明
- 权重初始化:Keras的
glorot_uniform和PyTorch的xavier_uniform_实现逻辑完全一致,采样区间均为[-sqrt(6/(输入维度+输出维度)), sqrt(6/(输入维度+输出维度))] - 延迟构建逻辑:实现了和Keras
build方法一致的运行时动态初始化权重的逻辑,不需要在实例化层时提前传入嵌入维度参数 - 前向逻辑修正:
- 替换仅支持二维运算的
torch.mm为支持批量运算的torch.matmul,适配带batch维度的输入 - 替换模块类
torch.nn.Tanh为函数式接口torch.tanh,避免实例化开销 - 移除冗余的
torch.Tensor(at)类型转换,输入本身已是张量类型
- 替换仅支持二维运算的
可选简化版本
如果不需要对齐Keras的延迟构建逻辑,也可以在实例化层时提前传入嵌入维度,代码更简洁:
class Attention_module(torch.nn.Module): def __init__(self, class_num, embedding_length): super().__init__() self.Ws = torch.nn.Parameter(torch.empty(class_num, embedding_length)) torch.nn.init.xavier_uniform_(self.Ws) def forward(self, inputs): sentence_trans = inputs.permute(0, 2, 1) at = torch.matmul(self.Ws, sentence_trans) at = torch.tanh(at) at = torch.exp(at - torch.max(at, dim=-1, keepdims=True).values) at = at / torch.sum(at, dim=-1, keepdims=True) v = torch.einsum('ijk,ikl->ijl', at, inputs) return v
内容的提问来源于stack exchange,提问作者Aaditya Ura
相关产品推荐
相关产品推荐

