From self-attention 2 flash-attention 数学原理与 cuda 实现优化

2024-06-09 07:44

本文主要是介绍From self-attention 2 flash-attention 数学原理与 cuda 实现优化,希望对大家解决编程问题提供一定的参考价值,需要的开发者们随着小编来一起学习吧!

self attension 是transformer 编码器和解码器中共同的一个计算环节,在整个transformer 网络体系中耗费的算力比例占主导。所以节省self attention 的正向和反向的计算时间,就可以加速 transormer 的训练和推理过程。

1,self attention 的数学提炼

两个矩阵乘法,加入一个列向的softmax

input   矩阵: \mathbf{Q}, \mathbf{K}, \mathbf{V} \in \mathbf{R}^{N \times d}

output 矩阵:\mathbf{O} \in \mathbf{R}^{N \times d}

 

\mathbf{self\ attention\ algorithm:}

        step1:        \mathbf{S} = \mathbf{Q}*\mathbf{K}^t

        step2:        \mathbf{P} = \mathbf{softmax_{column}(S)}

        step3:        \mathbf{O} = \mathbf{P}*\mathbf{V}

2,cpu 实现self attention

这里的数据类型使用了 float,实际网络中一般采用 fp16,数学过程是相同的;

cpu_self_attention.cpp

#include <stdio.h>
#include <string.h>#include "cpu_gemm.h"
#include "utils.h"
#include "soft_max.h"
//all matrices are row major.void cpu_self_attention(float* Q, int ldq,float* K, int ldk,float* V, int ldv,float* S, int lds,float* P, int ldp,float* O, int ldo,int N, int d)
{gemm_nt(Q, ldq, K, ldk, S, lds, N, N, d);// S = Q*K^t     (NxN) = (Nxd) * (dxN)printf("\nS =\n");	print_matrix(S, N, N, lds);soft_max_column(P, ldp, S, lds, N, N);// P(NxN) = softmax(S(NxN))printf("\nP =\n");	print_matrix(S, N, N, lds);gemm_nn(P, ldp, V, ldv, O, ldo, N, d, N);// O = P*V     (Nxd) = (NxN) * (Nxd)
}

cpu_gemm.cpp

#include "cpu_gemm.h"void gemm_nn(float *A, int lda,		//A(M x K) rowMjfloat *B, int ldb,		//B(K x N) rowMjfloat *C, int ldc,		//C(M x N) rowMjint M,int N,int K)
{for(int i=0; i<M; i++){for(int j=0; j<N; j++){float sigma = 0.0;for(int k=0; k<K; k++){sigma += A[i*lda + k] * B[k*ldb + j];}C[i*ldc + j] = sigma;}}
}void gemm_nt(float *A, int lda,		//A(M x K) rowMjfloat *B, int ldb,		//B(N x K) rowMjfloat *C, int ldc,		//C(M x N) rowMjint M,int N,int K)
{for(int i=0; i<M; i++){for(int j=0; j<N; j++){float sigma = 0.0;for(int k=0; k<K; k++){sigma += A[i*lda + k] * B[k + j*ldb];}C[i*ldc + j] = sigma;}}
}

cpu_softmax_column.cpp

这里使用的是未数值优化的方式,直接按照原始公式计算:

#include "soft_max.h"
void soft_max_column(float *P, int ldp, float* S, int lds, int M, int N)//P = softmax(S)  P(i,j) = exp(S(i,j))/sigma(exp(S(r,j)));  r=0,1,..,n-1 ;
{for(int j=0; j<N; j++){float sigma = 0.0f;for(int i=0; i<M; i++){sigma += exp(S[i*lds + j])}for(int i=0; i<M; i++){P[i*ldp + j] = S[i*lds + j]/sigma;}}
}

3, gpu 实现 self attention 正向

cuda 实现上述过程:

gpu_self_attention.cu

gpu_gemm.cu

gpu_softmax_column.cu

4,为什么不需要gpu 实现self attention 反向

融合上述过程

5, gpu 实现 flash attention 反向

融合算子

数学原理

cuda 实现

挖坑,未完待续 。。。

这篇关于From self-attention 2 flash-attention 数学原理与 cuda 实现优化的文章就介绍到这儿,希望我们推荐的文章对编程师们有所帮助!



http://www.chinasem.cn/article/1044578

相关文章

C#使用HttpClient进行Post请求出现超时问题的解决及优化

《C#使用HttpClient进行Post请求出现超时问题的解决及优化》最近我的控制台程序发现有时候总是出现请求超时等问题,通常好几分钟最多只有3-4个请求,在使用apipost发现并发10个5分钟也... 目录优化结论单例HttpClient连接池耗尽和并发并发异步最终优化后优化结论我直接上优化结论吧,

windos server2022里的DFS配置的实现

《windosserver2022里的DFS配置的实现》DFS是WindowsServer操作系统提供的一种功能,用于在多台服务器上集中管理共享文件夹和文件的分布式存储解决方案,本文就来介绍一下wi... 目录什么是DFS?优势:应用场景:DFS配置步骤什么是DFS?DFS指的是分布式文件系统(Distr

NFS实现多服务器文件的共享的方法步骤

《NFS实现多服务器文件的共享的方法步骤》NFS允许网络中的计算机之间共享资源,客户端可以透明地读写远端NFS服务器上的文件,本文就来介绍一下NFS实现多服务器文件的共享的方法步骤,感兴趣的可以了解一... 目录一、简介二、部署1、准备1、服务端和客户端:安装nfs-utils2、服务端:创建共享目录3、服

Java内存泄漏问题的排查、优化与最佳实践

《Java内存泄漏问题的排查、优化与最佳实践》在Java开发中,内存泄漏是一个常见且令人头疼的问题,内存泄漏指的是程序在运行过程中,已经不再使用的对象没有被及时释放,从而导致内存占用不断增加,最终... 目录引言1. 什么是内存泄漏?常见的内存泄漏情况2. 如何排查 Java 中的内存泄漏?2.1 使用 J

C#使用yield关键字实现提升迭代性能与效率

《C#使用yield关键字实现提升迭代性能与效率》yield关键字在C#中简化了数据迭代的方式,实现了按需生成数据,自动维护迭代状态,本文主要来聊聊如何使用yield关键字实现提升迭代性能与效率,感兴... 目录前言传统迭代和yield迭代方式对比yield延迟加载按需获取数据yield break显式示迭

Python实现高效地读写大型文件

《Python实现高效地读写大型文件》Python如何读写的是大型文件,有没有什么方法来提高效率呢,这篇文章就来和大家聊聊如何在Python中高效地读写大型文件,需要的可以了解下... 目录一、逐行读取大型文件二、分块读取大型文件三、使用 mmap 模块进行内存映射文件操作(适用于大文件)四、使用 pand

python实现pdf转word和excel的示例代码

《python实现pdf转word和excel的示例代码》本文主要介绍了python实现pdf转word和excel的示例代码,文中通过示例代码介绍的非常详细,对大家的学习或者工作具有一定的参考学习价... 目录一、引言二、python编程1,PDF转Word2,PDF转Excel三、前端页面效果展示总结一

Python xmltodict实现简化XML数据处理

《Pythonxmltodict实现简化XML数据处理》Python社区为提供了xmltodict库,它专为简化XML与Python数据结构的转换而设计,本文主要来为大家介绍一下如何使用xmltod... 目录一、引言二、XMLtodict介绍设计理念适用场景三、功能参数与属性1、parse函数2、unpa

C#实现获得某个枚举的所有名称

《C#实现获得某个枚举的所有名称》这篇文章主要为大家详细介绍了C#如何实现获得某个枚举的所有名称,文中的示例代码讲解详细,具有一定的借鉴价值,有需要的小伙伴可以参考一下... C#中获得某个枚举的所有名称using System;using System.Collections.Generic;usi

Go语言实现将中文转化为拼音功能

《Go语言实现将中文转化为拼音功能》这篇文章主要为大家详细介绍了Go语言中如何实现将中文转化为拼音功能,文中的示例代码讲解详细,感兴趣的小伙伴可以跟随小编一起学习一下... 有这么一个需求:新用户入职 创建一系列账号比较麻烦,打算通过接口传入姓名进行初始化。想把姓名转化成拼音。因为有些账号即需要中文也需要英