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

KonstantinWorkaround for the Microsoft’s bug in their MP3 decoder MFT9df2ee2

master
3.3 KiB128 linesraw
1#include "stdafx.h"
2#include "../API/iMediaFoundation.cl.h"
3#include "mfStartup.h"
4#include "../ComLightLib/comLightServer.h"
5#include "loadAudioFile.h"
6#include <mfidl.h>
7#include <mfreadwrite.h>
8#include "mfUtils.h"
9#include "AudioCapture.h"
10#include <mfapi.h>
11
12namespace Whisper
13{
14	class AudioReader : public ComLight::ObjectRoot<iAudioReader>
15	{
16		CComPtr<IMFSourceReader> reader;
17		bool wantStereo;
18		CComPtr<iMediaFoundation> mediaFoundation;
19		mutable int64_t preciseSamplesCount = 0;
20
21		HRESULT COMLIGHTCALL getReader( IMFSourceReader** pp ) const noexcept override final
22		{
23			if( pp == nullptr )
24				return E_POINTER;
25			CComPtr<IMFSourceReader> res = reader;
26			*pp = res.Detach();;
27			return S_OK;
28		}
29		HRESULT COMLIGHTCALL requestedStereo() const noexcept override final
30		{
31			return wantStereo ? S_OK : S_FALSE;
32		}
33		HRESULT COMLIGHTCALL getDuration( int64_t& rdi ) const noexcept override final
34		{
35			if( reader )
36			{
37				if( 0 == preciseSamplesCount )
38					return getStreamDuration( reader, rdi );
39				else
40				{	rdi = MFllMulDiv( preciseSamplesCount, 10'000'000, SAMPLE_RATE, 0 );
41					return S_OK;
42				}
43			}
44			return OLE_E_BLANK;
45		}
46	public:
47		HRESULT open( iMediaFoundation* owner, LPCTSTR path, bool stereo )
48		{
49			HRESULT hr = MFCreateSourceReaderFromURL( path, nullptr, &reader );
50			if( FAILED( hr ) )
51			{
52				logErrorHr( hr, u8"MFCreateSourceReaderFromURL failed" );
53				return hr;
54			}
55			wantStereo = stereo;
56			mediaFoundation = owner;
57			logDebug16( L"Created source reader from the file \"%s\"", path );
58			return S_OK;
59		}
60		void setPreciseSamplesCount( int64_t count ) const
61		{
62			preciseSamplesCount = count;
63		}
64	};
65
66	void setPreciseSamplesCount( const iAudioReader* ar, int64_t count )
67	{
68		const AudioReader* r = static_cast<const AudioReader*>( ar );
69		r->setPreciseSamplesCount( count );
70	}
71
72	class MediaFoundation : public ComLight::ObjectRoot<iMediaFoundation>
73	{
74		MfStartupRaii raii;
75		DWORD tid = ~(DWORD)0;
76
77		virtual HRESULT COMLIGHTCALL loadAudioFile( LPCTSTR path, bool stereo, iAudioBuffer** pp ) const noexcept override final
78		{
79			return Whisper::loadAudioFile( path, stereo, pp );
80		}
81		virtual HRESULT COMLIGHTCALL openAudioFile( LPCTSTR path, bool stereo, iAudioReader** pp ) noexcept override final
82		{
83			if( nullptr == path || nullptr == pp )
84				return E_POINTER;
85
86			ComLight::CComPtr<ComLight::Object<AudioReader>> res;
87			CHECK( ComLight::Object<AudioReader>::create( res ) );
88			CHECK( res->open( this, path, stereo ) );
89
90			res.detach( pp );
91			return S_OK;
92		}
93		HRESULT COMLIGHTCALL listCaptureDevices( pfnFoundCaptureDevices pfn, void* pv ) noexcept override final
94		{
95			return captureDeviceList( pfn, pv );
96		}
97		HRESULT COMLIGHTCALL openCaptureDevice( LPCTSTR endpoint, const sCaptureParams& captureParams, iAudioCapture** pp ) noexcept override final
98		{
99			return captureOpen( this, endpoint, captureParams, pp );
100		}
101	protected:
102
103		HRESULT FinalConstruct()
104		{
105			CHECK( raii.startup() );
106			tid = GetCurrentThreadId();
107			return S_OK;
108		}
109
110	public:
111
112		~MediaFoundation() override
113		{
114			assert( tid == GetCurrentThreadId() );
115		}
116	};
117}
118
119HRESULT COMLIGHTCALL Whisper::initMediaFoundation( iMediaFoundation** pp )
120{
121	if( nullptr == pp )
122		return E_POINTER;
123
124	ComLight::CComPtr<ComLight::Object<MediaFoundation>> obj;
125	CHECK( ComLight::Object<MediaFoundation>::create( obj ) );
126	obj.detach( pp );
127	return S_OK;
128}