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 "TensorEx.h" 3#include "../D3D/createBuffer.h" 4#include "../source/ggml.h" 5#include "../D3D/MappedResource.h" 6using namespace DirectCompute ; 7 8HRESULT TensorEx ::create (const ggml_tensor & ggml ,eBufferUse usage ,bool uploadData ) 9{ 10TensorGpuViews ::clear (); 11buffer = nullptr ; 12stagingBuffer = nullptr ; 13 14CHECK (TensorShape ::create (ggml ) ); 15const ggml_type dataType = ggml .type ; 16const uint32_t cbElement = (uint32_t )ggml_type_size (dataType ); 17 18const size_t totalBytes = ggml_nbytes (& ggml ); 19if (totalBytes > INT_MAX ) 20return DISP_E_OVERFLOW ; 21const uint32_t countElements = (uint32_t )(totalBytes /cbElement ); 22 23 { 24const void * const rsi = uploadData ?ggml .data :nullptr ; 25ID3D11Buffer ** ppStagingBuffer = (usage == eBufferUse::ReadWriteDownload ) ?& stagingBuffer :nullptr ; 26CHECK (createBuffer (usage ,totalBytes ,& buffer ,rsi ,ppStagingBuffer ) ); 27 } 28 29DXGI_FORMAT format ; 30switch (dataType ) 31 { 32case GGML_TYPE_F16 : 33format = DXGI_FORMAT_R16_FLOAT ; 34break ; 35case GGML_TYPE_F32 : 36format = DXGI_FORMAT_R32_FLOAT ; 37break ; 38default : 39return E_NOTIMPL ; 40 } 41 42const bool makeUav = usage == eBufferUse::ReadWrite || usage == eBufferUse::ReadWriteDownload ; 43return TensorGpuViews ::create (buffer ,format ,totalBytes /cbElement ,makeUav ); 44} 45 46HRESULT TensorEx ::create (eDataType type ,eBufferUse usage ,const std::array < uint32_t ,4 >& sizeElements ) 47{ 48TensorGpuViews ::clear (); 49buffer = nullptr ; 50stagingBuffer = nullptr ; 51 std::initializer_list < uint32_t > il (sizeElements .data (),sizeElements .data ()+ 4 ); 52 53ID3D11Buffer ** ppStaging = (usage == eBufferUse::ReadWriteDownload ) ?& stagingBuffer :nullptr ; 54return Tensor ::create (type ,il ,usage ,buffer ,nullptr ,ppStaging ); 55} 56 57HRESULT TensorEx ::getViewSize (uint32_t & cbElement ,uint32_t & countElements )const 58{ 59ID3D11ShaderResourceView * const srv = * this ; 60if (nullptr == srv ) 61return OLE_E_BLANK ; 62 63D3D11_SHADER_RESOURCE_VIEW_DESC viewDesc ; 64srv -> GetDesc (& viewDesc ); 65 66cbElement = dxgiSizeof (viewDesc .Format ); 67 68assert (viewDesc .ViewDimension == D3D_SRV_DIMENSION_BUFFER ); 69assert (viewDesc .Buffer .FirstElement == 0 ); 70countElements = viewDesc .Buffer .NumElements ; 71 72return S_OK ; 73} 74 75HRESULT TensorEx ::download (void * rdi ,size_t cb )const 76{ 77if (nullptr == stagingBuffer ) 78return HRESULT_FROM_WIN32 (ERROR_GPIO_OPERATION_DENIED );// The requested operation is not supported for the specified handle. 79 80ID3D11DeviceContext * const ctx = context (); 81ctx -> CopyResource (stagingBuffer ,buffer ); 82 83MappedResource mapped ; 84CHECK (mapped .map (stagingBuffer , true ) ); 85memcpy (rdi ,mapped .data (),cb ); 86 87return S_OK ; 88} 89 90HRESULT TensorEx ::download (void * rdi )const 91{ 92uint32_t cbElement ,numElements ; 93CHECK (getViewSize (cbElement ,numElements ) ); 94 95size_t cb = (size_t )cbElement * numElements ; 96return download (rdi ,cb ); 97}