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
18.3 KiB738 linesraw
1#include "stdafx.h"
2#include "simdUtils.h"
3#include "../ML/LookupTablesData.h"
4#include <cmath>
5#include <memory>
6
7namespace
8{
9	constexpr size_t maskAlign8 = ~(size_t)7;
10
11	__forceinline __m256 load8( const uint16_t* rsi )
12	{
13		__m128i i = _mm_loadu_si128( ( const __m128i* )rsi );
14		return _mm256_cvtph_ps( i );
15	}
16
17	__forceinline void loadPartial( const uint16_t* x, const uint16_t* y, size_t count, __m256& fx, __m256& fy )
18	{
19		assert( count < 8 );
20
21		__m128i ix, iy;
22		switch( count )
23		{
24		case 1: // load 2 bytes
25			ix = _mm_cvtsi32_si128( *x );
26			iy = _mm_cvtsi32_si128( *y );
27			break;
28		case 2: // load 4 bytes
29			ix = _mm_cvtsi32_si128( *(const int*)x );
30			iy = _mm_cvtsi32_si128( *(const int*)y );
31			break;
32		case 3: // load 6 bytes
33			ix = _mm_cvtsi32_si128( *(const int*)x );
34			iy = _mm_cvtsi32_si128( *(const int*)y );
35			ix = _mm_insert_epi16( ix, x[ 2 ], 2 );
36			iy = _mm_insert_epi16( iy, y[ 2 ], 2 );
37			break;
38		case 4: // load 8 bytes
39			ix = _mm_cvtsi64_si128( *(const int64_t*)x );
40			iy = _mm_cvtsi64_si128( *(const int64_t*)y );
41			break;
42		case 5: // load 10 bytes
43			ix = _mm_cvtsi64_si128( *(const int64_t*)x );
44			iy = _mm_cvtsi64_si128( *(const int64_t*)y );
45			ix = _mm_insert_epi16( ix, x[ 4 ], 4 );
46			iy = _mm_insert_epi16( iy, y[ 4 ], 4 );
47			break;
48		case 6: // load 12 bytes
49			ix = _mm_cvtsi64_si128( *(const int64_t*)x );
50			iy = _mm_cvtsi64_si128( *(const int64_t*)y );
51			ix = _mm_insert_epi32( ix, *(const int*)( x + 4 ), 2 );
52			iy = _mm_insert_epi32( iy, *(const int*)( y + 4 ), 2 );
53			break;
54		case 7: // load 14 bytes
55			ix = _mm_cvtsi64_si128( *(const int64_t*)x );
56			iy = _mm_cvtsi64_si128( *(const int64_t*)y );
57			ix = _mm_insert_epi32( ix, *(const int*)( x + 4 ), 2 );
58			iy = _mm_insert_epi32( iy, *(const int*)( y + 4 ), 2 );
59			ix = _mm_insert_epi16( ix, x[ 6 ], 6 );
60			iy = _mm_insert_epi16( iy, y[ 6 ], 6 );
61			break;
62		default:
63			fx = fy = _mm256_setzero_ps();
64			return;
65		}
66
67		fx = _mm256_cvtph_ps( ix );
68		fy = _mm256_cvtph_ps( iy );
69	}
70
71	__forceinline __m256 loadPartial( const uint16_t* x, size_t count )
72	{
73		assert( count < 8 );
74		__m128i ix;
75		switch( count )
76		{
77		case 1: // load 2 bytes
78			ix = _mm_cvtsi32_si128( *x );
79			break;
80		case 2: // load 4 bytes
81			ix = _mm_cvtsi32_si128( *(const int*)x );
82			break;
83		case 3: // load 6 bytes
84			ix = _mm_cvtsi32_si128( *(const int*)x );
85			ix = _mm_insert_epi16( ix, x[ 2 ], 2 );
86			break;
87		case 4: // load 8 bytes
88			ix = _mm_cvtsi64_si128( *(const int64_t*)x );
89			break;
90		case 5: // load 10 bytes
91			ix = _mm_cvtsi64_si128( *(const int64_t*)x );
92			ix = _mm_insert_epi16( ix, x[ 4 ], 4 );
93			break;
94		case 6: // load 12 bytes
95			ix = _mm_cvtsi64_si128( *(const int64_t*)x );
96			ix = _mm_insert_epi32( ix, *(const int*)( x + 4 ), 2 );
97			break;
98		case 7: // load 14 bytes
99			ix = _mm_cvtsi64_si128( *(const int64_t*)x );
100			ix = _mm_insert_epi32( ix, *(const int*)( x + 4 ), 2 );
101			ix = _mm_insert_epi16( ix, x[ 6 ], 6 );
102			break;
103		default:
104			return _mm256_setzero_ps();
105		}
106		return  _mm256_cvtph_ps( ix );
107	}
108
109	__forceinline __m128 loadFloat2( const float* rsi )
110	{
111		return _mm_castpd_ps( _mm_load_sd( (const double*)rsi ) );
112	}
113	__forceinline __m128 loadFloat3( const float* rsi )
114	{
115		__m128 f = loadFloat2( rsi );
116		f = _mm_insert_ps( f, _mm_load_ss( rsi + 2 ), 0x20 );
117		return f;
118	}
119
120	__forceinline __m256 loadPartial( const float* rsi, size_t count )
121	{
122		assert( count < 8 );
123		__m128 low = _mm_setzero_ps();
124		__m128 high = _mm_setzero_ps();
125		switch( count )
126		{
127		case 1:
128			low = _mm_load_ss( rsi );
129			break;
130		case 2:
131			low = loadFloat2( rsi );
132			break;
133		case 3:
134			low = loadFloat3( rsi );
135			break;
136		case 4:
137			low = _mm_loadu_ps( rsi );
138			break;
139		case 5:
140			low = _mm_loadu_ps( rsi );
141			high = _mm_load_ss( rsi + 4 );
142			break;
143		case 6:
144			low = _mm_loadu_ps( rsi );
145			high = loadFloat2( rsi + 4 );
146			break;
147		case 7:
148			low = _mm_loadu_ps( rsi );
149			high = loadFloat3( rsi + 4 );
150			break;
151		}
152		return _mm256_setr_m128( low, high );
153	}
154
155	__forceinline void storeFloat2( float* rdi, __m128 vec )
156	{
157		_mm_store_sd( (double*)rdi, _mm_castps_pd( vec ) );
158	}
159
160	__forceinline void storePartial( float* rdi, __m256 vec, size_t count )
161	{
162		assert( count < 8 );
163
164		__m128 tmp = _mm256_castps256_ps128( vec );
165		if( count >= 4 )
166		{
167			_mm_storeu_ps( rdi, tmp );
168			if( count == 4 )
169				return;
170			count -= 4;
171			rdi += 4;
172			tmp = _mm256_extractf128_ps( vec, 1 );
173		}
174
175		switch( count )
176		{
177		case 1:
178			_mm_store_ss( rdi, tmp );
179			return;
180		case 2:
181			storeFloat2( rdi, tmp );
182			return;
183		case 3:
184			storeFloat2( rdi, tmp );
185			( (int*)rdi )[ 2 ] = _mm_extract_ps( tmp, 2 );
186			return;
187		}
188	}
189}
190
191void addF16to32( float* rdi, const uint16_t* a, const uint16_t* b, size_t length )
192{
193	const uint16_t* const endAligned = a + ( length & maskAlign8 );
194	const size_t rem = length % 8;
195
196	for( ; a < endAligned; a += 8, b += 8, rdi += 8 )
197	{
198		__m256 f1 = load8( a );
199		__m256 f2 = load8( b );
200		__m256 res = _mm256_add_ps( f1, f2 );
201		_mm256_storeu_ps( rdi, res );
202	}
203
204	if( rem != 0 )
205	{
206		__m256 f1, f2;
207		loadPartial( a, b, rem, f1, f2 );
208		__m256 res = _mm256_add_ps( f1, f2 );
209		storePartial( rdi, res, rem );
210	}
211}
212
213void addF16to32( float* rdi, const uint16_t* a, const float* b, size_t length )
214{
215	const uint16_t* const endAligned = a + ( length & maskAlign8 );
216	const size_t rem = length % 8;
217
218	for( ; a < endAligned; a += 8, b += 8, rdi += 8 )
219	{
220		__m256 f1 = load8( a );
221		__m256 f2 = _mm256_loadu_ps( b );
222		__m256 res = _mm256_add_ps( f1, f2 );
223		_mm256_storeu_ps( rdi, res );
224	}
225
226	if( rem != 0 )
227	{
228		__m256 f1 = loadPartial( a, rem );
229		__m256 f2 = loadPartial( b, rem );
230		__m256 res = _mm256_add_ps( f1, f2 );
231		storePartial( rdi, res, rem );
232	}
233}
234
235alignas( 64 ) const std::array<int, 16> s_zeroTailMask =
236{
237	-1,-1,-1,-1,-1,-1,-1,-1,
238	0, 0, 0, 0, 0, 0, 0, 0,
239};
240
241namespace
242{
243	__forceinline float horizontalSum( __m256 vec )
244	{
245		__m128 v = _mm256_extractf128_ps( vec, 1 );
246		v = _mm_add_ps( v, _mm256_castps256_ps128( vec ) );
247		v = _mm_add_ps( v, _mm_movehl_ps( v, v ) );
248		v = _mm_add_ss( v, _mm_movehdup_ps( v ) );
249		return _mm_cvtss_f32( v );
250	}
251}
252
253void norm( float* rdi, float* temp, const float* rsi, size_t length )
254{
255	assert( (size_t)temp % 32 == 0 );
256	const float* rsiEndAligned = rsi + ( length & maskAlign8 );
257	const size_t rem = length % 8;
258
259	// First pass: copy to temp buffer, and compute the sum; computeVectorSum() in HLSL
260	__m256 sum = _mm256_setzero_ps();
261	float* t;
262	for( t = temp; rsi < rsiEndAligned; rsi += 8, t += 8 )
263	{
264		__m256 v = _mm256_loadu_ps( rsi );
265		sum = _mm256_add_ps( sum, v );
266		_mm256_store_ps( t, v );
267	}
268	float* const tEndAligned = t;
269	if( 0 != rem )
270	{
271		__m256 v = loadPartial( rsi, rem );
272		sum = _mm256_add_ps( sum, v );
273		_mm256_store_ps( t, v );
274		t += 8;
275	}
276
277	const float lengthFloat = (float)(int)length;
278	const float meanScalar = horizontalSum( sum ) / lengthFloat;
279	const __m256 mean = _mm256_set1_ps( meanScalar );
280
281	// Second pass, offsetAndComputeSumSquares() in HLSL
282	sum = _mm256_setzero_ps();
283	for( t = temp; t < tEndAligned; t += 8 )
284	{
285		__m256 v = _mm256_load_ps( t );
286		v = _mm256_sub_ps( v, mean );
287		_mm256_store_ps( t, v );
288		sum = _mm256_fmadd_ps( v, v, sum );
289	}
290	if( 0 != rem )
291	{
292		__m256 v = _mm256_load_ps( t );
293		v = _mm256_sub_ps( v, mean );
294		v = _mm256_and_ps( v, loadTailMaskFloats( rem ) );
295		_mm256_store_ps( t, v );
296		sum = _mm256_fmadd_ps( v, v, sum );
297	}
298
299	// Final pass: scale, and copy from temporary buffer into the destination row
300
301	constexpr float eps = 1e-5f; // TODO: make this a parameter
302	const float scaleScalar = 1.0f / std::sqrtf( horizontalSum( sum ) / lengthFloat + eps );
303	const __m256 scale = _mm256_set1_ps( scaleScalar );
304
305	for( t = temp; t < tEndAligned; t += 8, rdi += 8 )
306	{
307		__m256 v = _mm256_load_ps( t );
308		v = _mm256_mul_ps( v, scale );
309		_mm256_storeu_ps( rdi, v );
310	}
311	if( 0 != rem )
312	{
313		__m256 v = _mm256_load_ps( t );
314		v = _mm256_mul_ps( v, scale );
315		storePartial( rdi, v, rem );
316	}
317}
318
319void fmaRepeatRow( float* rdi, size_t len, const float* w, const float* b, size_t lenPattern )
320{
321	float* rdiEndAligned = rdi + ( len & maskAlign8 );
322	const size_t rem = len % 8;
323
324	if( 1 == lenPattern )
325	{
326		const __m256 v1 = _mm256_broadcast_ss( w );
327		const __m256 v2 = _mm256_broadcast_ss( b );
328		for( ; rdi < rdiEndAligned; rdi += 8 )
329		{
330			__m256 v = _mm256_loadu_ps( rdi );
331			v = _mm256_fmadd_ps( v, v1, v2 );
332			_mm256_storeu_ps( rdi, v );
333		}
334		if( 0 != rem )
335		{
336			const __m256i mask = loadTailMaskInt( rem );
337			__m256 v = _mm256_maskload_ps( rdi, mask );
338			v = _mm256_fmadd_ps( v, v1, v2 );
339			_mm256_maskstore_ps( rdi, mask, v );
340		}
341	}
342	else if( len == lenPattern )
343	{
344		for( ; rdi < rdiEndAligned; rdi += 8, w += 8, b += 8 )
345		{
346			__m256 v = _mm256_loadu_ps( rdi );
347			__m256 v1 = _mm256_loadu_ps( w );
348			__m256 v2 = _mm256_loadu_ps( b );
349			v = _mm256_fmadd_ps( v, v1, v2 );
350			_mm256_storeu_ps( rdi, v );
351		}
352		if( 0 != rem )
353		{
354			const __m256i mask = loadTailMaskInt( rem );
355			__m256 v = _mm256_maskload_ps( rdi, mask );
356			__m256 v1 = _mm256_maskload_ps( w, mask );
357			__m256 v2 = _mm256_maskload_ps( b, mask );
358			v = _mm256_fmadd_ps( v, v1, v2 );
359			_mm256_maskstore_ps( rdi, mask, v );
360		}
361	}
362	else
363	{
364		// TODO: implement if this actually happens
365		throw E_NOTIMPL;
366	}
367}
368
369void __vectorcall addRepeatScaleRow( float* rdi, size_t len, const float* b, size_t lenPattern, const __m256 scale )
370{
371	float* rdiEndAligned = rdi + ( len & maskAlign8 );
372	const size_t rem = len % 8;
373
374	if( 1 == lenPattern )
375	{
376		const __m256 v2 = _mm256_broadcast_ss( b );
377		for( ; rdi < rdiEndAligned; rdi += 8 )
378		{
379			__m256 v = _mm256_loadu_ps( rdi );
380			v = _mm256_add_ps( v, v2 );
381			v = _mm256_mul_ps( v, scale );
382			_mm256_storeu_ps( rdi, v );
383		}
384		if( 0 != rem )
385		{
386			const __m256i mask = loadTailMaskInt( rem );
387			__m256 v = _mm256_maskload_ps( rdi, mask );
388			v = _mm256_add_ps( v, v2 );
389			v = _mm256_mul_ps( v, scale );
390			_mm256_maskstore_ps( rdi, mask, v );
391		}
392		return;
393	}
394	else if( len == lenPattern )
395	{
396		for( ; rdi < rdiEndAligned; rdi += 8, b += 8 )
397		{
398			__m256 v = _mm256_loadu_ps( rdi );
399			__m256 v2 = _mm256_loadu_ps( b );
400			v = _mm256_add_ps( v, v2 );
401			v = _mm256_mul_ps( v, scale );
402			_mm256_storeu_ps( rdi, v );
403		}
404		if( 0 != rem )
405		{
406			const __m256i mask = loadTailMaskInt( rem );
407			__m256 v = _mm256_maskload_ps( rdi, mask );
408			__m256 v2 = _mm256_maskload_ps( b, mask );
409			v = _mm256_add_ps( v, v2 );
410			v = _mm256_mul_ps( v, scale );
411			_mm256_maskstore_ps( rdi, mask, v );
412		}
413		return;
414	}
415	else
416	{
417		// TODO: implement if this actually happens
418		throw E_NOTIMPL;
419	}
420}
421
422void addRepeatRow( float* rdi, size_t len, const float* b, size_t lenPattern )
423{
424	float* rdiEndAligned = rdi + ( len & maskAlign8 );
425	const size_t rem = len % 8;
426
427	if( 1 == lenPattern )
428	{
429		const __m256 v2 = _mm256_broadcast_ss( b );
430		for( ; rdi < rdiEndAligned; rdi += 8 )
431		{
432			__m256 v = _mm256_loadu_ps( rdi );
433			v = _mm256_add_ps( v, v2 );
434			_mm256_storeu_ps( rdi, v );
435		}
436		if( 0 != rem )
437		{
438			const __m256i mask = loadTailMaskInt( rem );
439			__m256 v = _mm256_maskload_ps( rdi, mask );
440			v = _mm256_add_ps( v, v2 );
441			_mm256_maskstore_ps( rdi, mask, v );
442		}
443		return;
444	}
445	else if( len == lenPattern )
446	{
447		for( ; rdi < rdiEndAligned; rdi += 8, b += 8 )
448		{
449			__m256 v = _mm256_loadu_ps( rdi );
450			__m256 v2 = _mm256_loadu_ps( b );
451			v = _mm256_add_ps( v, v2 );
452			_mm256_storeu_ps( rdi, v );
453		}
454		if( 0 != rem )
455		{
456			const __m256i mask = loadTailMaskInt( rem );
457			__m256 v = _mm256_maskload_ps( rdi, mask );
458			__m256 v2 = _mm256_maskload_ps( b, mask );
459			v = _mm256_add_ps( v, v2 );
460			_mm256_maskstore_ps( rdi, mask, v );
461		}
462		return;
463	}
464	else
465	{
466		// TODO: implement if this actually happens
467		throw E_NOTIMPL;
468	}
469}
470
471namespace
472{
473	__forceinline __m256 gelu( __m256 x, const DirectCompute::LookupTablesData& lookup )
474	{
475		__m128i iv = _mm256_cvtps_ph( x, 0 );
476		alignas( 16 ) std::array<uint16_t, 8> arr;
477		_mm_store_si128( ( __m128i* )arr.data(), iv );
478		for( uint16_t& a : arr )
479			a = lookup.gelu[ a ];
480		iv = _mm_load_si128( ( __m128i* )arr.data() );
481		return _mm256_cvtph_ps( iv );
482	}
483}
484
485void addRepeatGeluRow( float* rdi, size_t len, const float* b, size_t lenPattern, const DirectCompute::LookupTablesData& lookup )
486{
487	float* rdiEndAligned = rdi + ( len & maskAlign8 );
488	const size_t rem = len % 8;
489
490	if( 1 == lenPattern )
491	{
492		const __m256 v2 = _mm256_broadcast_ss( b );
493		for( ; rdi < rdiEndAligned; rdi += 8 )
494		{
495			__m256 v = _mm256_loadu_ps( rdi );
496			v = _mm256_add_ps( v, v2 );
497			v = gelu( v, lookup );
498			_mm256_storeu_ps( rdi, v );
499		}
500		if( 0 != rem )
501		{
502			const __m256i mask = loadTailMaskInt( rem );
503			__m256 v = _mm256_maskload_ps( rdi, mask );
504			v = _mm256_add_ps( v, v2 );
505			v = gelu( v, lookup );
506			_mm256_maskstore_ps( rdi, mask, v );
507		}
508		return;
509	}
510	else if( len == lenPattern )
511	{
512		for( ; rdi < rdiEndAligned; rdi += 8, b += 8 )
513		{
514			__m256 v = _mm256_loadu_ps( rdi );
515			__m256 v2 = _mm256_loadu_ps( b );
516			v = _mm256_add_ps( v, v2 );
517			v = gelu( v, lookup );
518			_mm256_storeu_ps( rdi, v );
519		}
520		if( 0 != rem )
521		{
522			const __m256i mask = loadTailMaskInt( rem );
523			__m256 v = _mm256_maskload_ps( rdi, mask );
524			__m256 v2 = _mm256_maskload_ps( b, mask );
525			v = _mm256_add_ps( v, v2 );
526			v = gelu( v, lookup );
527			_mm256_maskstore_ps( rdi, mask, v );
528		}
529		return;
530	}
531	else
532	{
533		// TODO: implement if this actually happens
534		throw E_NOTIMPL;
535	}
536}
537
538void __vectorcall scaleRow( float* rdi, size_t len, const __m256 scale )
539{
540	float* rdiEndAligned = rdi + ( len & maskAlign8 );
541	const size_t rem = len % 8;
542	for( ; rdi < rdiEndAligned; rdi += 8 )
543	{
544		__m256 v = _mm256_loadu_ps( rdi );
545		v = _mm256_mul_ps( v, scale );
546		_mm256_storeu_ps( rdi, v );
547	}
548	if( 0 != rem )
549	{
550		const __m256i mask = loadTailMaskInt( rem );
551		__m256 v = _mm256_maskload_ps( rdi, mask );
552		v = _mm256_mul_ps( v, scale );
553		_mm256_maskstore_ps( rdi, mask, v );
554	}
555}
556
557namespace
558{
559	using DirectCompute::LookupTablesData;
560
561	__forceinline float horizontalMax( __m256 vec )
562	{
563		__m128 v = _mm256_extractf128_ps( vec, 1 );
564		v = _mm_max_ps( v, _mm256_castps256_ps128( vec ) );
565		v = _mm_max_ps( v, _mm_movehl_ps( v, v ) );
566		v = _mm_max_ss( v, _mm_movehdup_ps( v ) );
567		return _mm_cvtss_f32( v );
568	}
569
570	__forceinline float _cvtsh_ss( uint16_t f16 )
571	{
572		__m128i i = _mm_cvtsi32_si128( f16 );
573		__m128 f = _mm_cvtph_ps( i );
574		return _mm_cvtss_f32( f );
575	}
576
577	__forceinline uint16_t _cvtss_sh( float f, int rounding )
578	{
579		assert( 0 == rounding );
580		__m128 v = _mm_set_ss( f );
581		__m128i i = _mm_cvtps_ph( v, 0 );
582		return (uint16_t)(uint32_t)_mm_cvtsi128_si32( i );
583	}
584}
585
586const LookupTablesData& getLookupTables()
587{
588	static const std::unique_ptr<LookupTablesData> res = std::make_unique<LookupTablesData>();
589	return *res;
590}
591
592void softMax( float* rdi, size_t length, const float inputScale )
593{
594	float* const rdiBegin = rdi;
595	float* const rdiEndAligned = rdi + ( length & maskAlign8 );
596	const size_t remainder = length % 8;
597	// First pass, compute maximum
598	__m256 max = _mm256_set1_ps( -INFINITY );
599	for( rdi = rdiBegin; rdi < rdiEndAligned; rdi += 8 )
600	{
601		__m256 v = _mm256_loadu_ps( rdi );
602		max = _mm256_max_ps( max, v );
603	}
604	__m256i tailMask;
605	if( 0 != remainder )
606	{
607		tailMask = loadTailMaskInt( remainder );
608		__m256 v = _mm256_maskload_ps( rdi, tailMask );
609		v = _mm256_max_ps( max, v );
610		max = _mm256_blendv_ps( max, v, _mm256_castsi256_ps( tailMask ) );
611	}
612
613	// Second pass: apply initial scale, compute the exponent, and compute total sum over the row
614	const LookupTablesData& lookup = getLookupTables();
615	const float maxScalar = horizontalMax( max );
616
617	float* const rdiEnd = rdiBegin + length;
618	double sum = 0;
619	for( rdi = rdiBegin; rdi < rdiEnd; rdi++ )
620	{
621		// Possible to vectorize, but relatively hard
622		// An easy way is upcast the complete lookup table to FP32 and then use two _mm256_i32gather_ps instructions per iteration
623		// However, that instruction is from AVX2 set. Let's hope this loop won't be a bottleneck.
624		float f = *rdi;
625		if( f != -INFINITY )
626		{
627			f = ( f - maxScalar ) * inputScale;
628			uint16_t f16 = _cvtss_sh( f, 0 );
629			f16 = lookup.exponent[ f16 ];
630			f = _cvtsh_ss( f16 );
631			sum += f;
632		}
633		else
634			f = 0;
635
636		*rdi = f;
637	}
638
639	// Final pass: apply the final scale
640	const __m256 finalScale = _mm256_set1_ps( (float)( 1.0 / sum ) );
641	for( rdi = rdiBegin; rdi < rdiEndAligned; rdi += 8 )
642	{
643		__m256 v = _mm256_loadu_ps( rdi );
644		v = _mm256_mul_ps( v, finalScale );
645		_mm256_storeu_ps( rdi, v );
646	}
647	if( 0 != remainder )
648	{
649		__m256 v = _mm256_maskload_ps( rdi, tailMask );
650		v = _mm256_mul_ps( v, finalScale );
651		_mm256_maskstore_ps( rdi, tailMask, v );
652	}
653}
654
655void floatsUpcast( float* rdi, const uint16_t* rsi, size_t length )
656{
657	const uint16_t* rsiEndAligned = rsi + ( length & maskAlign8 );
658	const size_t rem = length % 8;
659
660	for( ; rsi < rsiEndAligned; rsi += 8, rdi += 8 )
661		_mm256_storeu_ps( rdi, load8( rsi ) );
662
663	if( 0 != rem )
664	{
665		__m256 v = loadPartial( rsi, rem );
666		_mm256_maskstore_ps( rdi, loadTailMaskInt( rem ), v );
667	}
668}
669
670void floatsDowncast( uint16_t* rdi, const float* rsi, size_t length )
671{
672	const float* rsiEndAligned = rsi + ( length & maskAlign8 );
673	size_t rem = length % 8;
674
675	for( ; rsi < rsiEndAligned; rsi += 8, rdi += 8 )
676	{
677		__m256 vf = _mm256_loadu_ps( rsi );
678		__m128i vi = _mm256_cvtps_ph( vf, 0 );
679		store16( rdi, vi );
680	}
681
682	if( 0 != rem )
683	{
684		__m256 vf = _mm256_maskload_ps( rsi, loadTailMaskInt( rem ) );
685		__m128i vi = _mm256_cvtps_ph( vf, 0 );
686		for( size_t i = 0; i < rem; i++, rdi++ )
687		{
688			*rdi = (uint16_t)(uint32_t)_mm_cvtsi128_si32( vi );
689			vi = _mm_srli_si128( vi, 2 );
690		}
691	}
692}
693
694void addRowInPlace( float* rdi, const float* rsi, size_t length )
695{
696	const float* rdiEndAligned = rdi + ( length & maskAlign8 );
697	size_t rem = length % 8;
698
699	for( ; rdi < rdiEndAligned; rdi += 8, rsi += 8 )
700	{
701		__m256 a = _mm256_loadu_ps( rdi );
702		__m256 b = _mm256_loadu_ps( rsi );
703		a = _mm256_add_ps( a, b );
704		_mm256_storeu_ps( rdi, a );
705	}
706
707	if( 0 != rem )
708	{
709		const __m256i mask = loadTailMaskInt( rem );
710		__m256 a = _mm256_maskload_ps( rdi, mask );
711		__m256 b = _mm256_maskload_ps( rsi, mask );
712		a = _mm256_add_ps( a, b );
713		_mm256_maskstore_ps( rdi, mask, a );
714	}
715}
716
717void addRow( float* rdi, const float* a, const float* b, size_t length )
718{
719	const float* aEndAligned = a + ( length & maskAlign8 );
720	size_t rem = length % 8;
721
722	for( ; a < aEndAligned; a += 8, b += 8, rdi += 8 )
723	{
724		__m256 x = _mm256_loadu_ps( a );
725		__m256 y = _mm256_loadu_ps( b );
726		x = _mm256_add_ps( x, y );
727		_mm256_storeu_ps( rdi, x );
728	}
729
730	if( 0 != rem )
731	{
732		const __m256i mask = loadTailMaskInt( rem );
733		__m256 x = _mm256_maskload_ps( a, mask );
734		__m256 y = _mm256_maskload_ps( b, mask );
735		x = _mm256_add_ps( x, y );
736		_mm256_maskstore_ps( rdi, mask, x );
737	}
738}