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

KonstantinSource codes8c4603c

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