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
8c4603c
master
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{ 9static constexpr size_t totalRegs = (size_t )(tileWidthFloats )* panelHeightRegs ; 10 std::array < __m256 ,totalRegs > arr ; 11 12template < size_t idx > 13 __forceinlinevoid fmadd (__m256 a ,__m256 b ) 14 { 15arr [idx ]= _mm256_fmadd_ps (a ,b ,arr [idx ] ); 16 } 17 __forceinlinevoid kernel (const std::array < __m256 ,panelHeightRegs >& panel ,const float * rsi ,size_t stride ); 18 __forceinlinevoid kernelPartial (const std::array < __m256 ,panelHeightRegs >& panel ,const float * rsi ,size_t stride ,size_t rem ) 19 { 20throw E_UNEXPECTED ; 21 } 22 __forceinlinevoid store (float * rdi ,size_t w ,size_t h ,size_t stride )const ; 23}; 24 25#pragma region setZero functions 26__forceinlinevoid setZero ( std::array < __m256 ,1 >& dest ) 27{ 28dest [0 ]= _mm256_setzero_ps (); 29} 30__forceinlinevoid setZero ( std::array < __m256 ,2 >& dest ) 31{ 32dest [0 ]= _mm256_setzero_ps (); 33dest [1 ]= _mm256_setzero_ps (); 34} 35__forceinlinevoid setZero ( std::array < __m256 ,3 >& dest ) 36{ 37dest [0 ]= _mm256_setzero_ps (); 38dest [1 ]= _mm256_setzero_ps (); 39dest [2 ]= _mm256_setzero_ps (); 40} 41__forceinlinevoid setZero ( std::array < __m256 ,4 >& dest ) 42{ 43dest [0 ]= _mm256_setzero_ps (); 44dest [1 ]= _mm256_setzero_ps (); 45dest [2 ]= _mm256_setzero_ps (); 46dest [3 ]= _mm256_setzero_ps (); 47} 48__forceinlinevoid setZero ( std::array < __m256 ,6 >& dest ) 49{ 50dest [0 ]= _mm256_setzero_ps (); 51dest [1 ]= _mm256_setzero_ps (); 52dest [2 ]= _mm256_setzero_ps (); 53dest [3 ]= _mm256_setzero_ps (); 54dest [4 ]= _mm256_setzero_ps (); 55dest [5 ]= _mm256_setzero_ps (); 56} 57__forceinlinevoid setZero ( std::array < __m256 ,8 >& dest ) 58{ 59dest [0 ]= _mm256_setzero_ps (); 60dest [1 ]= _mm256_setzero_ps (); 61dest [2 ]= _mm256_setzero_ps (); 62dest [3 ]= _mm256_setzero_ps (); 63dest [4 ]= _mm256_setzero_ps (); 64dest [5 ]= _mm256_setzero_ps (); 65dest [6 ]= _mm256_setzero_ps (); 66dest [7 ]= _mm256_setzero_ps (); 67} 68#pragma endregion 69 70#pragma region Micro-kernels 71__forceinlinevoid ResultTile < 1 ,1 > ::kernel (const std::array < __m256 ,1 >& panel ,const float * rsi ,size_t stride ) 72{ 73__m256 b = _mm256_broadcast_ss (rsi ); 74fmadd < 0 > (panel [0 ],b ); 75} 76__forceinlinevoid ResultTile < 1 ,2 > ::kernel (const std::array < __m256 ,1 >& panel ,const float * rsi ,size_t stride ) 77{ 78__m256 b = _mm256_broadcast_ss (rsi ); 79fmadd < 0 > (panel [0 ],b ); 80b = _mm256_broadcast_ss (rsi + stride ); 81fmadd < 1 > (panel [0 ],b ); 82} 83__forceinlinevoid ResultTile < 1 ,2 > ::kernelPartial (const std::array < __m256 ,1 >& panel ,const float * rsi ,size_t stride ,size_t rem ) 84{ 85assert (1 == rem ); 86__m256 b = _mm256_broadcast_ss (rsi ); 87fmadd < 0 > (panel [0 ],b ); 88} 89__forceinlinevoid ResultTile < 1 ,3 > ::kernel (const std::array < __m256 ,1 >& panel ,const float * rsi ,size_t stride ) 90{ 91__m256 b = _mm256_broadcast_ss (rsi ); 92fmadd < 0 > (panel [0 ],b ); 93b = _mm256_broadcast_ss (rsi + stride ); 94fmadd < 1 > (panel [0 ],b ); 95b = _mm256_broadcast_ss (rsi + stride * 2 ); 96fmadd < 2 > (panel [0 ],b ); 97} 98__forceinlinevoid ResultTile < 1 ,3 > ::kernelPartial (const std::array < __m256 ,1 >& panel ,const float * rsi ,size_t stride ,size_t rem ) 99{ 100assert (rem > 0 && rem < 3 ); 101__m256 b = _mm256_broadcast_ss (rsi ); 102fmadd < 0 > (panel [0 ],b ); 103if (rem > 1 ) 104 { 105b = _mm256_broadcast_ss (rsi + stride ); 106fmadd < 1 > (panel [0 ],b ); 107 } 108} 109 110__forceinlinevoid ResultTile < 1 ,4 > ::kernel (const std::array < __m256 ,1 >& panel ,const float * rsi ,size_t stride ) 111{ 112__m256 b = _mm256_broadcast_ss (rsi ); 113fmadd < 0 > (panel [0 ],b ); 114b = _mm256_broadcast_ss (rsi + stride ); 115fmadd < 1 > (panel [0 ],b ); 116b = _mm256_broadcast_ss (rsi + stride * 2 ); 117fmadd < 2 > (panel [0 ],b ); 118b = _mm256_broadcast_ss (rsi + stride * 3 ); 119fmadd < 3 > (panel [0 ],b ); 120} 121__forceinlinevoid ResultTile < 1 ,4 > ::kernelPartial (const std::array < __m256 ,1 >& panel ,const float * rsi ,size_t stride ,size_t rem ) 122{ 123assert (rem > 0 && rem < 4 ); 124__m256 b = _mm256_broadcast_ss (rsi ); 125fmadd < 0 > (panel [0 ],b ); 126 127switch (rem ) 128 { 129case 3 : 130b = _mm256_broadcast_ss (rsi + stride * 2 ); 131fmadd < 2 > (panel [0 ],b ); 132case 2 : 133b = _mm256_broadcast_ss (rsi + stride ); 134fmadd < 1 > (panel [0 ],b ); 135 } 136} 137__forceinlinevoid ResultTile < 4 ,1 > ::kernel (const std::array < __m256 ,4 >& panel ,const float * rsi ,size_t stride ) 138{ 139__m256 b = _mm256_broadcast_ss (rsi ); 140fmadd < 0 > (panel [0 ],b ); 141fmadd < 1 > (panel [1 ],b ); 142fmadd < 2 > (panel [2 ],b ); 143fmadd < 3 > (panel [3 ],b ); 144} 145__forceinlinevoid ResultTile < 2 ,4 > ::kernel (const std::array < __m256 ,2 >& panel ,const float * rsi ,size_t stride ) 146{ 147__m256 b = _mm256_broadcast_ss (rsi ); 148fmadd < 0 > (panel [0 ],b ); 149fmadd < 1 > (panel [1 ],b ); 150 151b = _mm256_broadcast_ss (rsi + stride ); 152fmadd < 2 > (panel [0 ],b ); 153fmadd < 3 > (panel [1 ],b ); 154 155b = _mm256_broadcast_ss (rsi + stride * 2 ); 156fmadd < 4 > (panel [0 ],b ); 157fmadd < 5 > (panel [1 ],b ); 158 159b = _mm256_broadcast_ss (rsi + stride * 3 ); 160fmadd < 6 > (panel [0 ],b ); 161fmadd < 7 > (panel [1 ],b ); 162} 163 164__forceinlinevoid ResultTile < 2 ,4 > ::kernelPartial (const std::array < __m256 ,2 >& panel ,const float * rsi ,size_t stride ,size_t rem ) 165{ 166assert (rem > 0 && rem < 4 ); 167__m256 b = _mm256_broadcast_ss (rsi ); 168fmadd < 0 > (panel [0 ],b ); 169fmadd < 1 > (panel [1 ],b ); 170 171switch (rem ) 172 { 173case 3 : 174b = _mm256_broadcast_ss (rsi + stride * 2 ); 175fmadd < 4 > (panel [0 ],b ); 176fmadd < 5 > (panel [1 ],b ); 177case 2 : 178b = _mm256_broadcast_ss (rsi + stride ); 179fmadd < 2 > (panel [0 ],b ); 180fmadd < 3 > (panel [1 ],b ); 181 } 182} 183 184__forceinlinevoid ResultTile < 2 ,3 > ::kernel (const std::array < __m256 ,2 >& panel ,const float * rsi ,size_t stride ) 185{ 186__m256 b = _mm256_broadcast_ss (rsi ); 187fmadd < 0 > (panel [0 ],b ); 188fmadd < 1 > (panel [1 ],b ); 189 190b = _mm256_broadcast_ss (rsi + stride ); 191fmadd < 2 > (panel [0 ],b ); 192fmadd < 3 > (panel [1 ],b ); 193 194b = _mm256_broadcast_ss (rsi + stride * 2 ); 195fmadd < 4 > (panel [0 ],b ); 196fmadd < 5 > (panel [1 ],b ); 197} 198__forceinlinevoid ResultTile < 2 ,3 > ::kernelPartial (const std::array < __m256 ,2 >& panel ,const float * rsi ,size_t stride ,size_t rem ) 199{ 200assert (rem > 0 && rem < 3 ); 201__m256 b = _mm256_broadcast_ss (rsi ); 202fmadd < 0 > (panel [0 ],b ); 203fmadd < 1 > (panel [1 ],b ); 204if (rem > 1 ) 205 { 206b = _mm256_broadcast_ss (rsi + stride ); 207fmadd < 2 > (panel [0 ],b ); 208fmadd < 3 > (panel [1 ],b ); 209 } 210} 211 212__forceinlinevoid ResultTile < 4 ,2 > ::kernel (const std::array < __m256 ,4 >& panel ,const float * rsi ,size_t stride ) 213{ 214__m256 b = _mm256_broadcast_ss (rsi ); 215fmadd < 0 > (panel [0 ],b ); 216fmadd < 1 > (panel [1 ],b ); 217fmadd < 2 > (panel [2 ],b ); 218fmadd < 3 > (panel [3 ],b ); 219 220b = _mm256_broadcast_ss (rsi + stride ); 221fmadd < 4 > (panel [0 ],b ); 222fmadd < 5 > (panel [1 ],b ); 223fmadd < 6 > (panel [2 ],b ); 224fmadd < 7 > (panel [3 ],b ); 225} 226__forceinlinevoid ResultTile < 4 ,2 > ::kernelPartial (const std::array < __m256 ,4 >& panel ,const float * rsi ,size_t stride ,size_t rem ) 227{ 228assert (1 == rem ); 229__m256 b = _mm256_broadcast_ss (rsi ); 230fmadd < 0 > (panel [0 ],b ); 231fmadd < 1 > (panel [1 ],b ); 232fmadd < 2 > (panel [2 ],b ); 233fmadd < 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 ); 242return _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__forceinlinevoid loadPanel (const uint16_t * rsi , std::array < __m256 ,1 >& dest ) 249{ 250dest [0 ]= loadUpcasted (rsi ); 251} 252__forceinlinevoid loadPanel (const uint16_t * rsi , std::array < __m256 ,2 >& dest ) 253{ 254dest [0 ]= loadUpcasted (rsi ); 255dest [1 ]= loadUpcasted (rsi + 8 ); 256} 257__forceinlinevoid loadPanel (const uint16_t * rsi , std::array < __m256 ,3 >& dest ) 258{ 259dest [0 ]= loadUpcasted (rsi ); 260dest [1 ]= loadUpcasted (rsi + 8 ); 261dest [2 ]= loadUpcasted (rsi + 8 * 2 ); 262} 263__forceinlinevoid loadPanel (const uint16_t * rsi , std::array < __m256 ,4 >& dest ) 264{ 265dest [0 ]= loadUpcasted (rsi ); 266dest [1 ]= loadUpcasted (rsi + 8 ); 267dest [2 ]= loadUpcasted (rsi + 8 * 2 ); 268dest [3 ]= loadUpcasted (rsi + 8 * 3 ); 269} 270#pragma endregion 271 272#pragma region Stores 273__forceinlinevoid ResultTile < 1 ,1 > ::store (float * rdi ,size_t w ,size_t h ,size_t stride )const 274{ 275assert (h == 1 && w > 0 && w <=8 ); 276if (w == 8 ) 277_mm256_storeu_ps (rdi ,arr [0 ] ); 278else 279 { 280const __m256i mask = loadTailMaskInt (w ); 281_mm256_maskstore_ps (rdi ,mask ,arr [0 ] ); 282 } 283} 284 285__forceinlinevoid ResultTile < 1 ,2 > ::store (float * rdi ,size_t w ,size_t h ,size_t stride )const 286{ 287assert (h > 0 && w > 0 && h <=2 && w <=8 ); 288if (w == 8 ) 289 { 290switch (h ) 291 { 292case 2 : 293_mm256_storeu_ps (rdi + stride ,arr [1 ] ); 294case 1 : 295_mm256_storeu_ps (rdi ,arr [0 ] ); 296 } 297 } 298else 299 { 300const __m256i mask = loadTailMaskInt (w ); 301switch (h ) 302 { 303case 2 : 304_mm256_maskstore_ps (rdi + stride ,mask ,arr [1 ] ); 305case 1 : 306_mm256_maskstore_ps (rdi ,mask ,arr [0 ] ); 307 } 308 } 309} 310 311__forceinlinevoid ResultTile < 1 ,3 > ::store (float * rdi ,size_t w ,size_t h ,size_t stride )const 312{ 313assert (h > 0 && w > 0 && h <=3 && w <=8 ); 314if (w == 8 ) 315 { 316switch (h ) 317 { 318case 3 : 319_mm256_storeu_ps (rdi + stride * 2 ,arr [2 ] ); 320case 2 : 321_mm256_storeu_ps (rdi + stride ,arr [1 ] ); 322case 1 : 323_mm256_storeu_ps (rdi ,arr [0 ] ); 324 } 325 } 326else 327 { 328const __m256i mask = loadTailMaskInt (w ); 329switch (h ) 330 { 331case 3 : 332_mm256_maskstore_ps (rdi + stride * 2 ,mask ,arr [2 ] ); 333case 2 : 334_mm256_maskstore_ps (rdi + stride ,mask ,arr [1 ] ); 335case 1 : 336_mm256_maskstore_ps (rdi ,mask ,arr [0 ] ); 337 } 338 } 339} 340 341__forceinlinevoid ResultTile < 1 ,4 > ::store (float * rdi ,size_t w ,size_t h ,size_t stride )const 342{ 343assert (h > 0 && w > 0 && h <=4 && w <=8 ); 344 345if (w == 8 ) 346 { 347switch (h ) 348 { 349case 4 : 350_mm256_storeu_ps (rdi + stride * 3 ,arr [3 ] ); 351case 3 : 352_mm256_storeu_ps (rdi + stride * 2 ,arr [2 ] ); 353case 2 : 354_mm256_storeu_ps (rdi + stride ,arr [1 ] ); 355case 1 : 356_mm256_storeu_ps (rdi ,arr [0 ] ); 357 } 358 } 359else 360 { 361const __m256i mask = loadTailMaskInt (w ); 362switch (h ) 363 { 364case 4 : 365_mm256_maskstore_ps (rdi + stride * 3 ,mask ,arr [3 ] ); 366case 3 : 367_mm256_maskstore_ps (rdi + stride * 2 ,mask ,arr [2 ] ); 368case 2 : 369_mm256_maskstore_ps (rdi + stride ,mask ,arr [1 ] ); 370case 1 : 371_mm256_maskstore_ps (rdi ,mask ,arr [0 ] ); 372 } 373 } 374} 375 376__forceinlinevoid ResultTile < 4 ,1 > ::store (float * rdi ,size_t w ,size_t h ,size_t stride )const 377{ 378assert (h == 1 && w > 0 && w <=32 ); 379if (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 } 387else 388 { 389const size_t rem = w %8 ; 390const __m256i mask = loadTailMaskInt < false> (rem ); 391const size_t completeVectors = w /8 ; 392const size_t key = (completeVectors <<1 ) | ( (0 == rem ) ?0 :1 ); 393switch (key ) 394 { 395case 1 :// 0 complete vectors + remainder 396_mm256_maskstore_ps (rdi ,mask ,arr [0 ] ); 397break ; 398case 2 :// 1 complete vector 399_mm256_storeu_ps (rdi ,arr [0 ] ); 400break ; 401case 3 :// 1 complete vector + remainder 402_mm256_storeu_ps (rdi ,arr [0 ] ); 403_mm256_maskstore_ps (rdi + 8 ,mask ,arr [1 ] ); 404break ; 405case 4 :// 2 complete vectors 406_mm256_storeu_ps (rdi ,arr [0 ] ); 407_mm256_storeu_ps (rdi + 8 ,arr [1 ] ); 408break ; 409case 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 ] ); 413break ; 414case 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 ] ); 418break ; 419case 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 ] ); 424break ; 425default : 426throw E_UNEXPECTED ; 427 } 428 } 429} 430__forceinlinevoid ResultTile < 4 ,2 > ::store (float * rdi ,size_t w ,size_t h ,size_t stride )const 431{ 432assert (h > 0 && w > 0 && h <=2 && w <=32 ); 433const bool twoRows = h == 2 ; 434float * const rdi1 = rdi + stride ; 435if (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 442if (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 } 450else 451 { 452const size_t rem = w %8 ; 453const __m256i mask = loadTailMaskInt < false> (rem ); 454const 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 ] 458const size_t key = (completeVectors <<2 ) | ( (0 == rem ) ?0 :1 ) | (twoRows ?2 :0 ); 459switch (key ) 460 { 461case 1 :// 0 complete vectors + remainder, 1 row 462_mm256_maskstore_ps (rdi ,mask ,arr [0 ] ); 463break ; 464case 3 :// 0 complete vectors + remainder, 2 rows 465_mm256_maskstore_ps (rdi ,mask ,arr [0 ] ); 466_mm256_maskstore_ps (rdi1 ,mask ,arr [4 ] ); 467break ; 468case 4 :// 1 complete vector, 1 row 469_mm256_storeu_ps (rdi ,arr [0 ] ); 470break ; 471case 5 :// 1 complete vector + remainder, 1 row 472_mm256_storeu_ps (rdi ,arr [0 ] ); 473_mm256_maskstore_ps (rdi + 8 ,mask ,arr [1 ] ); 474break ; 475case 6 :// 1 complete vector, 2 rows 476_mm256_storeu_ps (rdi ,arr [0 ] ); 477_mm256_storeu_ps (rdi1 ,arr [4 ] ); 478break ; 479case 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 ] ); 485break ; 486case 8 :// 2 complete vectors, 1 row 487_mm256_storeu_ps (rdi ,arr [0 ] ); 488_mm256_storeu_ps (rdi + 8 ,arr [1 ] ); 489break ; 490case 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 ] ); 494break ; 495case 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 ] ); 501break ; 502case 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 ] ); 510break ; 511case 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 ] ); 515break ; 516case 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 ] ); 521break ; 522case 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 ] ); 530break ; 531case 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 ] ); 541break ; 542default : 543throw E_UNEXPECTED ; 544 } 545 } 546} 547 548__forceinlinevoid ResultTile < 2 ,4 > ::store (float * rdi ,size_t w ,size_t h ,size_t stride )const 549{ 550assert (h > 0 && w > 0 && h <=4 && w <=16 ); 551h -- ; 552float * const rdi1 = rdi + stride ; 553float * const rdi2 = rdi + stride * 2 ; 554float * const rdi3 = rdi + stride * 3 ; 555 556if (w == 16 ) 557 { 558switch (h ) 559 { 560case 3 : 561_mm256_storeu_ps (rdi3 ,arr [6 ] ); 562_mm256_storeu_ps (rdi3 + 8 ,arr [7 ] ); 563case 2 : 564_mm256_storeu_ps (rdi2 ,arr [4 ] ); 565_mm256_storeu_ps (rdi2 + 8 ,arr [5 ] ); 566case 1 : 567_mm256_storeu_ps (rdi1 ,arr [2 ] ); 568_mm256_storeu_ps (rdi1 + 8 ,arr [3 ] ); 569case 0 : 570_mm256_storeu_ps (rdi ,arr [0 ] ); 571_mm256_storeu_ps (rdi + 8 ,arr [1 ] ); 572 } 573 } 574else 575 { 576const size_t rem = w %8 ; 577const __m256i mask = loadTailMaskInt < false> (rem ); 578// 0 for partial first vector, 1 for exactly 1 complete vector, 2 for 1 complete vector with remainder 579const size_t partialCase = (w < 8 ) ?0 : ( (w == 8 ) ?1 :2 ); 580// Merge into a single integer for the switch statement 581const size_t key = partialCase + h * 3 ; 582 583switch (key ) 584 { 585// h = 1 586case 0 : 587_mm256_maskstore_ps (rdi ,mask ,arr [0 ] ); 588break ; 589case 1 : 590_mm256_storeu_ps (rdi ,arr [0 ] ); 591break ; 592case 2 : 593_mm256_storeu_ps (rdi ,arr [0 ] ); 594_mm256_maskstore_ps (rdi + 8 ,mask ,arr [1 ] ); 595break ; 596// h = 2 597case 3 : 598_mm256_maskstore_ps (rdi ,mask ,arr [0 ] ); 599_mm256_maskstore_ps (rdi1 ,mask ,arr [2 ] ); 600break ; 601case 4 : 602_mm256_storeu_ps (rdi ,arr [0 ] ); 603_mm256_storeu_ps (rdi1 ,arr [2 ] ); 604break ; 605case 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 ] ); 610break ; 611// h = 3 612case 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 ] ); 616break ; 617case 7 : 618_mm256_storeu_ps (rdi ,arr [0 ] ); 619_mm256_storeu_ps (rdi1 ,arr [2 ] ); 620_mm256_storeu_ps (rdi2 ,arr [4 ] ); 621break ; 622case 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 ] ); 629break ; 630// h = 4 631case 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 ] ); 636break ; 637case 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 ] ); 642break ; 643case 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 ] ); 652break ; 653default : 654throw E_UNEXPECTED ; 655 } 656 } 657} 658 659__forceinlinevoid ResultTile < 2 ,3 > ::store (float * rdi ,size_t w ,size_t h ,size_t stride )const 660{ 661assert (h > 0 && w > 0 && h <=3 && w <=16 ); 662float * const rdi1 = rdi + stride ; 663float * const rdi2 = rdi + stride * 2 ; 664h -- ; 665 666if (w == 16 ) 667 { 668switch (h ) 669 { 670case 2 : 671_mm256_storeu_ps (rdi2 ,arr [4 ] ); 672_mm256_storeu_ps (rdi2 + 8 ,arr [5 ] ); 673case 1 : 674_mm256_storeu_ps (rdi1 ,arr [2 ] ); 675_mm256_storeu_ps (rdi1 + 8 ,arr [3 ] ); 676case 0 : 677_mm256_storeu_ps (rdi ,arr [0 ] ); 678_mm256_storeu_ps (rdi + 8 ,arr [1 ] ); 679 } 680 } 681else 682 { 683const size_t rem = w %8 ; 684const __m256i mask = loadTailMaskInt < false> (rem ); 685// 0 for partial first vector, 1 for exactly 1 complete vector, 2 for 1 complete vector with remainder 686const size_t partialCase = (w < 8 ) ?0 : ( (w == 8 ) ?1 :2 ); 687// Merge into a single integer for the switch statement 688const size_t key = partialCase + h * 3 ; 689 690switch (key ) 691 { 692// h = 1 693case 0 : 694_mm256_maskstore_ps (rdi ,mask ,arr [0 ] ); 695break ; 696case 1 : 697_mm256_storeu_ps (rdi ,arr [0 ] ); 698break ; 699case 2 : 700_mm256_storeu_ps (rdi ,arr [0 ] ); 701_mm256_maskstore_ps (rdi + 8 ,mask ,arr [1 ] ); 702break ; 703// h = 2 704case 3 : 705_mm256_maskstore_ps (rdi ,mask ,arr [0 ] ); 706_mm256_maskstore_ps (rdi1 ,mask ,arr [2 ] ); 707break ; 708case 4 : 709_mm256_storeu_ps (rdi ,arr [0 ] ); 710_mm256_storeu_ps (rdi1 ,arr [2 ] ); 711break ; 712case 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 ] ); 717break ; 718// h = 3 719case 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 ] ); 723break ; 724case 7 : 725_mm256_storeu_ps (rdi ,arr [0 ] ); 726_mm256_storeu_ps (rdi1 ,arr [2 ] ); 727_mm256_storeu_ps (rdi2 ,arr [4 ] ); 728break ; 729case 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 ] ); 736break ; 737default : 738throw E_UNEXPECTED ; 739 } 740 } 741} 742#pragma endregion