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

如何将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/(输入维度+输出维度))]
  • 延迟构建逻辑:实现了和Kerasbuild方法一致的运行时动态初始化权重的逻辑,不需要在实例化层时提前传入嵌入维度参数
  • 前向逻辑修正:
    • 替换仅支持二维运算的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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.06 12:15:04