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

KonstantinBugfix, incorrect output of command-line examples when launched with multiple input files3ba8e63

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