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
9cfdfc7
master
1// An equivalent of "addRepeat.hlsl" followed by "addInPlace.hlsl". 2// Merging into a single shader saves some global memory bandwidth and reduces CPU overhead wasted binding resources and dispatching shaders 3RWBuffer < float > tensor :register ( u0 ); 4Buffer < float > pattern :register ( t0 ); 5Buffer < float > finalAdd :register ( t1 ); 6 7cbuffer Constants :register ( b0 ) 8{ 9uint4 tensorSize :packoffset ( c0 ); 10uint4 tensorStrides :packoffset ( c1 ); 11uint4 patternSize :packoffset ( c2 ); 12uint4 patternStrides :packoffset ( c3 ); 13// uint4 finalSize: packoffset( c4 ); 14uint4 finalStrides :packoffset ( c5 ); 15} 16 17#ifndef THREADS 18#define THREADS 256 19#endif 20 21#include "repeatUtils.hlsli" 22 23// The micro-kernel of the shader, computes tensor[ rsi.x ] += pattern + finalAdd[ rsi.y ] 24inline void add2 ( uint2 rsi , float pattern ) 25{ 26float f = tensor [ rsi . x ]; 27f += pattern ; 28f += finalAdd [ rsi . y ]; 29tensor [ rsi . x ] = f ; 30} 31 32[ numthreads ( THREADS , 1 , 1 ) ] 33void main ( uint3 group : SV_GroupID , uint thread : SV_GroupIndex ) 34{ 35const uint2 stridesX = uint2 ( tensorStrides . x , finalStrides . x ); 36uint2 rsi ; 37rsi . x = rowOffset ( group , tensorStrides ); 38rsi . y = rowOffset ( group , finalStrides ); 39const uint rsiEnd = rsi . x + tensorSize . x * stridesX . x ; 40rsi += stridesX * thread ; 41 42uint pat = rowOffset ( group % patternSize . yzw , patternStrides ); 43 44if ( patternSize . x == 1 ) 45{ 46// The pattern only has 1 column, broadcasting over the row 47const uint2 rsiInc = stridesX * THREADS ; 48const float p = pattern [ pat ]; 49for ( ; rsi . x < rsiEnd ; rsi += rsiInc ) 50add2 ( rsi , p ); 51} 52else if ( patternSize . x <= THREADS ) 53{ 54// pattern size doesn't exceed thread group size, load outside of the loop 55const uint threadsPerGroup = THREADS - ( THREADS % patternSize . x ); 56if ( thread >= threadsPerGroup ) 57return ; 58 59const uint2 rsiInc = stridesX * threadsPerGroup ; 60pat += ( thread % patternSize . x ) * patternStrides . x ; 61const float p = pattern [ pat ]; 62for ( ; rsi . x < rsiEnd ; rsi += rsiInc ) 63add2 ( rsi , p ); 64} 65else 66{ 67// Pattern rows are longer than the thread group, need to stream from both buffers 68uint3 rsi3 ; 69rsi3 . xy = rsi ; 70rsi3 . z = pat + thread * patternStrides . x ; 71 72const uint3 rsiInc = uint3 ( stridesX , patternStrides . x ) * THREADS ; 73while ( rsi3 . x < rsiEnd ) 74{ 75add2 ( rsi3 . xy , pattern [ rsi3 . z ] ); 76 77rsi3 += rsiInc ; 78if ( rsi3 . z >= patternSize . x ) 79rsi3 . z -= patternSize . x ; 80} 81} 82}