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
2.0 KiB69 linesraw
1#include "stdafx.h"
2#include "mfUtils.h"
3#include <mfapi.h>
4
5HRESULT Whisper::createMediaType( bool stereo, IMFMediaType** pp )
6{
7	if( nullptr == pp )
8		return E_POINTER;
9
10	CComPtr<IMFMediaType> mt;
11	CHECK( MFCreateMediaType( &mt ) );
12	CHECK( mt->SetGUID( MF_MT_MAJOR_TYPE, MFMediaType_Audio ) );
13	CHECK( mt->SetGUID( MF_MT_SUBTYPE, MFAudioFormat_Float ) );
14	CHECK( mt->SetUINT32( MF_MT_AUDIO_SAMPLES_PER_SECOND, SAMPLE_RATE ) );
15
16	const uint32_t channels = stereo ? 2 : 1;
17	CHECK( mt->SetUINT32( MF_MT_AUDIO_NUM_CHANNELS, channels ) );
18	CHECK( mt->SetUINT32( MF_MT_AUDIO_BLOCK_ALIGNMENT, channels * 4 ) );
19	CHECK( mt->SetUINT32( MF_MT_AUDIO_AVG_BYTES_PER_SECOND, channels * 4 * SAMPLE_RATE ) );
20	CHECK( mt->SetUINT32( MF_MT_AUDIO_BITS_PER_SAMPLE, 32 ) );
21	CHECK( mt->SetUINT32( MF_MT_ALL_SAMPLES_INDEPENDENT, TRUE ) );
22
23	*pp = mt.Detach();
24
25	return S_OK;
26}
27
28HRESULT Whisper::getStreamDuration( IMFSourceReader* reader, int64_t& duration )
29{
30	PROPVARIANT var;
31	PropVariantInit( &var );
32	CHECK( reader->GetPresentationAttribute( MF_SOURCE_READER_MEDIASOURCE, MF_PD_DURATION, &var ) );
33
34	if( var.vt == VT_UI8 )
35	{
36		// The documentation says the type of that attribute is UINT64
37		// https://learn.microsoft.com/en-us/windows/win32/medfound/mf-pd-duration-attribute
38		duration = var.uhVal.QuadPart;
39		return S_OK;
40	}
41	logError( u8"Unexpected type of MF_PD_DURATION attribute" );
42	return E_INVALIDARG;
43}
44
45HRESULT Whisper::validateCurrentMediaType( IMFSourceReader* reader, uint32_t expectedChannels )
46{
47	CComPtr<IMFMediaType> mt;
48	CHECK( reader->GetCurrentMediaType( MF_SOURCE_READER_FIRST_AUDIO_STREAM, &mt ) );
49
50	GUID guid;
51	CHECK( mt->GetGUID( MF_MT_MAJOR_TYPE, &guid ) );
52	if( guid != MFMediaType_Audio )
53		return E_FAIL;
54
55	CHECK( mt->GetGUID( MF_MT_SUBTYPE, &guid ) );
56	if( guid != MFAudioFormat_Float )
57		return E_FAIL;
58
59	UINT32 u32;
60	CHECK( mt->GetUINT32( MF_MT_AUDIO_SAMPLES_PER_SECOND, &u32 ) );
61	if( u32 != SAMPLE_RATE )
62		return E_FAIL;
63
64	CHECK( mt->GetUINT32( MF_MT_AUDIO_NUM_CHANNELS, &u32 ) );
65	if( u32 != expectedChannels )
66		return E_FAIL;
67
68	return S_OK;
69}