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 "Tensor.h" 3 4namespace DirectCompute 5{ 6using pfnNewCapacity = uint32_t (* )(uint32_t current ,uint32_t requested ); 7 8uint32_t defaultNewCapacity (uint32_t current ,uint32_t requested ); 9 10class PooledTensor 11 { 12TensorGpuViews views ; 13uint32_t capacity = 0 ; 14public : 15Tensor tensor (eDataType type ,const std ::array < uint32_t ,4 >& ne ,pfnNewCapacity pfnNewCap ); 16size_t getCapacity ()const {return capacity ; } 17void clear () 18 { 19views .clear (); 20capacity = 0 ; 21 } 22HRESULT zeroMemory (CComPtr < ID3D11Buffer >& cb ); 23 }; 24 25__interface iTensorArena 26 { 27Tensor tensor (eDataType type ,const std ::array < uint32_t ,4 >& ne ); 28void reset (); 29 }; 30 31class TensorsArena :public iTensorArena 32 { 33public : 34struct sArenaConfig 35 { 36pfnNewCapacity pfnCapInner ; 37size_t initialCapOuter ; 38 }; 39 40struct sArenaConfigs 41{ 42sArenaConfig fp16 , fp32 ; 43}; 44 45TensorsArena( const sArenaConfigs & configs ); 46 47Tensor tensor ( eDataType type, const std ::array < uint32_t , 4 >& ne ) override final; 48void reset () override final; 49 50void clear (); 51__m128i getMemoryUse () const; 52HRESULT zeroMemory ( CComPtr < ID3D11Buffer >& cb ); 53 54private : 55 56struct ArenaImpl 57{ 58ArenaImpl( eDataType dataType , const sArenaConfig & config ); 59 60void reset () 61{ 62index = 0 ; 63} 64 65void clear () 66{ 67index = 0 ; 68pool . clear (); 69} 70 71Tensor tensor ( const std ::array < uint32_t , 4 >& ne ); 72__m128i getMemoryUse () const; 73HRESULT zeroMemory ( CComPtr < ID3D11Buffer >& cb ); 74 75private : 76 77const eDataType type; 78const pfnNewCapacity pfnNewCap; 79 80std ::vector < PooledTensor > pool; 81size_t index = 0 ; 82}; 83 84static constexpr size_t countTypes = 2 ; 85std ::array < ArenaImpl, countTypes > arenas ; 86 }; 87}