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