GQA分组注意力机制

一、目录

  1. 定义
  2. demo

二、实现

  1. 定义
    grouped query attention(GQA)
    1 GQA 原理与优点:将query 进行分组,每组query 参数共享一份key,value, 从而使key, value 矩阵变小。
    2. 优点: 降低内存读取模型权重的时间开销:由于Key矩阵和Value矩阵数量变少了,因此权重参数量也减少了,需要读取到内存的数量量少了,因此减少了读取权重的等待时间。
    3. 效果(并未降低模型性能):GQA通过设置合适的分组大小,可以和MQA的推理性能几乎相等,同时逼近MHA的模型性能。

  2. llama3 分组数为4, chatglm2 分组数为2 .


    参考:https://zhuanlan.zhihu.com/p/693928854
    demo

    import torch
    import torch.nn as nn
    import math

    #GQA
    bs=3
    seq_len =5
    hidden_size= 32
    n_heads=4
    n_kv_heads = 2
    head_dim = hidden_size//n_heads #
    groups = n_heads//n_kv_heads # 4/2
    print("groups=",groups)
    x=torch.randn((bs,seq_len,hidden_size))
    print("x:", x.shape)
    wq = nn.Linear(hidden_size,n_heads*head_dim,bias=False)
    wk = nn.Linear(hidden_size, n_kv_heads * head_dim, bias=False)
    wv = nn.Linear(hidden_size, n_kv_heads * head_dim, bias=False)
    xq,xk,xv=wq(x),wk(x),wv(x)
    xq = xq.view(bs,seq_len, n_heads, head_dim).transpose(1, 2)
    xk = xk.view(bs,seq_len, n_kv_heads, head_dim).transpose(1, 2)
    xv = xv.view(bs,seq_len, n_kv_heads, head_dim).transpose(1, 2)
    print("xq:",xq.shape) #[bs,n_heads,seq_len, head_dim]
    print("xk:", xk.shape)#[bs,n_kv_heads,seq_len, head_dim]
    print("xv:", xv.shape)#[bs,n_kv_heads,seq_len, head_dim]
    def repeat_kv(keys: torch.Tensor, values: torch.Tensor, repeats: int, dim: int):
    keys = torch.repeat_interleave(keys, repeats=repeats, dim=dim)
    values = torch.repeat_interleave(values, repeats=repeats, dim=dim)
    return keys, values
    #复制kv head
    key,val = repeat_kv(xk,xv, groups,dim=1)
    print("key:", key.shape)
    print("val:", val.shape)
    attn_weights = torch.matmul(xq, key.transpose(2, 3)) / math.sqrt(head_dim)
    print("attn_weights:", attn_weights.shape) #[bs,n_heads,seq_len,seq_len]
    attn_output = torch.matmul(attn_weights, val)
    print("attn_output:", attn_output.shape) # [bs,n_heads,seq_len,head_dim]

相关推荐
羊小猪~~31 分钟前
神经网络基础--什么是正向传播??什么是方向传播??
人工智能·pytorch·python·深度学习·神经网络·算法·机器学习
软工菜鸡1 小时前
预训练语言模型BERT——PaddleNLP中的预训练模型
大数据·人工智能·深度学习·算法·语言模型·自然语言处理·bert
哔哩哔哩技术2 小时前
B站S赛直播中的关键事件识别与应用
深度学习
deephub2 小时前
Tokenformer:基于参数标记化的高效可扩展Transformer架构
人工智能·python·深度学习·架构·transformer
___Dream2 小时前
【CTFN】基于耦合翻译融合网络的多模态情感分析的层次学习
人工智能·深度学习·机器学习·transformer·人机交互
极客代码2 小时前
【Python TensorFlow】入门到精通
开发语言·人工智能·python·深度学习·tensorflow
王哈哈^_^4 小时前
【数据集】【YOLO】【VOC】目标检测数据集,查找数据集,yolo目标检测算法详细实战训练步骤!
人工智能·深度学习·算法·yolo·目标检测·计算机视觉·pyqt
写代码的小阿帆4 小时前
pytorch实现深度神经网络DNN与卷积神经网络CNN
pytorch·cnn·dnn
是瑶瑶子啦4 小时前
【深度学习】论文笔记:空间变换网络(Spatial Transformer Networks)
论文阅读·人工智能·深度学习·视觉检测·空间变换
wangyue45 小时前
c# 深度模型入门
深度学习