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
3ba8e63
master
1#pragma once 2#include <vector> 3#include "TempBuffers.h" 4#include "ConstantBuffer.h" 5#include "Tensor.h" 6#include "../Utils/GpuProfiler.h" 7#include "../Utils/ProfileCollection.h" 8 9namespace DirectCompute 10{ 11enum struct eComputeShader :uint16_t ; 12 13class MlContext 14 { 15// When false, the implementation is 100% compatible with the CPU-running code written by Georgi Gerganov 16// When true, the implementation is much faster, and doesn't require FP64 support in the compute shaders. 17// FP64 is an optional feature, not all GPUs support that. 18static constexprbool enableInexactOptimizations = true; 19 20ConstantBuffer cb ; 21TempBuffers temp ; 22CComPtr < ID3D11Buffer > flashAttentionConstants ; 23 24void convolutionImpl (const Tensor & a ,const Tensor & b ,Tensor & res ,bool is2 ); 25 26void cwiseBinary (const Tensor & a ,const Tensor & b ,Tensor & res ,eComputeShader cs ); 27Tensor cwiseBinary (const Tensor & a ,const Tensor & b ,eComputeShader cs ); 28 29void mulMatDot (const Tensor & a ,const Tensor & b ,Tensor & res ); 30void mulMatMad (const Tensor & a ,const Tensor & b ,Tensor & res ); 31void mulMatTiled (const Tensor & a ,const Tensor & b ,Tensor & res ); 32 33void bindShader (eComputeShader cs ); 34 35protected : 36void copyImpl (const Tensor & a ,Tensor & res ,bool downcastFp32 ); 37 38// Create a dense output tensor for the results of a computation 39// Override this method to implement a pool of these tensors 40virtual Tensor createTensor (eDataType type ,const std ::array < uint32_t ,4 >& ne ); 41 42Tensor createTensor (eDataType type ,std ::initializer_list < uint32_t > ne ); 43 44GpuProfiler profiler ; 45 46CComPtr < ID3D11Buffer >& getSmallConstantBuffer () {return temp .smallCb ; } 47 48public : 49MlContext (Whisper ::ProfileCollection & profileColl ); 50MlContext (const MlContext & )= delete ; 51 52// res = a * b 53void mulMat (const Tensor & a ,const Tensor & b ,Tensor & res ); 54 55void flashAttention (const Tensor & q ,const Tensor & k ,const Tensor & v ,Tensor & res ,bool masked ); 56 57inline void convolution (const Tensor & a ,const Tensor & b ,Tensor & res ) 58 { 59convolutionImpl (a ,b ,res , false ); 60 } 61void convolution2 (const Tensor & a ,const Tensor & b ,Tensor & res ) 62 { 63convolutionImpl (a ,b ,res , true ); 64 } 65 66void norm (const Tensor & a ,Tensor & res ); 67 68Tensor conv_1d_1s (const Tensor & a ,const Tensor & b ); 69Tensor conv_1d_2s (const Tensor & a ,const Tensor & b ); 70 71Tensor add (const Tensor & a ,const Tensor & b ); 72void addInPlace (Tensor & a ,const Tensor & b ); 73 74Tensor view2d (const Tensor & a ,uint32_t ne0 ,uint32_t ne1 ,uint32_t nb1 ,uint32_t offset ); 75Tensor transpose (const Tensor & a ); 76 77Tensor norm (const Tensor & a ); 78Tensor mulMat (const Tensor & a ,const Tensor & b ); 79Tensor mulMatEx (const Tensor & a ,const Tensor & b ,const char * tagName ); 80Tensor permute (const Tensor & a ,uint8_t axis0 ,uint8_t axis1 ,uint8_t axis2 ,uint8_t axis3 ); 81Tensor flashAttention (const Tensor & q ,const Tensor & k ,const Tensor & v ,bool masked ); 82 83Tensor copy (const Tensor & a ,eDataType type ,std ::initializer_list < uint32_t > size ); 84void copyInPlace (Tensor & dest ,const Tensor & a ,eDataType type ,std ::initializer_list < uint32_t > size ); 85 86void dbgPrintDifference (const ggml_tensor * reference ,const Tensor & gpu ,const char * what ,bool trapToDebugger = true ); 87 88void scale (Tensor & a ,float mul ); 89 90void addRepeat (Tensor & a ,const Tensor & b ); 91void addRepeatScale (Tensor & a ,const Tensor & b ,float scale ); 92void fmaRepeat (Tensor & a ,const Tensor & mul ,const Tensor & add ); 93 94// ggml_diag_mask_inf 95void diagMaskInf (Tensor & a ,uint32_t n_past ); 96// ggml_soft_max 97void softMax (Tensor & a ,float inputScale = 1.0f ); 98 99void addRepeatGelu (Tensor & a ,const Tensor & b ); 100 101// Extract rows from tokenEmbedding matrix, row indices are taken from the `embd` R32_UINT row vector 102// Extract same count of rows from positionalEmbedding matrix, starting at the `pastTokensCount` row 103// Return a new FP32 matrix with the sum of these rows 104Tensor addRows (const Tensor & tokenEmbedding ,const Tensor & positionalEmbedding ,const Tensor & embd ,uint32_t pastTokensCount ); 105 106Tensor reshapePanels (const Tensor & a ); 107 108Tensor mulMatTiledEx (const Tensor & a ,const Tensor & b ); 109Tensor mulMatByRowTiledEx (const Tensor & a ,const Tensor & b ); 110 111// An equivalent of addRepeat( dest, pattern ) followed by addInPlace( dest, finalAdd ) 112void addRepeatEx (Tensor & dest ,const Tensor & pattern ,const Tensor & finalAdd ); 113 114__m128i getMemoryUse ()const ; 115 }; 116}