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#include "stdafx.h" 2#include "mulMatImpl.h" 3#include <immintrin.h> 4#include "mulMatUtils.hpp" 5using namespace CpuCompute ; 6 7namespace 8{ 9constexpr size_t prefetchBytes = 96 ; 10constexpr int prefetchHint = _MM_HINT_T0 ; 11 12constexpr size_t maskAlign16 = ~(size_t )15 ; 13 14 __forceinline__m256i load (const void * rsi ) 15 { 16return _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 __forceinlinevoid 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 __forceinlinevoid 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 __forceinlinevoid prefetch (const uint8_t * p ) 87 { 88_mm_prefetch ( (const char * )p ,prefetchHint ); 89 } 90 91 __forceinlinevoid transpose8Avx2 (uint16_t * rdiWords ,size_t w ,const uint16_t * rsiWords ,size_t sourceStride ,size_t destStride ) 92 { 93assert (0 == ( (size_t )rdiWords ) %16 ); 94assert (0 == destStride %8 ); 95assert (w <=sourceStride ); 96 97// Scale strides to bytes, and cast the pointers 98sourceStride *=2 ; 99destStride *=2 ; 100uint8_t * rdi = (uint8_t * )rdiWords ; 101const uint8_t * rsi = (const uint8_t * )rsiWords ; 102 103const uint8_t * const rsiEndAligned = rsi + (w & maskAlign16 )* 2 ; 104const uint8_t * const rsiEnd = rsi + w * 2 ; 105const uint8_t * rsiMid = rsi + sourceStride * 3 ; 106const uint8_t * rsiLast = rsi + sourceStride * 6 ; 107uint8_t * rdiMid = rdi + destStride * 3 ; 108uint8_t * rdiLast = rdi + destStride * 6 ; 109 110while (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 ); 116rsi += 32 ; 117__m256i r3 = load (rsiMid ); 118__m256i r4 = load (rsiMid + sourceStride ); 119__m256i r5 = load (rsiMid + sourceStride * 2 ); 120rsiMid += 32 ; 121__m256i r6 = load (rsiLast ); 122__m256i r7 = load (rsiLast + sourceStride ); 123rsiLast += 32 ; 124 125// Transpose FP16 values in registers 126TRANSPOSE_8X16 (); 127 128// Store 129STORE_8X16_LOW (); 130STORE_8X16_HIGH (); 131 132if constexpr (prefetchBytes > 0 ) 133 { 134if (rsi + prefetchBytes < rsiEnd ) 135 { 136prefetch (rsi + prefetchBytes ); 137prefetch (rsi + sourceStride + prefetchBytes ); 138prefetch (rsi + sourceStride * 2 + prefetchBytes ); 139prefetch (rsiMid + prefetchBytes ); 140prefetch (rsiMid + sourceStride + prefetchBytes ); 141prefetch (rsiMid + sourceStride * 2 + prefetchBytes ); 142prefetch (rsiLast + prefetchBytes ); 143prefetch (rsiLast + sourceStride + prefetchBytes ); 144 } 145 } 146 } 147 148if (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 ) 165while (rsi < rsiEnd ) 166 { 167r0 = _mm_cvtsi32_si128 (* (const uint16_t * )rsi ); 168r1 = _mm_insert_epi16 (r1 ,* (const int16_t * )(rsi + sourceStride ),1 ); 169r2 = _mm_insert_epi16 (r2 ,* (const int16_t * )(rsi + sourceStride * 2 ),2 ); 170rsi += 2 ; 171r3 = _mm_insert_epi16 (r3 ,* (const int16_t * )(rsiMid ),3 ); 172r4 = _mm_insert_epi16 (r4 ,* (const int16_t * )(rsiMid + sourceStride ),4 ); 173r5 = _mm_insert_epi16 (r5 ,* (const int16_t * )(rsiMid + sourceStride * 2 ),5 ); 174rsiMid += 2 ; 175r6 = _mm_insert_epi16 (r6 ,* (const int16_t * )(rsiLast ),6 ); 176r7 = _mm_insert_epi16 (r7 ,* (const int16_t * )(rsiLast + sourceStride ),7 ); 177rsiLast += 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 181t0 = _mm_or_si128 (r0 ,r1 ); 182t1 = _mm_or_si128 (r2 ,r3 ); 183t2 = _mm_or_si128 (r4 ,r5 ); 184t3 = _mm_or_si128 (r6 ,r7 ); 185 186t4 = _mm_or_si128 (t0 ,t1 ); 187t5 = _mm_or_si128 (t2 ,t3 ); 188 189t6 = _mm_or_si128 (t4 ,t5 ); 190// Store 8 FP16 values, the destination is aligned 191_mm_store_si128 ( (__m128i * )rdi ,t6 ); 192rdi += destStride ; 193 } 194 } 195 } 196 197 __forceinlinevoid transpose8PartialAvx2 (uint16_t * rdiWords ,size_t w ,size_t h ,const uint16_t * rsiWords ,size_t sourceStride ,size_t destStride ) 198 { 199assert (0 == ( (size_t )rdiWords ) %16 ); 200assert (0 == destStride %8 ); 201assert (w <=sourceStride ); 202assert (h > 0 && h < 8 ); 203 204// Scale strides to bytes, and cast the pointers 205sourceStride *=2 ; 206destStride *=2 ; 207uint8_t * rdi = (uint8_t * )rdiWords ; 208const uint8_t * rsi = (const uint8_t * )rsiWords ; 209 210const uint8_t * const rsiEndAligned = rsi + (w & maskAlign16 )* 2 ; 211const uint8_t * const rsiEnd = rsi + w * 2 ; 212const uint8_t * rsiMid = rsi + sourceStride * 3 ; 213const uint8_t * rsiLast = rsi + sourceStride * 6 ; 214uint8_t * rdiMid = rdi + destStride * 3 ; 215uint8_t * rdiLast = rdi + destStride * 6 ; 216 217while (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 228switch (h ) 229 { 230case 7 : 231r6 = load (rsiLast ); 232case 6 : 233r5 = load (rsiMid + sourceStride * 2 ); 234case 5 : 235r4 = load (rsiMid + sourceStride ); 236case 4 : 237r3 = load (rsiMid ); 238case 3 : 239r2 = load (rsi + sourceStride * 2 ); 240case 2 : 241r1 = load (rsi + sourceStride ); 242 } 243rsi += 32 ; 244rsiMid += 32 ; 245rsiLast += 32 ; 246 247__m256i r7 = _mm256_setzero_si256 (); 248 249// Transpose FP16 values in registers 250TRANSPOSE_8X16 (); 251 252// Store 253STORE_8X16_LOW (); 254 255STORE_8X16_HIGH (); 256 } 257 258if (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 ) 274while (rsi < rsiEnd ) 275 { 276r0 = _mm_cvtsi32_si128 (* (const uint16_t * )rsi ); 277 278switch (h ) 279 { 280case 7 : 281r6 = _mm_insert_epi16 (r6 ,* (const int16_t * )(rsiLast ),6 ); 282case 6 : 283r5 = _mm_insert_epi16 (r5 ,* (const int16_t * )(rsiMid + sourceStride * 2 ),5 ); 284case 5 : 285r4 = _mm_insert_epi16 (r4 ,* (const int16_t * )(rsiMid + sourceStride ),4 ); 286case 4 : 287r3 = _mm_insert_epi16 (r3 ,* (const int16_t * )(rsiMid ),3 ); 288case 3 : 289r2 = _mm_insert_epi16 (r2 ,* (const int16_t * )(rsi + sourceStride * 2 ),2 ); 290case 2 : 291r1 = _mm_insert_epi16 (r1 ,* (const int16_t * )(rsi + sourceStride ),1 ); 292 } 293rsi += 2 ; 294rsiMid += 2 ; 295rsiLast += 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 299t0 = _mm_or_si128 (r0 ,r1 ); 300t1 = _mm_or_si128 (r2 ,r3 ); 301t2 = _mm_or_si128 (r4 ,r5 ); 302 303t3 = _mm_or_si128 (t0 ,t1 ); 304t4 = _mm_or_si128 (t2 ,r6 ); 305 306t5 = _mm_or_si128 (t3 ,t4 ); 307// Store 8 FP16 values, the destination is aligned 308_mm_store_si128 ( (__m128i * )rdi ,t5 ); 309rdi += 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{ 319assert (stridesA [0 ]== 1 ); 320 321const size_t heightFloats = (size_t )panelHeightRegisters * 8 ; 322i *=heightFloats ; 323 324const uint16_t * rsi = (const uint16_t * )pa ; 325rsi += m3 * stridesA [3 ]; 326rsi += m2 * stridesA [2 ]; 327rsi += i * stridesA [1 ]; 328 329const size_t resultStride = heightFloats ; 330 331if (i + heightFloats <=resultSize [0 ] ) 332 { 333// A complete panel 334for (size_t i = 0 ;i < panelHeightRegisters ;i ++ ) 335 { 336transpose8Avx2 (rdi ,length ,rsi ,stridesA [1 ],resultStride ); 337// Advance by 8 floats in the output buffer 338rdi += 8 ; 339// Advance by 8 rows in the source matrix 340rsi += 8 * stridesA [1 ]; 341 } 342 } 343else 344 { 345// A partial panel, at the bottom of the first argument matrix 346const size_t remainder = resultSize [0 ]- i ; 347assert (remainder > 0 && remainder < heightFloats ); 348zeroAlignedMemory (rdi ,resultStride * length * sizeof (uint16_t ) ); 349 350const size_t completePanels = remainder /8 ; 351for (size_t i = 0 ;i < completePanels ;i ++ ) 352 { 353transpose8Avx2 (rdi ,length ,rsi ,stridesA [1 ],resultStride ); 354rdi += 8 ; 355rsi += 8 * stridesA [1 ]; 356 } 357const size_t lastPanel = remainder %8 ; 358if (0 != lastPanel ) 359transpose8PartialAvx2 (rdi ,length ,lastPanel ,rsi ,stridesA [1 ],resultStride ); 360 } 361return S_OK ; 362}