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.0 KiB301 linesraw
1#pragma once
2#include <immintrin.h>
3#include <stdint.h>
4#include <assert.h>
5
6__forceinline __m128i f16Load( const uint16_t* rsi )
7{
8	return _mm_loadu_si128( ( const __m128i* )rsi );
9}
10
11constexpr size_t maskAlign8 = ~(size_t)7;
12
13__forceinline void transpose8( uint16_t* rdi, size_t w, const uint16_t* rsi, size_t sourceStride, size_t destStride )
14{
15	assert( 0 == ( (size_t)rdi ) % 16 );
16	assert( 0 == destStride % 8 );
17	assert( w <= sourceStride );
18
19	const uint16_t* const rsiEndAligned = rsi + ( w & maskAlign8 );
20	const uint16_t* rsi5 = rsi + sourceStride * 5;
21	uint16_t* rdi5 = rdi + destStride * 5;
22	const size_t rem = w % 8;
23	for( ; rsi < rsiEndAligned; rsi += 8, rsi5 += 8, rdi += 8 * destStride, rdi5 += 8 * destStride )
24	{
25		// Load 8x8 block into 8 registers
26		__m128i r0 = f16Load( rsi );                     // 00, 01, 02, 03, 04, 05, 06, 07
27		__m128i r1 = f16Load( rsi + sourceStride );      // 10, 11, 12, 13, 14, 15, 16, 17
28		__m128i r2 = f16Load( rsi + sourceStride * 2 );  // 20, 21, 22, 23, 24, 25, 26, 27
29		__m128i r3 = f16Load( rsi5 - sourceStride * 2 ); // 30, 31, 32, 33, 34, 35, 36, 37
30		__m128i r4 = f16Load( rsi5 - sourceStride );     // 40, 41, 42, 43, 44, 45, 46, 47
31		__m128i r5 = f16Load( rsi5 );                    // 50, 51, 52, 53, 54, 55, 56, 57
32		__m128i r6 = f16Load( rsi5 + sourceStride );     // 60, 61, 62, 63, 64, 65, 66, 67
33		__m128i r7 = f16Load( rsi5 + sourceStride * 2 ); // 70, 71, 72, 73, 74, 75, 76, 77
34
35		// Transpose FP16 values in registers
36		__m128i t0 = _mm_unpacklo_epi16( r0, r1 ); // 00, 10, 01, 11, 02, 12, 03, 13
37		__m128i t1 = _mm_unpackhi_epi16( r0, r1 ); // 04, 14, 05, 15, 06, 16, 07, 17
38		__m128i t2 = _mm_unpacklo_epi16( r2, r3 ); // 20, 30, 21, 31, 22, 32, 23, 33
39		__m128i t3 = _mm_unpackhi_epi16( r2, r3 ); // 24, 34, 25, 35, 26, 36, 27, 37
40		__m128i t4 = _mm_unpacklo_epi16( r4, r5 ); // 40, 50, 41, 52, 42, 52, 43, 53
41		__m128i t5 = _mm_unpackhi_epi16( r4, r5 ); // 44, 54, 45, 55, 46, 56, 47, 57
42		__m128i t6 = _mm_unpacklo_epi16( r6, r7 ); // 60, 70, 61, 71, 62, 72, 63, 73
43		__m128i t7 = _mm_unpackhi_epi16( r6, r7 ); // 64, 74, 65, 75, 66, 76, 67, 77
44
45		r0 = _mm_unpacklo_epi32( t0, t2 ); // 00, 10, 20, 30, 01, 11, 21, 31
46		r1 = _mm_unpackhi_epi32( t0, t2 ); // 02, 12, 22, 32, 03, 13, 23, 33
47		r2 = _mm_unpacklo_epi32( t1, t3 ); // 04, 14, 24, 34, 05, 15, 25, 35
48		r3 = _mm_unpackhi_epi32( t1, t3 ); // 06, 16, 26, 36, 07, 17, 27, 37
49		r4 = _mm_unpacklo_epi32( t4, t6 ); // 40, 50, 60, 70, 41, 51, 61, 71
50		r5 = _mm_unpackhi_epi32( t4, t6 ); // 42, 52, 62, 72, 43, 53, 63, 73
51		r6 = _mm_unpacklo_epi32( t5, t7 ); // 44, 54, 64, 74, 45, 55, 65, 75
52		r7 = _mm_unpackhi_epi32( t5, t7 ); // 46, 56, 66, 76, 47, 57, 67, 77
53
54		t0 = _mm_unpacklo_epi64( r0, r4 ); // 00, 10, 20, 30, 40, 50, 60, 70
55		t1 = _mm_unpackhi_epi64( r0, r4 ); // 01, 11, 21, 31, 41, 52, 61, 71
56		t2 = _mm_unpacklo_epi64( r1, r5 ); // 02, 12, 22, 32, 42, 52, 62, 72
57		t3 = _mm_unpackhi_epi64( r1, r5 ); // 03, 13, 23, 33, 43, 53, 63, 73
58		t4 = _mm_unpacklo_epi64( r2, r6 );
59		t5 = _mm_unpackhi_epi64( r2, r6 );
60		t6 = _mm_unpacklo_epi64( r3, r7 );
61		t7 = _mm_unpackhi_epi64( r3, r7 );
62
63		// Store
64		store16( rdi, t0 );
65		store16( rdi + destStride, t1 );
66		store16( rdi + destStride * 2, t2 );
67		store16( rdi5 - destStride * 2, t3 );
68		store16( rdi5 - destStride, t4 );
69		store16( rdi5, t5 );
70		store16( rdi5 + destStride, t6 );
71		store16( rdi5 + destStride * 2, t7 );
72	}
73
74#pragma loop( no_vector )
75	for( size_t i = 0; i < rem; rsi++, rsi5++, rdi += destStride )
76	{
77		const int16_t* p0 = (const int16_t*)rsi;
78		const int16_t* p5 = (const int16_t*)rsi5;
79		// Load a complete column into a vector
80		__m128i v = _mm_cvtsi32_si128( *rsi );
81		v = _mm_insert_epi16( v, *( p0 + sourceStride ), 1 );
82		v = _mm_insert_epi16( v, *( p0 + sourceStride * 2 ), 2 );
83		v = _mm_insert_epi16( v, *( p5 - sourceStride * 2 ), 3 );
84		v = _mm_insert_epi16( v, *( p5 - sourceStride ), 4 );
85		v = _mm_insert_epi16( v, *( p5 ), 5 );
86		v = _mm_insert_epi16( v, *( p5 + sourceStride ), 6 );
87		v = _mm_insert_epi16( v, *( p5 + sourceStride * 2 ), 7 );
88		// Store 8 FP16 values
89		store16( rdi, v );
90	}
91}
92
93inline void transpose8Partial( uint16_t* rdi, size_t w, size_t h, const uint16_t* rsi, size_t sourceStride, size_t destStride )
94{
95	assert( 0 == ( (size_t)rdi ) % 16 );
96	assert( 0 == destStride % 8 );
97	assert( w <= sourceStride );
98	assert( h > 0 && h < 8 );
99
100	const uint16_t* const rsiEndAligned = rsi + ( w & maskAlign8 );
101	const uint16_t* rsi5 = rsi + sourceStride * 5;
102	uint16_t* rdi5 = rdi + destStride * 5;
103	const size_t rem = w % 8;
104	for( ; rsi < rsiEndAligned; rsi += 8, rsi5 += 8, rdi += 8 * destStride, rdi5 += 8 * destStride )
105	{
106		// Load the block into 8 registers, set unused rows to zero
107		__m128i r0 = f16Load( rsi );
108		__m128i r1 = _mm_setzero_si128();
109		__m128i r2 = _mm_setzero_si128();
110		__m128i r3 = _mm_setzero_si128();
111		__m128i r4 = _mm_setzero_si128();
112		__m128i r5 = _mm_setzero_si128();
113		__m128i r6 = _mm_setzero_si128();
114		// These branches, whether direct or indirect, are very predictable: same outcome for all iterations of the outer loop
115		switch( h )
116		{
117		case 7:
118			r6 = f16Load( rsi5 + sourceStride );
119		case 6:
120			r5 = f16Load( rsi5 );
121		case 5:
122			r4 = f16Load( rsi5 - sourceStride );
123		case 4:
124			r3 = f16Load( rsi5 - sourceStride * 2 );
125		case 3:
126			r2 = f16Load( rsi + sourceStride * 2 );
127		case 2:
128			r1 = f16Load( rsi + sourceStride );
129		}
130		__m128i r7 = _mm_setzero_si128();
131
132		// Transpose FP16 values in registers
133		__m128i t0 = _mm_unpacklo_epi16( r0, r1 ); // 00, 10, 01, 11, 02, 12, 03, 13
134		__m128i t1 = _mm_unpackhi_epi16( r0, r1 ); // 04, 14, 05, 15, 06, 16, 07, 17
135		__m128i t2 = _mm_unpacklo_epi16( r2, r3 ); // 20, 30, 21, 31, 22, 32, 23, 33
136		__m128i t3 = _mm_unpackhi_epi16( r2, r3 ); // 24, 34, 25, 35, 26, 36, 27, 37
137		__m128i t4 = _mm_unpacklo_epi16( r4, r5 ); // 40, 50, 41, 52, 42, 52, 43, 53
138		__m128i t5 = _mm_unpackhi_epi16( r4, r5 ); // 44, 54, 45, 55, 46, 56, 47, 57
139		__m128i t6 = _mm_unpacklo_epi16( r6, r7 ); // 60, 70, 61, 71, 62, 72, 63, 73
140		__m128i t7 = _mm_unpackhi_epi16( r6, r7 ); // 64, 74, 65, 75, 66, 76, 67, 77
141
142		r0 = _mm_unpacklo_epi32( t0, t2 ); // 00, 10, 20, 30, 01, 11, 21, 31
143		r1 = _mm_unpackhi_epi32( t0, t2 ); // 02, 12, 22, 32, 03, 13, 23, 33
144		r2 = _mm_unpacklo_epi32( t1, t3 ); // 04, 14, 24, 34, 05, 15, 25, 35
145		r3 = _mm_unpackhi_epi32( t1, t3 ); // 06, 16, 26, 36, 07, 17, 27, 37
146		r4 = _mm_unpacklo_epi32( t4, t6 ); // 40, 50, 60, 70, 41, 51, 61, 71
147		r5 = _mm_unpackhi_epi32( t4, t6 ); // 42, 52, 62, 72, 43, 53, 63, 73
148		r6 = _mm_unpacklo_epi32( t5, t7 ); // 44, 54, 64, 74, 45, 55, 65, 75
149		r7 = _mm_unpackhi_epi32( t5, t7 ); // 46, 56, 66, 76, 47, 57, 67, 77
150
151		t0 = _mm_unpacklo_epi64( r0, r4 ); // 00, 10, 20, 30, 40, 50, 60, 70
152		t1 = _mm_unpackhi_epi64( r0, r4 ); // 01, 11, 21, 31, 41, 52, 61, 71
153		t2 = _mm_unpacklo_epi64( r1, r5 ); // 02, 12, 22, 32, 42, 52, 62, 72
154		t3 = _mm_unpackhi_epi64( r1, r5 ); // 03, 13, 23, 33, 43, 53, 63, 73
155		t4 = _mm_unpacklo_epi64( r2, r6 );
156		t5 = _mm_unpackhi_epi64( r2, r6 );
157		t6 = _mm_unpacklo_epi64( r3, r7 );
158		t7 = _mm_unpackhi_epi64( r3, r7 );
159
160		// Store
161		store16( rdi, t0 );
162		store16( rdi + destStride, t1 );
163		store16( rdi + destStride * 2, t2 );
164		store16( rdi5 - destStride * 2, t3 );
165		store16( rdi5 - destStride, t4 );
166		store16( rdi5, t5 );
167		store16( rdi5 + destStride, t6 );
168		store16( rdi5 + destStride * 2, t7 );
169	}
170
171#pragma loop( no_vector )
172	for( size_t i = 0; i < rem; rsi++, rsi5++, rdi += destStride )
173	{
174		const int16_t* p0 = (const int16_t*)rsi;
175		const int16_t* p5 = (const int16_t*)rsi5;
176		// Load a partial column into vector
177		__m128i v = _mm_cvtsi32_si128( *rsi );
178		switch( h )
179		{
180		case 7:
181			v = _mm_insert_epi16( v, *( p5 + sourceStride ), 6 );
182		case 6:
183			v = _mm_insert_epi16( v, *( p5 ), 5 );
184		case 5:
185			v = _mm_insert_epi16( v, *( p5 - sourceStride ), 4 );
186		case 4:
187			v = _mm_insert_epi16( v, *( p5 - sourceStride * 2 ), 3 );
188		case 3:
189			v = _mm_insert_epi16( v, *( p0 + sourceStride * 2 ), 2 );
190		case 2:
191			v = _mm_insert_epi16( v, *( p0 + sourceStride ), 1 );
192		}
193		// Store 8 FP16 values
194		store16( rdi, v );
195	}
196}
197
198// Same as above, but skip the transpose. The source stride is distance between columns of the matrix.
199__forceinline void copyColumnMajor( uint16_t* rdi, size_t w, const uint16_t* rsi, size_t sourceStride, size_t destStride )
200{
201	assert( 0 == ( (size_t)rdi ) % 16 );
202	assert( 0 == destStride % 8 );
203
204	constexpr size_t maskAlign4 = ~(size_t)3;
205
206	const uint16_t* const rsiEndAligned = rsi + sourceStride * ( w & maskAlign4 );
207	const uint16_t* const rsiEnd = rsi + sourceStride * w;
208	for( ; rsi < rsiEndAligned; rsi += sourceStride * 4, rdi += destStride * 4 )
209	{
210		__m128i c = f16Load( rsi );
211		store16( rdi, c );
212
213		c = f16Load( rsi + sourceStride );
214		store16( rdi + destStride, c );
215
216		c = f16Load( rsi + sourceStride * 2 );
217		store16( rdi + destStride * 2, c );
218
219		c = f16Load( rsi + sourceStride * 3 );
220		store16( rdi + destStride * 3, c );
221	}
222
223	for( ; rsi < rsiEnd; rsi += sourceStride, rdi += destStride )
224	{
225		__m128i c = f16Load( rsi );
226		store16( rdi, c );
227	}
228}
229
230__forceinline __m128i loadPartial( const uint16_t* x, size_t count )
231{
232	assert( count < 8 );
233	__m128i ix;
234	switch( count )
235	{
236	case 1: // load 2 bytes
237		ix = _mm_cvtsi32_si128( *x );
238		break;
239	case 2: // load 4 bytes
240		ix = _mm_cvtsi32_si128( *(const int*)x );
241		break;
242	case 3: // load 6 bytes
243		ix = _mm_cvtsi32_si128( *(const int*)x );
244		ix = _mm_insert_epi16( ix, x[ 2 ], 2 );
245		break;
246	case 4: // load 8 bytes
247		ix = _mm_cvtsi64_si128( *(const int64_t*)x );
248		break;
249	case 5: // load 10 bytes
250		ix = _mm_cvtsi64_si128( *(const int64_t*)x );
251		ix = _mm_insert_epi16( ix, x[ 4 ], 4 );
252		break;
253	case 6: // load 12 bytes
254		ix = _mm_cvtsi64_si128( *(const int64_t*)x );
255		ix = _mm_insert_epi32( ix, *(const int*)( x + 4 ), 2 );
256		break;
257	case 7: // load 14 bytes
258		ix = _mm_cvtsi64_si128( *(const int64_t*)x );
259		ix = _mm_insert_epi32( ix, *(const int*)( x + 4 ), 2 );
260		ix = _mm_insert_epi16( ix, x[ 6 ], 6 );
261		break;
262	default:
263		return _mm_setzero_si128();
264	}
265	return ix;
266}
267
268inline void copyColumnMajorPartial( uint16_t* rdi, size_t w, size_t h, const uint16_t* rsi, size_t sourceStride, size_t destStride )
269{
270	assert( 0 == ( (size_t)rdi ) % 32 );
271	assert( 0 == destStride % 8 );
272	assert( h > 0 && h < 8 );
273
274	const uint16_t* const rsiEnd = rsi + sourceStride * w;
275	for( ; rsi < rsiEnd; rsi += sourceStride, rdi += destStride )
276	{
277		// Can't use mask loads because loading 2-byte elements
278		// Still, that switch() in loadPartial makes a very predictable branch, same outcome for all iterations of this loop.
279		__m128i c = loadPartial( rsi, h );
280		store16( rdi, c );
281	}
282}
283
284// Store zeros into block of memory, with aligned AVX store instructions
285__forceinline void zeroAlignedMemory( void* pv, size_t cb )
286{
287	assert( 0 == cb % 16 );
288	assert( 0 == ( (size_t)pv % 32 ) );
289
290	uint8_t* rdi = (uint8_t*)pv;
291	constexpr size_t maskAlign32 = ~(size_t)31;
292	uint8_t* const rdiEndAligned = rdi + ( cb & maskAlign32 );
293	uint8_t* const rdiEnd = rdi + cb;
294
295	const __m256 zero = _mm256_setzero_ps();
296	for( ; rdi < rdiEndAligned; rdi += 32 )
297		_mm256_store_ps( (float*)rdi, zero );
298
299	if( rdi < rdiEnd )
300		_mm_store_ps( (float*)rdi, _mm_setzero_ps() );
301}