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
21.0 KiB742 linesraw
1#pragma once
2#include <stdint.h>
3#include <immintrin.h>
4#include "simdUtils.h"
5
6template<uint8_t panelHeightRegs, uint8_t tileWidthFloats>
7struct ResultTile
8{
9	static constexpr size_t totalRegs = (size_t)(tileWidthFloats)*panelHeightRegs;
10	std::array<__m256, totalRegs> arr;
11
12	template<size_t idx>
13	__forceinline void fmadd( __m256 a, __m256 b )
14	{
15		arr[ idx ] = _mm256_fmadd_ps( a, b, arr[ idx ] );
16	}
17	__forceinline void kernel( const std::array<__m256, panelHeightRegs>& panel, const float* rsi, size_t stride );
18	__forceinline void kernelPartial( const std::array<__m256, panelHeightRegs>& panel, const float* rsi, size_t stride, size_t rem )
19	{
20		throw E_UNEXPECTED;
21	}
22	__forceinline void store( float* rdi, size_t w, size_t h, size_t stride ) const;
23};
24
25#pragma region setZero functions
26__forceinline void setZero( std::array<__m256, 1>& dest )
27{
28	dest[ 0 ] = _mm256_setzero_ps();
29}
30__forceinline void setZero( std::array<__m256, 2>& dest )
31{
32	dest[ 0 ] = _mm256_setzero_ps();
33	dest[ 1 ] = _mm256_setzero_ps();
34}
35__forceinline void setZero( std::array<__m256, 3>& dest )
36{
37	dest[ 0 ] = _mm256_setzero_ps();
38	dest[ 1 ] = _mm256_setzero_ps();
39	dest[ 2 ] = _mm256_setzero_ps();
40}
41__forceinline void setZero( std::array<__m256, 4>& dest )
42{
43	dest[ 0 ] = _mm256_setzero_ps();
44	dest[ 1 ] = _mm256_setzero_ps();
45	dest[ 2 ] = _mm256_setzero_ps();
46	dest[ 3 ] = _mm256_setzero_ps();
47}
48__forceinline void setZero( std::array<__m256, 6>& dest )
49{
50	dest[ 0 ] = _mm256_setzero_ps();
51	dest[ 1 ] = _mm256_setzero_ps();
52	dest[ 2 ] = _mm256_setzero_ps();
53	dest[ 3 ] = _mm256_setzero_ps();
54	dest[ 4 ] = _mm256_setzero_ps();
55	dest[ 5 ] = _mm256_setzero_ps();
56}
57__forceinline void setZero( std::array<__m256, 8>& dest )
58{
59	dest[ 0 ] = _mm256_setzero_ps();
60	dest[ 1 ] = _mm256_setzero_ps();
61	dest[ 2 ] = _mm256_setzero_ps();
62	dest[ 3 ] = _mm256_setzero_ps();
63	dest[ 4 ] = _mm256_setzero_ps();
64	dest[ 5 ] = _mm256_setzero_ps();
65	dest[ 6 ] = _mm256_setzero_ps();
66	dest[ 7 ] = _mm256_setzero_ps();
67}
68#pragma endregion
69
70#pragma region Micro-kernels
71__forceinline void ResultTile<1, 1>::kernel( const std::array<__m256, 1>& panel, const float* rsi, size_t stride )
72{
73	__m256 b = _mm256_broadcast_ss( rsi );
74	fmadd<0>( panel[ 0 ], b );
75}
76__forceinline void ResultTile<1, 2>::kernel( const std::array<__m256, 1>& panel, const float* rsi, size_t stride )
77{
78	__m256 b = _mm256_broadcast_ss( rsi );
79	fmadd<0>( panel[ 0 ], b );
80	b = _mm256_broadcast_ss( rsi + stride );
81	fmadd<1>( panel[ 0 ], b );
82}
83__forceinline void ResultTile<1, 2>::kernelPartial( const std::array<__m256, 1>& panel, const float* rsi, size_t stride, size_t rem )
84{
85	assert( 1 == rem );
86	__m256 b = _mm256_broadcast_ss( rsi );
87	fmadd<0>( panel[ 0 ], b );
88}
89__forceinline void ResultTile<1, 3>::kernel( const std::array<__m256, 1>& panel, const float* rsi, size_t stride )
90{
91	__m256 b = _mm256_broadcast_ss( rsi );
92	fmadd<0>( panel[ 0 ], b );
93	b = _mm256_broadcast_ss( rsi + stride );
94	fmadd<1>( panel[ 0 ], b );
95	b = _mm256_broadcast_ss( rsi + stride * 2 );
96	fmadd<2>( panel[ 0 ], b );
97}
98__forceinline void ResultTile<1, 3>::kernelPartial( const std::array<__m256, 1>& panel, const float* rsi, size_t stride, size_t rem )
99{
100	assert( rem > 0 && rem < 3 );
101	__m256 b = _mm256_broadcast_ss( rsi );
102	fmadd<0>( panel[ 0 ], b );
103	if( rem > 1 )
104	{
105		b = _mm256_broadcast_ss( rsi + stride );
106		fmadd<1>( panel[ 0 ], b );
107	}
108}
109
110__forceinline void ResultTile<1, 4>::kernel( const std::array<__m256, 1>& panel, const float* rsi, size_t stride )
111{
112	__m256 b = _mm256_broadcast_ss( rsi );
113	fmadd<0>( panel[ 0 ], b );
114	b = _mm256_broadcast_ss( rsi + stride );
115	fmadd<1>( panel[ 0 ], b );
116	b = _mm256_broadcast_ss( rsi + stride * 2 );
117	fmadd<2>( panel[ 0 ], b );
118	b = _mm256_broadcast_ss( rsi + stride * 3 );
119	fmadd<3>( panel[ 0 ], b );
120}
121__forceinline void ResultTile<1, 4>::kernelPartial( const std::array<__m256, 1>& panel, const float* rsi, size_t stride, size_t rem )
122{
123	assert( rem > 0 && rem < 4 );
124	__m256 b = _mm256_broadcast_ss( rsi );
125	fmadd<0>( panel[ 0 ], b );
126
127	switch( rem )
128	{
129	case 3:
130		b = _mm256_broadcast_ss( rsi + stride * 2 );
131		fmadd<2>( panel[ 0 ], b );
132	case 2:
133		b = _mm256_broadcast_ss( rsi + stride );
134		fmadd<1>( panel[ 0 ], b );
135	}
136}
137__forceinline void ResultTile<4, 1>::kernel( const std::array<__m256, 4>& panel, const float* rsi, size_t stride )
138{
139	__m256 b = _mm256_broadcast_ss( rsi );
140	fmadd<0>( panel[ 0 ], b );
141	fmadd<1>( panel[ 1 ], b );
142	fmadd<2>( panel[ 2 ], b );
143	fmadd<3>( panel[ 3 ], b );
144}
145__forceinline void ResultTile<2, 4>::kernel( const std::array<__m256, 2>& panel, const float* rsi, size_t stride )
146{
147	__m256 b = _mm256_broadcast_ss( rsi );
148	fmadd<0>( panel[ 0 ], b );
149	fmadd<1>( panel[ 1 ], b );
150
151	b = _mm256_broadcast_ss( rsi + stride );
152	fmadd<2>( panel[ 0 ], b );
153	fmadd<3>( panel[ 1 ], b );
154
155	b = _mm256_broadcast_ss( rsi + stride * 2 );
156	fmadd<4>( panel[ 0 ], b );
157	fmadd<5>( panel[ 1 ], b );
158
159	b = _mm256_broadcast_ss( rsi + stride * 3 );
160	fmadd<6>( panel[ 0 ], b );
161	fmadd<7>( panel[ 1 ], b );
162}
163
164__forceinline void ResultTile<2, 4>::kernelPartial( const std::array<__m256, 2>& panel, const float* rsi, size_t stride, size_t rem )
165{
166	assert( rem > 0 && rem < 4 );
167	__m256 b = _mm256_broadcast_ss( rsi );
168	fmadd<0>( panel[ 0 ], b );
169	fmadd<1>( panel[ 1 ], b );
170
171	switch( rem )
172	{
173	case 3:
174		b = _mm256_broadcast_ss( rsi + stride * 2 );
175		fmadd<4>( panel[ 0 ], b );
176		fmadd<5>( panel[ 1 ], b );
177	case 2:
178		b = _mm256_broadcast_ss( rsi + stride );
179		fmadd<2>( panel[ 0 ], b );
180		fmadd<3>( panel[ 1 ], b );
181	}
182}
183
184__forceinline void ResultTile<2, 3>::kernel( const std::array<__m256, 2>& panel, const float* rsi, size_t stride )
185{
186	__m256 b = _mm256_broadcast_ss( rsi );
187	fmadd<0>( panel[ 0 ], b );
188	fmadd<1>( panel[ 1 ], b );
189
190	b = _mm256_broadcast_ss( rsi + stride );
191	fmadd<2>( panel[ 0 ], b );
192	fmadd<3>( panel[ 1 ], b );
193
194	b = _mm256_broadcast_ss( rsi + stride * 2 );
195	fmadd<4>( panel[ 0 ], b );
196	fmadd<5>( panel[ 1 ], b );
197}
198__forceinline void ResultTile<2, 3>::kernelPartial( const std::array<__m256, 2>& panel, const float* rsi, size_t stride, size_t rem )
199{
200	assert( rem > 0 && rem < 3 );
201	__m256 b = _mm256_broadcast_ss( rsi );
202	fmadd<0>( panel[ 0 ], b );
203	fmadd<1>( panel[ 1 ], b );
204	if( rem > 1 )
205	{
206		b = _mm256_broadcast_ss( rsi + stride );
207		fmadd<2>( panel[ 0 ], b );
208		fmadd<3>( panel[ 1 ], b );
209	}
210}
211
212__forceinline void ResultTile<4, 2>::kernel( const std::array<__m256, 4>& panel, const float* rsi, size_t stride )
213{
214	__m256 b = _mm256_broadcast_ss( rsi );
215	fmadd<0>( panel[ 0 ], b );
216	fmadd<1>( panel[ 1 ], b );
217	fmadd<2>( panel[ 2 ], b );
218	fmadd<3>( panel[ 3 ], b );
219
220	b = _mm256_broadcast_ss( rsi + stride );
221	fmadd<4>( panel[ 0 ], b );
222	fmadd<5>( panel[ 1 ], b );
223	fmadd<6>( panel[ 2 ], b );
224	fmadd<7>( panel[ 3 ], b );
225}
226__forceinline void ResultTile<4, 2>::kernelPartial( const std::array<__m256, 4>& panel, const float* rsi, size_t stride, size_t rem )
227{
228	assert( 1 == rem );
229	__m256 b = _mm256_broadcast_ss( rsi );
230	fmadd<0>( panel[ 0 ], b );
231	fmadd<1>( panel[ 1 ], b );
232	fmadd<2>( panel[ 2 ], b );
233	fmadd<3>( panel[ 3 ], b );
234}
235#pragma endregion
236
237#pragma region Loads
238// This function should compile into a single `vcvtph2ps` instruction, with memory operand
239__forceinline __m256 loadUpcasted( const uint16_t* rsi )
240{
241	__m128i i = _mm_load_si128( ( const __m128i* )rsi );
242	return _mm256_cvtph_ps( i );
243}
244
245// We loading the panel from the temporary buffer.
246// For this reason, we don't need to handle remainders, the code which made the buffer wrote zeros into the remainder elements
247// We can even use aligned load instructions.
248__forceinline void loadPanel( const uint16_t* rsi, std::array<__m256, 1>& dest )
249{
250	dest[ 0 ] = loadUpcasted( rsi );
251}
252__forceinline void loadPanel( const uint16_t* rsi, std::array<__m256, 2>& dest )
253{
254	dest[ 0 ] = loadUpcasted( rsi );
255	dest[ 1 ] = loadUpcasted( rsi + 8 );
256}
257__forceinline void loadPanel( const uint16_t* rsi, std::array<__m256, 3>& dest )
258{
259	dest[ 0 ] = loadUpcasted( rsi );
260	dest[ 1 ] = loadUpcasted( rsi + 8 );
261	dest[ 2 ] = loadUpcasted( rsi + 8 * 2 );
262}
263__forceinline void loadPanel( const uint16_t* rsi, std::array<__m256, 4>& dest )
264{
265	dest[ 0 ] = loadUpcasted( rsi );
266	dest[ 1 ] = loadUpcasted( rsi + 8 );
267	dest[ 2 ] = loadUpcasted( rsi + 8 * 2 );
268	dest[ 3 ] = loadUpcasted( rsi + 8 * 3 );
269}
270#pragma endregion
271
272#pragma region Stores
273__forceinline void ResultTile<1, 1>::store( float* rdi, size_t w, size_t h, size_t stride ) const
274{
275	assert( h == 1 && w > 0 && w <= 8 );
276	if( w == 8 )
277		_mm256_storeu_ps( rdi, arr[ 0 ] );
278	else
279	{
280		const __m256i mask = loadTailMaskInt( w );
281		_mm256_maskstore_ps( rdi, mask, arr[ 0 ] );
282	}
283}
284
285__forceinline void ResultTile<1, 2>::store( float* rdi, size_t w, size_t h, size_t stride ) const
286{
287	assert( h > 0 && w > 0 && h <= 2 && w <= 8 );
288	if( w == 8 )
289	{
290		switch( h )
291		{
292		case 2:
293			_mm256_storeu_ps( rdi + stride, arr[ 1 ] );
294		case 1:
295			_mm256_storeu_ps( rdi, arr[ 0 ] );
296		}
297	}
298	else
299	{
300		const __m256i mask = loadTailMaskInt( w );
301		switch( h )
302		{
303		case 2:
304			_mm256_maskstore_ps( rdi + stride, mask, arr[ 1 ] );
305		case 1:
306			_mm256_maskstore_ps( rdi, mask, arr[ 0 ] );
307		}
308	}
309}
310
311__forceinline void ResultTile<1, 3>::store( float* rdi, size_t w, size_t h, size_t stride ) const
312{
313	assert( h > 0 && w > 0 && h <= 3 && w <= 8 );
314	if( w == 8 )
315	{
316		switch( h )
317		{
318		case 3:
319			_mm256_storeu_ps( rdi + stride * 2, arr[ 2 ] );
320		case 2:
321			_mm256_storeu_ps( rdi + stride, arr[ 1 ] );
322		case 1:
323			_mm256_storeu_ps( rdi, arr[ 0 ] );
324		}
325	}
326	else
327	{
328		const __m256i mask = loadTailMaskInt( w );
329		switch( h )
330		{
331		case 3:
332			_mm256_maskstore_ps( rdi + stride * 2, mask, arr[ 2 ] );
333		case 2:
334			_mm256_maskstore_ps( rdi + stride, mask, arr[ 1 ] );
335		case 1:
336			_mm256_maskstore_ps( rdi, mask, arr[ 0 ] );
337		}
338	}
339}
340
341__forceinline void ResultTile<1, 4>::store( float* rdi, size_t w, size_t h, size_t stride ) const
342{
343	assert( h > 0 && w > 0 && h <= 4 && w <= 8 );
344
345	if( w == 8 )
346	{
347		switch( h )
348		{
349		case 4:
350			_mm256_storeu_ps( rdi + stride * 3, arr[ 3 ] );
351		case 3:
352			_mm256_storeu_ps( rdi + stride * 2, arr[ 2 ] );
353		case 2:
354			_mm256_storeu_ps( rdi + stride, arr[ 1 ] );
355		case 1:
356			_mm256_storeu_ps( rdi, arr[ 0 ] );
357		}
358	}
359	else
360	{
361		const __m256i mask = loadTailMaskInt( w );
362		switch( h )
363		{
364		case 4:
365			_mm256_maskstore_ps( rdi + stride * 3, mask, arr[ 3 ] );
366		case 3:
367			_mm256_maskstore_ps( rdi + stride * 2, mask, arr[ 2 ] );
368		case 2:
369			_mm256_maskstore_ps( rdi + stride, mask, arr[ 1 ] );
370		case 1:
371			_mm256_maskstore_ps( rdi, mask, arr[ 0 ] );
372		}
373	}
374}
375
376__forceinline void ResultTile<4, 1>::store( float* rdi, size_t w, size_t h, size_t stride ) const
377{
378	assert( h == 1 && w > 0 && w <= 32 );
379	if( w == 32 )
380	{
381		// 4 complete vectors, this branch is very likely to be taken
382		_mm256_storeu_ps( rdi, arr[ 0 ] );
383		_mm256_storeu_ps( rdi + 8, arr[ 1 ] );
384		_mm256_storeu_ps( rdi + 8 * 2, arr[ 2 ] );
385		_mm256_storeu_ps( rdi + 8 * 3, arr[ 3 ] );
386	}
387	else
388	{
389		const size_t rem = w % 8;
390		const __m256i mask = loadTailMaskInt<false>( rem );
391		const size_t completeVectors = w / 8;
392		const size_t key = ( completeVectors << 1 ) | ( ( 0 == rem ) ? 0 : 1 );
393		switch( key )
394		{
395		case 1:	// 0 complete vectors + remainder
396			_mm256_maskstore_ps( rdi, mask, arr[ 0 ] );
397			break;
398		case 2:	// 1 complete vector
399			_mm256_storeu_ps( rdi, arr[ 0 ] );
400			break;
401		case 3:	// 1 complete vector + remainder
402			_mm256_storeu_ps( rdi, arr[ 0 ] );
403			_mm256_maskstore_ps( rdi + 8, mask, arr[ 1 ] );
404			break;
405		case 4:	// 2 complete vectors
406			_mm256_storeu_ps( rdi, arr[ 0 ] );
407			_mm256_storeu_ps( rdi + 8, arr[ 1 ] );
408			break;
409		case 5:	// 2 complete vectors + remainder
410			_mm256_storeu_ps( rdi, arr[ 0 ] );
411			_mm256_storeu_ps( rdi + 8, arr[ 1 ] );
412			_mm256_maskstore_ps( rdi + 8 * 2, mask, arr[ 2 ] );
413			break;
414		case 6:	// 3 complete vectors
415			_mm256_storeu_ps( rdi, arr[ 0 ] );
416			_mm256_storeu_ps( rdi + 8, arr[ 1 ] );
417			_mm256_storeu_ps( rdi + 8 * 2, arr[ 2 ] );
418			break;
419		case 7:	// 3 complete vectors + remainder
420			_mm256_storeu_ps( rdi, arr[ 0 ] );
421			_mm256_storeu_ps( rdi + 8, arr[ 1 ] );
422			_mm256_storeu_ps( rdi + 8 * 2, arr[ 2 ] );
423			_mm256_maskstore_ps( rdi + 8 * 3, mask, arr[ 3 ] );
424			break;
425		default:
426			throw E_UNEXPECTED;
427		}
428	}
429}
430__forceinline void ResultTile<4, 2>::store( float* rdi, size_t w, size_t h, size_t stride ) const
431{
432	assert( h > 0 && w > 0 && h <= 2 && w <= 32 );
433	const bool twoRows = h == 2;
434	float* const rdi1 = rdi + stride;
435	if( w == 32 )
436	{
437		_mm256_storeu_ps( rdi, arr[ 0 ] );
438		_mm256_storeu_ps( rdi + 8, arr[ 1 ] );
439		_mm256_storeu_ps( rdi + 8 * 2, arr[ 2 ] );
440		_mm256_storeu_ps( rdi + 8 * 3, arr[ 3 ] );
441
442		if( twoRows )
443		{
444			_mm256_storeu_ps( rdi1, arr[ 4 ] );
445			_mm256_storeu_ps( rdi1 + 8, arr[ 5 ] );
446			_mm256_storeu_ps( rdi1 + 8 * 2, arr[ 6 ] );
447			_mm256_storeu_ps( rdi1 + 8 * 3, arr[ 7 ] );
448		}
449	}
450	else
451	{
452		const size_t rem = w % 8;
453		const __m256i mask = loadTailMaskInt<false>( rem );
454		const size_t completeVectors = w / 8;
455		// Lowest bit: remainder
456		// Next bit: set when storing 2 rows
457		// Next 2 bits: count of complete vectors in X direction, [ 0..3 ]
458		const size_t key = ( completeVectors << 2 ) | ( ( 0 == rem ) ? 0 : 1 ) | ( twoRows ? 2 : 0 );
459		switch( key )
460		{
461		case 1:	// 0 complete vectors + remainder, 1 row
462			_mm256_maskstore_ps( rdi, mask, arr[ 0 ] );
463			break;
464		case 3:	// 0 complete vectors + remainder, 2 rows
465			_mm256_maskstore_ps( rdi, mask, arr[ 0 ] );
466			_mm256_maskstore_ps( rdi1, mask, arr[ 4 ] );
467			break;
468		case 4:	// 1 complete vector, 1 row
469			_mm256_storeu_ps( rdi, arr[ 0 ] );
470			break;
471		case 5:	// 1 complete vector + remainder, 1 row
472			_mm256_storeu_ps( rdi, arr[ 0 ] );
473			_mm256_maskstore_ps( rdi + 8, mask, arr[ 1 ] );
474			break;
475		case 6:	// 1 complete vector, 2 rows
476			_mm256_storeu_ps( rdi, arr[ 0 ] );
477			_mm256_storeu_ps( rdi1, arr[ 4 ] );
478			break;
479		case 7:	// 1 complete vector + remainder, 2 rows
480			_mm256_storeu_ps( rdi, arr[ 0 ] );
481			_mm256_maskstore_ps( rdi + 8, mask, arr[ 1 ] );
482
483			_mm256_storeu_ps( rdi1, arr[ 4 ] );
484			_mm256_maskstore_ps( rdi1 + 8, mask, arr[ 5 ] );
485			break;
486		case 8:	// 2 complete vectors, 1 row
487			_mm256_storeu_ps( rdi, arr[ 0 ] );
488			_mm256_storeu_ps( rdi + 8, arr[ 1 ] );
489			break;
490		case 9:	// 2 complete vectors + remainder, 1 row
491			_mm256_storeu_ps( rdi, arr[ 0 ] );
492			_mm256_storeu_ps( rdi + 8, arr[ 1 ] );
493			_mm256_maskstore_ps( rdi + 8 * 2, mask, arr[ 2 ] );
494			break;
495		case 10:	// 2 complete vectors, 2 rows
496			_mm256_storeu_ps( rdi, arr[ 0 ] );
497			_mm256_storeu_ps( rdi + 8, arr[ 1 ] );
498
499			_mm256_storeu_ps( rdi1, arr[ 4 ] );
500			_mm256_storeu_ps( rdi1 + 8, arr[ 5 ] );
501			break;
502		case 11:	// 2 complete vectors + remainder, 2 rows
503			_mm256_storeu_ps( rdi, arr[ 0 ] );
504			_mm256_storeu_ps( rdi + 8, arr[ 1 ] );
505			_mm256_maskstore_ps( rdi + 8 * 2, mask, arr[ 2 ] );
506
507			_mm256_storeu_ps( rdi1, arr[ 4 ] );
508			_mm256_storeu_ps( rdi1 + 8, arr[ 5 ] );
509			_mm256_maskstore_ps( rdi1 + 8 * 2, mask, arr[ 6 ] );
510			break;
511		case 12:	// 3 complete vectors, 1 row
512			_mm256_storeu_ps( rdi, arr[ 0 ] );
513			_mm256_storeu_ps( rdi + 8, arr[ 1 ] );
514			_mm256_storeu_ps( rdi + 8 * 2, arr[ 2 ] );
515			break;
516		case 13:	// 3 complete vectors + remainder, 1 row
517			_mm256_storeu_ps( rdi, arr[ 0 ] );
518			_mm256_storeu_ps( rdi + 8, arr[ 1 ] );
519			_mm256_storeu_ps( rdi + 8 * 2, arr[ 2 ] );
520			_mm256_maskstore_ps( rdi + 8 * 3, mask, arr[ 3 ] );
521			break;
522		case 14:	// 3 complete vectors, 2 rows
523			_mm256_storeu_ps( rdi, arr[ 0 ] );
524			_mm256_storeu_ps( rdi + 8, arr[ 1 ] );
525			_mm256_storeu_ps( rdi + 8 * 2, arr[ 2 ] );
526
527			_mm256_storeu_ps( rdi1, arr[ 4 ] );
528			_mm256_storeu_ps( rdi1 + 8, arr[ 5 ] );
529			_mm256_storeu_ps( rdi1 + 8 * 2, arr[ 6 ] );
530			break;
531		case 15:	// 3 complete vectors + remainder, 2 rows
532			_mm256_storeu_ps( rdi, arr[ 0 ] );
533			_mm256_storeu_ps( rdi + 8, arr[ 1 ] );
534			_mm256_storeu_ps( rdi + 8 * 2, arr[ 2 ] );
535			_mm256_maskstore_ps( rdi + 8 * 3, mask, arr[ 3 ] );
536
537			_mm256_storeu_ps( rdi1, arr[ 4 ] );
538			_mm256_storeu_ps( rdi1 + 8, arr[ 5 ] );
539			_mm256_storeu_ps( rdi1 + 8 * 2, arr[ 6 ] );
540			_mm256_maskstore_ps( rdi1 + 8 * 3, mask, arr[ 7 ] );
541			break;
542		default:
543			throw E_UNEXPECTED;
544		}
545	}
546}
547
548__forceinline void ResultTile<2, 4>::store( float* rdi, size_t w, size_t h, size_t stride ) const
549{
550	assert( h > 0 && w > 0 && h <= 4 && w <= 16 );
551	h--;
552	float* const rdi1 = rdi + stride;
553	float* const rdi2 = rdi + stride * 2;
554	float* const rdi3 = rdi + stride * 3;
555
556	if( w == 16 )
557	{
558		switch( h )
559		{
560		case 3:
561			_mm256_storeu_ps( rdi3, arr[ 6 ] );
562			_mm256_storeu_ps( rdi3 + 8, arr[ 7 ] );
563		case 2:
564			_mm256_storeu_ps( rdi2, arr[ 4 ] );
565			_mm256_storeu_ps( rdi2 + 8, arr[ 5 ] );
566		case 1:
567			_mm256_storeu_ps( rdi1, arr[ 2 ] );
568			_mm256_storeu_ps( rdi1 + 8, arr[ 3 ] );
569		case 0:
570			_mm256_storeu_ps( rdi, arr[ 0 ] );
571			_mm256_storeu_ps( rdi + 8, arr[ 1 ] );
572		}
573	}
574	else
575	{
576		const size_t rem = w % 8;
577		const __m256i mask = loadTailMaskInt<false>( rem );
578		// 0 for partial first vector, 1 for exactly 1 complete vector, 2 for 1 complete vector with remainder
579		const size_t partialCase = ( w < 8 ) ? 0 : ( ( w == 8 ) ? 1 : 2 );
580		// Merge into a single integer for the switch statement
581		const size_t key = partialCase + h * 3;
582
583		switch( key )
584		{
585			// h = 1
586		case 0:
587			_mm256_maskstore_ps( rdi, mask, arr[ 0 ] );
588			break;
589		case 1:
590			_mm256_storeu_ps( rdi, arr[ 0 ] );
591			break;
592		case 2:
593			_mm256_storeu_ps( rdi, arr[ 0 ] );
594			_mm256_maskstore_ps( rdi + 8, mask, arr[ 1 ] );
595			break;
596			// h = 2
597		case 3:
598			_mm256_maskstore_ps( rdi, mask, arr[ 0 ] );
599			_mm256_maskstore_ps( rdi1, mask, arr[ 2 ] );
600			break;
601		case 4:
602			_mm256_storeu_ps( rdi, arr[ 0 ] );
603			_mm256_storeu_ps( rdi1, arr[ 2 ] );
604			break;
605		case 5:
606			_mm256_storeu_ps( rdi, arr[ 0 ] );
607			_mm256_maskstore_ps( rdi + 8, mask, arr[ 1 ] );
608			_mm256_storeu_ps( rdi1, arr[ 2 ] );
609			_mm256_maskstore_ps( rdi1 + 8, mask, arr[ 3 ] );
610			break;
611			// h = 3
612		case 6:
613			_mm256_maskstore_ps( rdi, mask, arr[ 0 ] );
614			_mm256_maskstore_ps( rdi1, mask, arr[ 2 ] );
615			_mm256_maskstore_ps( rdi2, mask, arr[ 4 ] );
616			break;
617		case 7:
618			_mm256_storeu_ps( rdi, arr[ 0 ] );
619			_mm256_storeu_ps( rdi1, arr[ 2 ] );
620			_mm256_storeu_ps( rdi2, arr[ 4 ] );
621			break;
622		case 8:
623			_mm256_storeu_ps( rdi, arr[ 0 ] );
624			_mm256_maskstore_ps( rdi + 8, mask, arr[ 1 ] );
625			_mm256_storeu_ps( rdi1, arr[ 2 ] );
626			_mm256_maskstore_ps( rdi1 + 8, mask, arr[ 3 ] );
627			_mm256_storeu_ps( rdi2, arr[ 4 ] );
628			_mm256_maskstore_ps( rdi2 + 8, mask, arr[ 5 ] );
629			break;
630			// h = 4
631		case 9:
632			_mm256_maskstore_ps( rdi, mask, arr[ 0 ] );
633			_mm256_maskstore_ps( rdi1, mask, arr[ 2 ] );
634			_mm256_maskstore_ps( rdi2, mask, arr[ 4 ] );
635			_mm256_maskstore_ps( rdi3, mask, arr[ 6 ] );
636			break;
637		case 10:
638			_mm256_storeu_ps( rdi, arr[ 0 ] );
639			_mm256_storeu_ps( rdi1, arr[ 2 ] );
640			_mm256_storeu_ps( rdi2, arr[ 4 ] );
641			_mm256_storeu_ps( rdi3, arr[ 6 ] );
642			break;
643		case 11:
644			_mm256_storeu_ps( rdi, arr[ 0 ] );
645			_mm256_maskstore_ps( rdi + 8, mask, arr[ 1 ] );
646			_mm256_storeu_ps( rdi1, arr[ 2 ] );
647			_mm256_maskstore_ps( rdi1 + 8, mask, arr[ 3 ] );
648			_mm256_storeu_ps( rdi2, arr[ 4 ] );
649			_mm256_maskstore_ps( rdi2 + 8, mask, arr[ 5 ] );
650			_mm256_storeu_ps( rdi3, arr[ 6 ] );
651			_mm256_maskstore_ps( rdi3 + 8, mask, arr[ 7 ] );
652			break;
653		default:
654			throw E_UNEXPECTED;
655		}
656	}
657}
658
659__forceinline void ResultTile<2, 3>::store( float* rdi, size_t w, size_t h, size_t stride ) const
660{
661	assert( h > 0 && w > 0 && h <= 3 && w <= 16 );
662	float* const rdi1 = rdi + stride;
663	float* const rdi2 = rdi + stride * 2;
664	h--;
665
666	if( w == 16 )
667	{
668		switch( h )
669		{
670		case 2:
671			_mm256_storeu_ps( rdi2, arr[ 4 ] );
672			_mm256_storeu_ps( rdi2 + 8, arr[ 5 ] );
673		case 1:
674			_mm256_storeu_ps( rdi1, arr[ 2 ] );
675			_mm256_storeu_ps( rdi1 + 8, arr[ 3 ] );
676		case 0:
677			_mm256_storeu_ps( rdi, arr[ 0 ] );
678			_mm256_storeu_ps( rdi + 8, arr[ 1 ] );
679		}
680	}
681	else
682	{
683		const size_t rem = w % 8;
684		const __m256i mask = loadTailMaskInt<false>( rem );
685		// 0 for partial first vector, 1 for exactly 1 complete vector, 2 for 1 complete vector with remainder
686		const size_t partialCase = ( w < 8 ) ? 0 : ( ( w == 8 ) ? 1 : 2 );
687		// Merge into a single integer for the switch statement
688		const size_t key = partialCase + h * 3;
689
690		switch( key )
691		{
692			// h = 1
693		case 0:
694			_mm256_maskstore_ps( rdi, mask, arr[ 0 ] );
695			break;
696		case 1:
697			_mm256_storeu_ps( rdi, arr[ 0 ] );
698			break;
699		case 2:
700			_mm256_storeu_ps( rdi, arr[ 0 ] );
701			_mm256_maskstore_ps( rdi + 8, mask, arr[ 1 ] );
702			break;
703			// h = 2
704		case 3:
705			_mm256_maskstore_ps( rdi, mask, arr[ 0 ] );
706			_mm256_maskstore_ps( rdi1, mask, arr[ 2 ] );
707			break;
708		case 4:
709			_mm256_storeu_ps( rdi, arr[ 0 ] );
710			_mm256_storeu_ps( rdi1, arr[ 2 ] );
711			break;
712		case 5:
713			_mm256_storeu_ps( rdi, arr[ 0 ] );
714			_mm256_maskstore_ps( rdi + 8, mask, arr[ 1 ] );
715			_mm256_storeu_ps( rdi1, arr[ 2 ] );
716			_mm256_maskstore_ps( rdi1 + 8, mask, arr[ 3 ] );
717			break;
718			// h = 3
719		case 6:
720			_mm256_maskstore_ps( rdi, mask, arr[ 0 ] );
721			_mm256_maskstore_ps( rdi1, mask, arr[ 2 ] );
722			_mm256_maskstore_ps( rdi2, mask, arr[ 4 ] );
723			break;
724		case 7:
725			_mm256_storeu_ps( rdi, arr[ 0 ] );
726			_mm256_storeu_ps( rdi1, arr[ 2 ] );
727			_mm256_storeu_ps( rdi2, arr[ 4 ] );
728			break;
729		case 8:
730			_mm256_storeu_ps( rdi, arr[ 0 ] );
731			_mm256_maskstore_ps( rdi + 8, mask, arr[ 1 ] );
732			_mm256_storeu_ps( rdi1, arr[ 2 ] );
733			_mm256_maskstore_ps( rdi1 + 8, mask, arr[ 3 ] );
734			_mm256_storeu_ps( rdi2, arr[ 4 ] );
735			_mm256_maskstore_ps( rdi2 + 8, mask, arr[ 5 ] );
736			break;
737		default:
738			throw E_UNEXPECTED;
739		}
740	}
741}
742#pragma endregion