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

KonstantinBugfix, addRepeatEx compute shader9cfdfc7

master
2.3 KiB82 linesraw
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{
9	uint4 tensorSize: packoffset( c0 );
10	uint4 tensorStrides: packoffset( c1 );
11	uint4 patternSize: packoffset( c2 );
12	uint4 patternStrides: packoffset( c3 );
13	// uint4 finalSize: packoffset( c4 );
14	uint4 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{
26	float f = tensor[ rsi.x ];
27	f += pattern;
28	f += finalAdd[ rsi.y ];
29	tensor[ rsi.x ] = f;
30}
31
32[ numthreads( THREADS, 1, 1 ) ]
33void main( uint3 group: SV_GroupID, uint thread : SV_GroupIndex )
34{
35	const uint2 stridesX = uint2( tensorStrides.x, finalStrides.x );
36	uint2 rsi;
37	rsi.x = rowOffset( group, tensorStrides );
38	rsi.y = rowOffset( group, finalStrides );
39	const uint rsiEnd = rsi.x + tensorSize.x * stridesX.x;
40	rsi += stridesX * thread;
41
42	uint pat = rowOffset( group % patternSize.yzw, patternStrides );
43
44	if( patternSize.x == 1 )
45	{
46		// The pattern only has 1 column, broadcasting over the row
47		const uint2 rsiInc = stridesX * THREADS;
48		const float p = pattern[ pat ];
49		for( ; rsi.x < rsiEnd; rsi += rsiInc )
50			add2( rsi, p );
51	}
52	else if( patternSize.x <= THREADS )
53	{
54		// pattern size doesn't exceed thread group size, load outside of the loop
55		const uint threadsPerGroup = THREADS - ( THREADS % patternSize.x );
56		if( thread >= threadsPerGroup )
57			return;
58
59		const uint2 rsiInc = stridesX * threadsPerGroup;
60		pat += ( thread % patternSize.x ) * patternStrides.x;
61		const float p = pattern[ pat ];
62		for( ; rsi.x < rsiEnd; rsi += rsiInc )
63			add2( rsi, p );
64	}
65	else
66	{
67		// Pattern rows are longer than the thread group, need to stream from both buffers
68		uint3 rsi3;
69		rsi3.xy = rsi;
70		rsi3.z = pat + thread * patternStrides.x;
71
72		const uint3 rsiInc = uint3( stridesX, patternStrides.x ) * THREADS;
73		while( rsi3.x < rsiEnd )
74		{
75			add2( rsi3.xy, pattern[ rsi3.z ] );
76
77			rsi3 += rsiInc;
78			if( rsi3.z >= patternSize.x )
79				rsi3.z -= patternSize.x;
80		}
81	}
82}