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 "Tensor.h" 3#include "../D3D/MappedResource.h" 4#include "../D3D/createBuffer.h" 5#include "../source/ggml.h" 6using namespace DirectCompute ; 7 8Tensor ::Tensor (const Tensor & that ) 9{ 10ne = that .ne ; 11nb = that .nb ; 12srv = that .srv ; 13uav = that .uav ; 14#ifdef _DEBUG 15dbgType = that .dbgType ; 16#endif 17} 18 19Tensor ::Tensor (Tensor && that )noexcept 20{ 21ne = that .ne ; 22nb = that .nb ; 23srv .Attach (that .srv .Detach () ); 24uav .Attach (that .uav .Detach () ); 25#ifdef _DEBUG 26dbgType = that .dbgType ; 27#endif 28} 29 30Tensor & Tensor ::operator= (const Tensor & that ) 31{ 32ne = that .ne ; 33nb = that .nb ; 34srv = that .srv ; 35uav = that .uav ; 36#ifdef _DEBUG 37dbgType = that .dbgType ; 38#endif 39return * this ; 40} 41 42Tensor & Tensor ::operator= (Tensor && that )noexcept 43{ 44ne = that .ne ; 45nb = that .nb ; 46srv .Attach (that .srv .Detach () ); 47uav .Attach (that .uav .Detach () ); 48#ifdef _DEBUG 49dbgType = that .dbgType ; 50#endif 51return * this ; 52} 53 54Tensor ::Tensor (const TensorShape & shape ,CComPtr < ID3D11ShaderResourceView >& srv ,CComPtr < ID3D11UnorderedAccessView >& uav )noexcept : 55TensorShape (shape ) 56{ 57TensorGpuViews ::srv .Attach (srv .Detach () ); 58TensorGpuViews ::uav .Attach (uav .Detach () ); 59} 60 61Tensor ::Tensor (const TensorShape & shape ,const TensorGpuViews & views ) : 62TensorShape (shape ) 63{ 64srv = views ; 65uav = views ; 66} 67 68HRESULT Tensor ::create (const ggml_tensor & ggml ,eBufferUse usage ,bool uploadData ) 69{ 70TensorGpuViews ::clear (); 71 72switch (usage ) 73 { 74case eBufferUse::Immutable : 75case eBufferUse::ReadWriteDownload : 76break ; 77default : 78return E_INVALIDARG ; 79 } 80 81CComPtr < ID3D11Buffer > buffer ; 82 83CHECK (TensorShape ::create (ggml ) ); 84const ggml_type dataType = ggml .type ; 85const uint32_t cbElement = (uint32_t )ggml_type_size (dataType ); 86 87const size_t totalBytes = ggml_nbytes (& ggml ); 88if (totalBytes > INT_MAX ) 89return DISP_E_OVERFLOW ; 90const uint32_t countElements = (uint32_t )(totalBytes /cbElement ); 91 92 { 93const void * const rsi = uploadData ?ggml .data :nullptr ; 94CHECK (createBuffer (usage ,totalBytes ,& buffer ,rsi ,nullptr ) ); 95 } 96 97DXGI_FORMAT format ; 98eDataType type ; 99switch (dataType ) 100 { 101case GGML_TYPE_F16 : 102format = DXGI_FORMAT_R16_FLOAT ; 103type = eDataType::FP16 ; 104break ; 105case GGML_TYPE_F32 : 106format = DXGI_FORMAT_R32_FLOAT ; 107type = eDataType::FP32 ; 108break ; 109default : 110return E_NOTIMPL ; 111 } 112 113const bool makeUav = (usage == eBufferUse::ReadWrite ); 114 115CHECK (TensorGpuViews ::create (buffer ,format ,totalBytes /cbElement ,makeUav ) ); 116#ifdef _DEBUG 117dbgType .type = type ; 118dbgType .usage = usage ; 119dbgType .hasInitialData = uploadData ; 120#endif 121return S_OK ; 122} 123 124HRESULT Tensor ::createImmutable (eDataType type ,const std::array < int ,4 >& size ,const void * rsi ) 125{ 126size_t elts = (uint32_t )size [0 ]; 127elts *= (uint32_t )size [1 ]; 128elts *= (uint32_t )size [2 ]; 129elts *= (uint32_t )size [3 ]; 130 131DXGI_FORMAT format ; 132size_t cbElement ; 133switch (type ) 134 { 135case eDataType::FP16 : 136format = DXGI_FORMAT_R16_FLOAT ; 137cbElement = 2 ; 138break ; 139case eDataType::FP32 : 140format = DXGI_FORMAT_R32_FLOAT ; 141cbElement = 4 ; 142break ; 143default : 144return E_NOTIMPL ; 145 } 146 147CComPtr < ID3D11Buffer > buffer ; 148CHECK (createBuffer ( eBufferUse::Immutable ,cbElement * elts ,& buffer ,rsi ,nullptr ) ); 149CHECK (TensorGpuViews ::create (buffer ,format ,elts , false ) ); 150 151__m128i v = _mm_loadu_si128 ( (const __m128i * )size .data () ); 152_mm_storeu_si128 ( (__m128i * )ne .data (),v ); 153setDenseStrides (); 154return S_OK ; 155} 156 157HRESULT Tensor ::create (eDataType type , std::initializer_list < uint32_t > sizeElements ,eBufferUse usage ,CComPtr < ID3D11Buffer >& buffer ,const void * rsi ,ID3D11Buffer ** ppStagingBuffer ) 158{ 159TensorGpuViews ::clear (); 160 161size_t nDims = sizeElements .size (); 162if (0 == nDims || nDims > 4 ) 163return E_INVALIDARG ; 164nDims = std::min (nDims , (size_t )4 ); 165size_t totalElements = 1 ; 166for (size_t i = 0 ;i < nDims ;i ++ ) 167 { 168uint32_t n = sizeElements .begin ()[i ]; 169if (n == 0 ) 170return E_INVALIDARG ; 171ne [i ]= n ; 172totalElements *=n ; 173 } 174 175DXGI_FORMAT format ; 176size_t cbElement ; 177switch (type ) 178 { 179case eDataType::FP32 : 180format = DXGI_FORMAT_R32_FLOAT ; 181cbElement = 4 ; 182break ; 183case eDataType::FP16 : 184format = DXGI_FORMAT_R16_FLOAT ; 185cbElement = 2 ; 186break ; 187case eDataType::U32 : 188format = DXGI_FORMAT_R32_UINT ; 189cbElement = 4 ; 190break ; 191default : 192return E_NOTIMPL ; 193 } 194 195const size_t totalBytes = cbElement * totalElements ; 196if (totalBytes > INT_MAX ) 197return DISP_E_OVERFLOW ; 198 199for (size_t i = nDims ;i < 4 ;i ++ ) 200ne [i ]= 1 ; 201TensorShape ::setDenseStrides (); 202 203CHECK (createBuffer (usage ,totalBytes ,& buffer ,rsi ,ppStagingBuffer ) ); 204 205CHECK (TensorGpuViews ::create (buffer ,format ,totalBytes /cbElement , true ) ); 206#ifdef _DEBUG 207dbgType .type = type ; 208dbgType .usage = usage ; 209dbgType .hasInitialData = (nullptr != rsi ); 210#endif 211return S_OK ; 212} 213 214HRESULT Tensor ::create (eDataType type , std::initializer_list < uint32_t > sizeElements ) 215{ 216CComPtr < ID3D11Buffer > buffer ; 217return create (type ,sizeElements , eBufferUse::ReadWrite ,buffer ,nullptr ,nullptr ); 218} 219 220HRESULT Tensor ::create (eDataType type ,const std::array < uint32_t ,4 >& sizeElements ) 221{ 222 std::initializer_list < uint32_t > il (sizeElements .data (),sizeElements .data ()+ 4 ); 223return create (type ,il ); 224} 225 226eDataType Tensor ::getType ()const 227{ 228ID3D11ShaderResourceView * const srv = * this ; 229if (nullptr == srv ) 230throw OLE_E_BLANK ; 231 232D3D11_SHADER_RESOURCE_VIEW_DESC viewDesc ; 233srv -> GetDesc (& viewDesc ); 234const DXGI_FORMAT format = viewDesc .Format ; 235switch (format ) 236 { 237case DXGI_FORMAT_R32_FLOAT : 238return eDataType::FP32 ; 239case DXGI_FORMAT_R16_FLOAT : 240return eDataType::FP16 ; 241case DXGI_FORMAT_R32_UINT : 242return eDataType::U32 ; 243 } 244throw E_NOTIMPL ; 245} 246 247CComPtr < ID3D11Buffer > Tensor ::getBuffer ()const 248{ 249ID3D11ShaderResourceView * const srv = * this ; 250if (nullptr == srv ) 251throw OLE_E_BLANK ; 252 253CComPtr < ID3D11Resource > res ; 254srv -> GetResource (& res ); 255 256CComPtr < ID3D11Buffer > buff ; 257check (res .QueryInterface (& buff ) ); 258return buff ; 259} 260 261uint32_t Tensor ::dxgiSizeof (DXGI_FORMAT format ) 262{ 263switch (format ) 264 { 265case DXGI_FORMAT_R16_FLOAT : 266return 2 ; 267case DXGI_FORMAT_R32_FLOAT : 268case DXGI_FORMAT_R32_UINT : 269return 4 ; 270 } 271throw E_INVALIDARG ; 272} 273 274void Tensor ::downloadImpl (const D3D11_SHADER_RESOURCE_VIEW_DESC & viewDesc ,uint32_t countElements ,size_t cbElement ,void * rdi )const 275{ 276assert (viewDesc .ViewDimension == D3D_SRV_DIMENSION_BUFFER ); 277const uint32_t idxFirst = viewDesc .Buffer .FirstElement ; 278 279CComPtr < ID3D11Buffer > buff = getBuffer (); 280D3D11_BUFFER_DESC desc ; 281buff -> GetDesc (& desc ); 282desc .BindFlags = 0 ; 283desc .Usage = D3D11_USAGE_STAGING ; 284desc .CPUAccessFlags = D3D11_CPU_ACCESS_READ ; 285 286CComPtr < ID3D11Buffer > staging ; 287check (device ()-> CreateBuffer (& desc ,nullptr ,& staging ) ); 288context ()-> CopyResource (staging ,buff ); 289 290MappedResource mapped ; 291check (mapped .map (staging , true ) ); 292const uint8_t * rsi = (const uint8_t * )mapped .data (); 293rsi += cbElement * idxFirst ; 294memcpy (rdi ,rsi ,cbElement * countElements ); 295} 296 297void Tensor ::download ( std::vector < float >& vec )const 298{ 299ID3D11ShaderResourceView * const srv = * this ; 300if (nullptr == srv ) 301throw OLE_E_BLANK ; 302 303D3D11_SHADER_RESOURCE_VIEW_DESC viewDesc ; 304srv -> GetDesc (& viewDesc ); 305if (viewDesc .Format != DXGI_FORMAT_R32_FLOAT ) 306throw E_INVALIDARG ; 307 308uint32_t countElements = viewDesc .Buffer .NumElements ; 309vec .resize (countElements ); 310downloadImpl (viewDesc ,countElements ,4 ,vec .data () ); 311} 312 313void Tensor ::download ( std::vector < uint16_t >& vec )const 314{ 315ID3D11ShaderResourceView * const srv = * this ; 316if (nullptr == srv ) 317throw OLE_E_BLANK ; 318 319D3D11_SHADER_RESOURCE_VIEW_DESC viewDesc ; 320srv -> GetDesc (& viewDesc ); 321if (viewDesc .Format != DXGI_FORMAT_R16_FLOAT ) 322throw E_INVALIDARG ; 323 324uint32_t countElements = viewDesc .Buffer .NumElements ; 325vec .resize (countElements ); 326downloadImpl (viewDesc ,countElements ,2 ,vec .data () ); 327} 328 329Tensor Tensor ::reshape3d (uint32_t ne0 ,uint32_t ne1 ,uint32_t ne2 )const 330{ 331if ( !isContinuous () ) 332throw E_NOTIMPL ; 333if (countElements ()!= ne0 * ne1 * ne2 ) 334throw E_INVALIDARG ; 335 336Tensor res = * this ; 337res .ne = {ne0 ,ne1 ,ne2 ,1 }; 338res .setDenseStrides (); 339return res ; 340}