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
11.6 KiB362 linesraw
1#include "stdafx.h"
2#include "mulMatImpl.h"
3#include <immintrin.h>
4#include "mulMatUtils.hpp"
5using namespace CpuCompute;
6
7namespace
8{
9	constexpr size_t prefetchBytes = 96;
10	constexpr int prefetchHint = _MM_HINT_T0;
11
12	constexpr size_t maskAlign16 = ~(size_t)15;
13
14	__forceinline __m256i load( const void* rsi )
15	{
16		return _mm256_loadu_si256( ( const __m256i* )rsi );
17	}
18
19#define TRANSPOSE_8X16()                           \
20                                                   \
21	__m256i t0 = _mm256_unpacklo_epi16( r0, r1 );  \
22	__m256i t1 = _mm256_unpackhi_epi16( r0, r1 );  \
23	__m256i t2 = _mm256_unpacklo_epi16( r2, r3 );  \
24	__m256i t3 = _mm256_unpackhi_epi16( r2, r3 );  \
25	__m256i t4 = _mm256_unpacklo_epi16( r4, r5 );  \
26	__m256i t5 = _mm256_unpackhi_epi16( r4, r5 );  \
27	__m256i t6 = _mm256_unpacklo_epi16( r6, r7 );  \
28	__m256i t7 = _mm256_unpackhi_epi16( r6, r7 );  \
29                                                   \
30	r0 = _mm256_unpacklo_epi32( t0, t2 );          \
31	r1 = _mm256_unpackhi_epi32( t0, t2 );          \
32	r2 = _mm256_unpacklo_epi32( t1, t3 );          \
33	r3 = _mm256_unpackhi_epi32( t1, t3 );          \
34	r4 = _mm256_unpacklo_epi32( t4, t6 );          \
35	r5 = _mm256_unpackhi_epi32( t4, t6 );          \
36	r6 = _mm256_unpacklo_epi32( t5, t7 );          \
37	r7 = _mm256_unpackhi_epi32( t5, t7 );          \
38                                                   \
39	t0 = _mm256_unpacklo_epi64( r0, r4 );          \
40	t1 = _mm256_unpackhi_epi64( r0, r4 );          \
41	t2 = _mm256_unpacklo_epi64( r1, r5 );          \
42	t3 = _mm256_unpackhi_epi64( r1, r5 );          \
43	t4 = _mm256_unpacklo_epi64( r2, r6 );          \
44	t5 = _mm256_unpackhi_epi64( r2, r6 );          \
45	t6 = _mm256_unpacklo_epi64( r3, r7 );          \
46	t7 = _mm256_unpackhi_epi64( r3, r7 )
47
48	__forceinline void storeLow( void* rdi, __m256i v )
49	{
50		__m128i i = _mm256_castsi256_si128( v );
51		_mm_store_si128( ( __m128i* )rdi, i );
52	}
53
54#define STORE_8X16_LOW()                     \
55	storeLow( rdi, t0 );                     \
56	storeLow( rdi + destStride, t1 );        \
57	storeLow( rdi + destStride * 2, t2 );    \
58	rdi += destStride * 8;                   \
59	storeLow( rdiMid, t3 );                  \
60	storeLow( rdiMid + destStride, t4 );     \
61	storeLow( rdiMid + destStride * 2, t5 ); \
62	rdiMid += destStride * 8;                \
63	storeLow( rdiLast, t6 );                 \
64	storeLow( rdiLast + destStride, t7 );    \
65	rdiLast += destStride * 8
66
67	__forceinline void storeHigh( void* rdi, __m256i v )
68	{
69		__m128i i = _mm256_extracti128_si256( v, 1 );
70		_mm_store_si128( ( __m128i* )rdi, i );
71	}
72
73#define STORE_8X16_HIGH()                     \
74	storeHigh( rdi, t0 );                     \
75	storeHigh( rdi + destStride, t1 );        \
76	storeHigh( rdi + destStride * 2, t2 );    \
77	rdi += destStride * 8;                    \
78	storeHigh( rdiMid, t3 );                  \
79	storeHigh( rdiMid + destStride, t4 );     \
80	storeHigh( rdiMid + destStride * 2, t5 ); \
81	rdiMid += destStride * 8;                 \
82	storeHigh( rdiLast, t6 );                 \
83	storeHigh( rdiLast + destStride, t7 );    \
84	rdiLast += destStride * 8
85
86	__forceinline void prefetch( const uint8_t* p )
87	{
88		_mm_prefetch( (const char*)p, prefetchHint );
89	}
90
91	__forceinline void transpose8Avx2( uint16_t* rdiWords, size_t w, const uint16_t* rsiWords, size_t sourceStride, size_t destStride )
92	{
93		assert( 0 == ( (size_t)rdiWords ) % 16 );
94		assert( 0 == destStride % 8 );
95		assert( w <= sourceStride );
96
97		// Scale strides to bytes, and cast the pointers
98		sourceStride *= 2;
99		destStride *= 2;
100		uint8_t* rdi = (uint8_t*)rdiWords;
101		const uint8_t* rsi = (const uint8_t*)rsiWords;
102
103		const uint8_t* const rsiEndAligned = rsi + ( w & maskAlign16 ) * 2;
104		const uint8_t* const rsiEnd = rsi + w * 2;
105		const uint8_t* rsiMid = rsi + sourceStride * 3;
106		const uint8_t* rsiLast = rsi + sourceStride * 6;
107		uint8_t* rdiMid = rdi + destStride * 3;
108		uint8_t* rdiLast = rdi + destStride * 6;
109
110		while( rsi < rsiEndAligned )
111		{
112			// Load 16x8 block into 8 registers
113			__m256i r0 = load( rsi );
114			__m256i r1 = load( rsi + sourceStride );
115			__m256i r2 = load( rsi + sourceStride * 2 );
116			rsi += 32;
117			__m256i r3 = load( rsiMid );
118			__m256i r4 = load( rsiMid + sourceStride );
119			__m256i r5 = load( rsiMid + sourceStride * 2 );
120			rsiMid += 32;
121			__m256i r6 = load( rsiLast );
122			__m256i r7 = load( rsiLast + sourceStride );
123			rsiLast += 32;
124
125			// Transpose FP16 values in registers
126			TRANSPOSE_8X16();
127
128			// Store
129			STORE_8X16_LOW();
130			STORE_8X16_HIGH();
131
132			if constexpr( prefetchBytes > 0 )
133			{
134				if( rsi + prefetchBytes < rsiEnd )
135				{
136					prefetch( rsi + prefetchBytes );
137					prefetch( rsi + sourceStride + prefetchBytes );
138					prefetch( rsi + sourceStride * 2 + prefetchBytes );
139					prefetch( rsiMid + prefetchBytes );
140					prefetch( rsiMid + sourceStride + prefetchBytes );
141					prefetch( rsiMid + sourceStride * 2 + prefetchBytes );
142					prefetch( rsiLast + prefetchBytes );
143					prefetch( rsiLast + sourceStride + prefetchBytes );
144				}
145			}
146		}
147
148		if( rsi < rsiEnd )
149		{
150			// Loading 8 elements into corresponding lanes of 8 vectors
151			// This way there's no data dependencies between these load instructions
152			// Out of order execution should hopefully do it's magic in the CPU, running all these loads in parallel.
153			__m128i r0;
154			__m128i r1 = _mm_setzero_si128();
155			__m128i r2 = _mm_setzero_si128();
156			__m128i r3 = _mm_setzero_si128();
157			__m128i r4 = _mm_setzero_si128();
158			__m128i r5 = _mm_setzero_si128();
159			__m128i r6 = _mm_setzero_si128();
160			__m128i r7 = _mm_setzero_si128();
161
162			__m128i t0, t1, t2, t3, t4, t5, t6;
163
164#pragma loop( no_vector )
165			while( rsi < rsiEnd )
166			{
167				r0 = _mm_cvtsi32_si128( *(const uint16_t*)rsi );
168				r1 = _mm_insert_epi16( r1, *(const int16_t*)( rsi + sourceStride ), 1 );
169				r2 = _mm_insert_epi16( r2, *(const int16_t*)( rsi + sourceStride * 2 ), 2 );
170				rsi += 2;
171				r3 = _mm_insert_epi16( r3, *(const int16_t*)( rsiMid ), 3 );
172				r4 = _mm_insert_epi16( r4, *(const int16_t*)( rsiMid + sourceStride ), 4 );
173				r5 = _mm_insert_epi16( r5, *(const int16_t*)( rsiMid + sourceStride * 2 ), 5 );
174				rsiMid += 2;
175				r6 = _mm_insert_epi16( r6, *(const int16_t*)( rsiLast ), 6 );
176				r7 = _mm_insert_epi16( r7, *(const int16_t*)( rsiLast + sourceStride ), 7 );
177				rsiLast += 2;
178
179				// Bitwise operations are pretty fast, AMD Zen3 CPU can run 4 of them every clock cycle
180				// Combine 8 vectors into one
181				t0 = _mm_or_si128( r0, r1 );
182				t1 = _mm_or_si128( r2, r3 );
183				t2 = _mm_or_si128( r4, r5 );
184				t3 = _mm_or_si128( r6, r7 );
185
186				t4 = _mm_or_si128( t0, t1 );
187				t5 = _mm_or_si128( t2, t3 );
188
189				t6 = _mm_or_si128( t4, t5 );
190				// Store 8 FP16 values, the destination is aligned
191				_mm_store_si128( ( __m128i* )rdi, t6 );
192				rdi += destStride;
193			}
194		}
195	}
196
197	__forceinline void transpose8PartialAvx2( uint16_t* rdiWords, size_t w, size_t h, const uint16_t* rsiWords, size_t sourceStride, size_t destStride )
198	{
199		assert( 0 == ( (size_t)rdiWords ) % 16 );
200		assert( 0 == destStride % 8 );
201		assert( w <= sourceStride );
202		assert( h > 0 && h < 8 );
203
204		// Scale strides to bytes, and cast the pointers
205		sourceStride *= 2;
206		destStride *= 2;
207		uint8_t* rdi = (uint8_t*)rdiWords;
208		const uint8_t* rsi = (const uint8_t*)rsiWords;
209
210		const uint8_t* const rsiEndAligned = rsi + ( w & maskAlign16 ) * 2;
211		const uint8_t* const rsiEnd = rsi + w * 2;
212		const uint8_t* rsiMid = rsi + sourceStride * 3;
213		const uint8_t* rsiLast = rsi + sourceStride * 6;
214		uint8_t* rdiMid = rdi + destStride * 3;
215		uint8_t* rdiLast = rdi + destStride * 6;
216
217		while( rsi < rsiEndAligned )
218		{
219			// Load the block into 8 registers, set unused rows to zero
220			__m256i r0 = load( rsi );
221			__m256i r1 = _mm256_setzero_si256();
222			__m256i r2 = _mm256_setzero_si256();
223			__m256i r3 = _mm256_setzero_si256();
224			__m256i r4 = _mm256_setzero_si256();
225			__m256i r5 = _mm256_setzero_si256();
226			__m256i r6 = _mm256_setzero_si256();
227			// These branches, whether direct or indirect, are very predictable: same outcome for all iterations of the outer loop
228			switch( h )
229			{
230			case 7:
231				r6 = load( rsiLast );
232			case 6:
233				r5 = load( rsiMid + sourceStride * 2 );
234			case 5:
235				r4 = load( rsiMid + sourceStride );
236			case 4:
237				r3 = load( rsiMid );
238			case 3:
239				r2 = load( rsi + sourceStride * 2 );
240			case 2:
241				r1 = load( rsi + sourceStride );
242			}
243			rsi += 32;
244			rsiMid += 32;
245			rsiLast += 32;
246
247			__m256i r7 = _mm256_setzero_si256();
248
249			// Transpose FP16 values in registers
250			TRANSPOSE_8X16();
251
252			// Store
253			STORE_8X16_LOW();
254
255			STORE_8X16_HIGH();
256		}
257
258		if( rsi < rsiEnd )
259		{
260			// Loading 8 elements into corresponding lanes of 8 vectors
261			// This way there's no data dependencies between these load instructions
262			// Out of order execution should hopefully do it's magic in the CPU, running all these loads in parallel.
263			__m128i r0;
264			__m128i r1 = _mm_setzero_si128();
265			__m128i r2 = _mm_setzero_si128();
266			__m128i r3 = _mm_setzero_si128();
267			__m128i r4 = _mm_setzero_si128();
268			__m128i r5 = _mm_setzero_si128();
269			__m128i r6 = _mm_setzero_si128();
270
271			__m128i t0, t1, t2, t3, t4, t5;
272
273#pragma loop( no_vector )
274			while( rsi < rsiEnd )
275			{
276				r0 = _mm_cvtsi32_si128( *(const uint16_t*)rsi );
277
278				switch( h )
279				{
280				case 7:
281					r6 = _mm_insert_epi16( r6, *(const int16_t*)( rsiLast ), 6 );
282				case 6:
283					r5 = _mm_insert_epi16( r5, *(const int16_t*)( rsiMid + sourceStride * 2 ), 5 );
284				case 5:
285					r4 = _mm_insert_epi16( r4, *(const int16_t*)( rsiMid + sourceStride ), 4 );
286				case 4:
287					r3 = _mm_insert_epi16( r3, *(const int16_t*)( rsiMid ), 3 );
288				case 3:
289					r2 = _mm_insert_epi16( r2, *(const int16_t*)( rsi + sourceStride * 2 ), 2 );
290				case 2:
291					r1 = _mm_insert_epi16( r1, *(const int16_t*)( rsi + sourceStride ), 1 );
292				}
293				rsi += 2;
294				rsiMid += 2;
295				rsiLast += 2;
296
297				// Bitwise operations are pretty fast, AMD Zen3 CPU can run 4 of them every clock cycle
298				// Combine 7 vectors into one
299				t0 = _mm_or_si128( r0, r1 );
300				t1 = _mm_or_si128( r2, r3 );
301				t2 = _mm_or_si128( r4, r5 );
302
303				t3 = _mm_or_si128( t0, t1 );
304				t4 = _mm_or_si128( t2, r6 );
305
306				t5 = _mm_or_si128( t3, t4 );
307				// Store 8 FP16 values, the destination is aligned
308				_mm_store_si128( ( __m128i* )rdi, t5 );
309				rdi += destStride;
310			}
311		}
312	}
313}
314
315// At least for the hybrid decoder, this method absolutely dominates the CPU time.
316// And not due to the integer shuffles - the bottleneck is loading data from the source matrix.
317HRESULT MulMatBase::transposePanelAvx2( uint16_t* rdi, size_t i, size_t m2, size_t m3 ) const
318{
319	assert( stridesA[ 0 ] == 1 );
320
321	const size_t heightFloats = (size_t)panelHeightRegisters * 8;
322	i *= heightFloats;
323
324	const uint16_t* rsi = (const uint16_t*)pa;
325	rsi += m3 * stridesA[ 3 ];
326	rsi += m2 * stridesA[ 2 ];
327	rsi += i * stridesA[ 1 ];
328
329	const size_t resultStride = heightFloats;
330
331	if( i + heightFloats <= resultSize[ 0 ] )
332	{
333		// A complete panel
334		for( size_t i = 0; i < panelHeightRegisters; i++ )
335		{
336			transpose8Avx2( rdi, length, rsi, stridesA[ 1 ], resultStride );
337			// Advance by 8 floats in the output buffer
338			rdi += 8;
339			// Advance by 8 rows in the source matrix
340			rsi += 8 * stridesA[ 1 ];
341		}
342	}
343	else
344	{
345		// A partial panel, at the bottom of the first argument matrix
346		const size_t remainder = resultSize[ 0 ] - i;
347		assert( remainder > 0 && remainder < heightFloats );
348		zeroAlignedMemory( rdi, resultStride * length * sizeof( uint16_t ) );
349
350		const size_t completePanels = remainder / 8;
351		for( size_t i = 0; i < completePanels; i++ )
352		{
353			transpose8Avx2( rdi, length, rsi, stridesA[ 1 ], resultStride );
354			rdi += 8;
355			rsi += 8 * stridesA[ 1 ];
356		}
357		const size_t lastPanel = remainder % 8;
358		if( 0 != lastPanel )
359			transpose8PartialAvx2( rdi, length, lastPanel, rsi, stridesA[ 1 ], resultStride );
360	}
361	return S_OK;
362}