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
4.4 KiB116 linesraw
1#pragma once
2#include <vector>
3#include "TempBuffers.h"
4#include "ConstantBuffer.h"
5#include "Tensor.h"
6#include "../Utils/GpuProfiler.h"
7#include "../Utils/ProfileCollection.h"
8
9namespace DirectCompute
10{
11	enum struct eComputeShader : uint16_t;
12
13	class MlContext
14	{
15		// When false, the implementation is 100% compatible with the CPU-running code written by Georgi Gerganov
16		// When true, the implementation is much faster, and doesn't require FP64 support in the compute shaders.
17		// FP64 is an optional feature, not all GPUs support that.
18		static constexpr bool enableInexactOptimizations = true;
19
20		ConstantBuffer cb;
21		TempBuffers temp;
22		CComPtr<ID3D11Buffer> flashAttentionConstants;
23
24		void convolutionImpl( const Tensor& a, const Tensor& b, Tensor& res, bool is2 );
25
26		void cwiseBinary( const Tensor& a, const Tensor& b, Tensor& res, eComputeShader cs );
27		Tensor cwiseBinary( const Tensor& a, const Tensor& b, eComputeShader cs );
28
29		void mulMatDot( const Tensor& a, const Tensor& b, Tensor& res );
30		void mulMatMad( const Tensor& a, const Tensor& b, Tensor& res );
31		void mulMatTiled( const Tensor& a, const Tensor& b, Tensor& res );
32
33		void bindShader( eComputeShader cs );
34
35	protected:
36		void copyImpl( const Tensor& a, Tensor& res, bool downcastFp32 );
37
38		// Create a dense output tensor for the results of a computation
39		// Override this method to implement a pool of these tensors
40		virtual Tensor createTensor( eDataType type, const std::array<uint32_t, 4>& ne );
41
42		Tensor createTensor( eDataType type, std::initializer_list<uint32_t> ne );
43
44		GpuProfiler profiler;
45
46		CComPtr<ID3D11Buffer>& getSmallConstantBuffer() { return temp.smallCb; }
47
48	public:
49		MlContext( Whisper::ProfileCollection& profileColl );
50		MlContext( const MlContext& ) = delete;
51
52		// res = a * b
53		void mulMat( const Tensor& a, const Tensor& b, Tensor& res );
54
55		void flashAttention( const Tensor& q, const Tensor& k, const Tensor& v, Tensor& res, bool masked );
56
57		inline void convolution( const Tensor& a, const Tensor& b, Tensor& res )
58		{
59			convolutionImpl( a, b, res, false );
60		}
61		void convolution2( const Tensor& a, const Tensor& b, Tensor& res )
62		{
63			convolutionImpl( a, b, res, true );
64		}
65
66		void norm( const Tensor& a, Tensor& res );
67
68		Tensor conv_1d_1s( const Tensor& a, const Tensor& b );
69		Tensor conv_1d_2s( const Tensor& a, const Tensor& b );
70
71		Tensor add( const Tensor& a, const Tensor& b );
72		void addInPlace( Tensor& a, const Tensor& b );
73
74		Tensor view2d( const Tensor& a, uint32_t ne0, uint32_t ne1, uint32_t nb1, uint32_t offset );
75		Tensor transpose( const Tensor& a );
76
77		Tensor norm( const Tensor& a );
78		Tensor mulMat( const Tensor& a, const Tensor& b );
79		Tensor mulMatEx( const Tensor& a, const Tensor& b, const char* tagName );
80		Tensor permute( const Tensor& a, uint8_t axis0, uint8_t axis1, uint8_t axis2, uint8_t axis3 );
81		Tensor flashAttention( const Tensor& q, const Tensor& k, const Tensor& v, bool masked );
82
83		Tensor copy( const Tensor& a, eDataType type, std::initializer_list<uint32_t> size );
84		void copyInPlace( Tensor& dest, const Tensor& a, eDataType type, std::initializer_list<uint32_t> size );
85
86		void dbgPrintDifference( const ggml_tensor* reference, const Tensor& gpu, const char * what, bool trapToDebugger = true );
87
88		void scale( Tensor& a, float mul );
89
90		void addRepeat( Tensor& a, const Tensor& b );
91		void addRepeatScale( Tensor& a, const Tensor& b, float scale );
92		void fmaRepeat( Tensor& a, const Tensor& mul, const Tensor& add );
93
94		// ggml_diag_mask_inf
95		void diagMaskInf( Tensor& a, uint32_t n_past );
96		// ggml_soft_max
97		void softMax( Tensor& a, float inputScale = 1.0f );
98
99		void addRepeatGelu( Tensor& a, const Tensor& b );
100
101		// Extract rows from tokenEmbedding matrix, row indices are taken from the `embd` R32_UINT row vector
102		// Extract same count of rows from positionalEmbedding matrix, starting at the `pastTokensCount` row
103		// Return a new FP32 matrix with the sum of these rows
104		Tensor addRows( const Tensor& tokenEmbedding, const Tensor& positionalEmbedding, const Tensor& embd, uint32_t pastTokensCount );
105
106		Tensor reshapePanels( const Tensor& a );
107
108		Tensor mulMatTiledEx( const Tensor& a, const Tensor& b );
109		Tensor mulMatByRowTiledEx( const Tensor& a, const Tensor& b );
110
111		// An equivalent of addRepeat( dest, pattern ) followed by addInPlace( dest, finalAdd )
112		void addRepeatEx( Tensor& dest, const Tensor& pattern, const Tensor& finalAdd );
113
114		__m128i getMemoryUse() const;
115	};
116}