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

KonstantinSource codes8c4603c

master
1.5 KiB45 linesraw
1// Implementation of fmaRepeat() when source arguments have different shape or VRAM layout
2// Dispatch [ nb[ 1 ], nb[ 2 ], nb[ 3 ] ] thread groups of this shader, where nb is size of the destination tensor
3RWBuffer<float> tensor: register( u0 );
4Buffer<float> patternMul: register( t0 );
5Buffer<float> patternAdd: register( t1 );
6
7cbuffer Constants: register( b0 )
8{
9	uint4 tensorSize: packoffset( c0 );
10	uint4 tensorStrides: packoffset( c1 );
11	uint4 patternSizeMul: packoffset( c2 );
12	uint4 patternStridesMul: packoffset( c3 );
13	uint4 patternSizeAdd: packoffset( c4 );
14	uint4 patternStridesAdd: packoffset( c5 );
15}
16
17#ifndef THREADS
18#define THREADS 32
19#endif
20
21#include "repeatUtils.hlsli"
22
23inline float loadPattern( Buffer<float> buffer, uint rowStart, uint i, uint4 size, uint4 stride )
24{
25	i %= size.x;
26	return buffer[ i * stride.x + rowStart ];
27}
28
29[ numthreads( THREADS, 1, 1 ) ]
30void main( uint3 group: SV_GroupID, uint thread : SV_GroupIndex )
31{
32	uint3 it = tensorIteratorState( group, thread, tensorSize, tensorStrides );
33	const uint rsiMul = rowOffset( group % patternSizeMul.yzw, patternStridesMul );
34	const uint rsiAdd = rowOffset( group % patternSizeAdd.yzw, patternStridesAdd );
35
36	for( uint i = thread; it.x < it.z; it.x += it.y, i++ )
37	{
38		precise float f = tensor[ it.x ];
39		float mul = loadPattern( patternMul, rsiMul, i, patternSizeMul, patternStridesMul );
40		float add = loadPattern( patternAdd, rsiAdd, i, patternSizeAdd, patternStridesAdd );
41		f *= mul;
42		f += add;
43		tensor[ it.x ] = f;
44	}
45}