yum-archive/TaSTT-Whisper
High-performance GPGPU inference of OpenAI's Whisper automatic speech recognition (ASR) model
git clone https://git.yummers.dev/yum-archive/TaSTT-Whisper
8c4603c
master
1// Dispatch with [ neq1*neq2*neq3, 1, 1 ] thread groups 2#include "flashAttentionCommon.hlsli" 3#include "groupReduce.hlsli" 4 5inline void computeDotProduct ( Buffer < float > buff0 , Buffer < float > buff1 , uint s0 , uint s1 , const uint len , const uint thread , inout float acc ) 6{ 7acc = 0 ; 8/* 9const uint s0End = s0 + len; 10s0 += thread; 11s1 += thread; 12for( ; s0 < s0End; s0 += 32, s1 += 32 ) 13acc = mad( buff0[ s0 ], buff1[ s1 ], acc ); 14horizontalSumCompatNew( thread, acc ); 15*/ 16 17const uint completeVectors = len / 32 ; 18uint i ; 19for ( i = 0 ; i < completeVectors ; i ++ , s0 += 32 , s1 += 32 ) 20acc = mad ( buff0 [ s0 + thread ], buff1 [ s1 + thread ], acc ); 21 22horizontalSumCompatNew ( thread , acc ); 23 24if ( 0 == thread ) 25{ 26const uint rem = len % 32 ; 27for ( i = 0 ; i < rem ; i ++ ) 28{ 29 precisefloat a = buff0 [ s0 + i ]; 30 precisefloat b = buff1 [ s1 + i ]; 31 precisefloat prod = a * b ; 32acc += prod ; 33} 34} 35} 36 37void scaleTempVector ( uint i , const uint length , const uint thread , const float multiplier ) 38{ 39const uint end = i + length ; 40for ( i += thread ; i < end ; i += 32 ) 41{ 42float f = temp [ i ]; 43f *= multiplier ; 44temp [ i ] = f ; 45} 46} 47 48[ numthreads ( 32 , 1 , 1 ) ] 49void main ( uint3 group : SV_GroupID , uint thread : SV_GroupIndex ) 50{ 51const uint neq0 = q_elements [ 0 ]; 52const uint neq1 = q_elements [ 1 ]; 53const uint neq2 = q_elements [ 2 ]; 54const uint neq3 = q_elements [ 3 ]; 55 56const uint nek0 = k_elements [ 0 ]; 57const uint nek1 = k_elements [ 1 ]; 58 59const uint nev1 = v_elements [ 1 ]; 60 61const uint ne0 = res_elements [ 0 ]; 62const uint ne1 = res_elements [ 1 ]; 63 64const uint nbk0 = k_strides [ 0 ]; 65const uint nbk1 = k_strides [ 1 ]; 66const uint nbk2 = k_strides [ 2 ]; 67const uint nbk3 = k_strides [ 3 ]; 68 69const uint nbq0 = q_strides [ 0 ]; 70const uint nbq1 = q_strides [ 1 ]; 71const uint nbq2 = q_strides [ 2 ]; 72const uint nbq3 = q_strides [ 3 ]; 73 74const uint nbv0 = v_strides [ 0 ]; 75const uint nbv1 = v_strides [ 1 ]; 76const uint nbv2 = v_strides [ 2 ]; 77const uint nbv3 = v_strides [ 3 ]; 78 79const uint nb0 = res_strides [ 0 ]; 80const uint nb1 = res_strides [ 1 ]; 81const uint nb2 = res_strides [ 2 ]; 82const uint nb3 = res_strides [ 3 ]; 83 84const uint D = neq0 ; 85const uint N = neq1 ; 86const uint P = nek1 - N ; 87// const uint M = P + N; 88const uint M = nek1 ; 89 90const uint ir = group . x ; 91const uint iq3 = ir / ( neq2 * neq1 ); 92const uint iq2 = ( ir - iq3 * neq2 * neq1 ) / neq1 ; 93const uint iq1 = ( ir - iq3 * neq2 * neq1 - iq2 * neq1 ); 94 95const uint tempIndex = ir * tempBufferStride ; 96 97uint ic ; 98for ( ic = 0 ; ic < nek1 ; ic ++ ) 99{ 100// k indices 101const uint ik3 = iq3 ; 102const uint ik2 = iq2 ; 103const uint ik1 = ic ; 104 105// S indices 106const uint i1 = ik1 ; 107 108if ( masked ) 109{ 110if ( ic > P + iq1 ) 111{ 112if ( 0 == thread ) 113temp [ tempIndex + ic ] = negativeInfinity ; 114continue ; 115} 116} 117 118const uint s0 = ik1 * nbk1 + ik2 * nbk2 + ik3 * nbk3 ; 119const uint s1 = iq1 * nbq1 + iq2 * nbq2 + iq3 * nbq3 ; 120float dp ; 121computeDotProduct ( k , q , s0 , s1 , neq0 , thread , dp ); 122if ( 0 == thread ) 123temp [ tempIndex + ic ] = dp * scale ; 124} 125}