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

KonstantinRefactor, removed a redundant function3f3a9a1

master
1.9 KiB74 linesraw
1#include "stdafx.h"
2#include "logger.h"
3#include "miscUtils.h"
4
5namespace
6{
7	using namespace Whisper;
8
9	// Terminal color map. 10 colors grouped in ranges [0.0, 0.1, ..., 0.9]
10	// Lowest is red, middle is yellow, highest is green.
11	static const std::array<const char*, 10> k_colors =
12	{
13		"\033[38;5;196m", "\033[38;5;202m", "\033[38;5;208m", "\033[38;5;214m", "\033[38;5;220m",
14		"\033[38;5;226m", "\033[38;5;190m", "\033[38;5;154m", "\033[38;5;118m", "\033[38;5;82m",
15	};
16
17	static int colorIndex( const sToken& tok )
18	{
19		const float p = tok.probability;
20		const float p3 = p * p * p;
21		int col = (int)( p3 * float( k_colors.size() ) );
22		col = std::max( 0, std::min( (int)k_colors.size() - 1, col ) );
23		return col;
24	}
25}
26
27void printTime( CStringA& rdi, Whisper::sTimeSpan time, bool comma )
28{
29	Whisper::sTimeSpanFields fields = time;
30	const uint32_t hours = fields.days * 24 + fields.hours;
31	const char separator = comma ? ',' : '.';
32	rdi.AppendFormat( "%02d:%02d:%02d%c%03d",
33		(int)hours,
34		(int)fields.minutes,
35		(int)fields.seconds,
36		separator,
37		fields.ticks / 10'000 );
38}
39
40HRESULT logNewSegments( const iTranscribeResult* results, size_t newSegments, bool printSpecial )
41{
42	sTranscribeLength length;
43	CHECK( results->getSize( length ) );
44
45	const size_t len = length.countSegments;
46	size_t i = len - newSegments;
47
48	const sSegment* const segments = results->getSegments();
49	const sToken* const tokens = results->getTokens();
50
51	CStringA str;
52	for( ; i < len; i++ )
53	{
54		const sSegment& seg = segments[ i ];
55		str = "[";
56		printTime( str, seg.time.begin );
57		str += " --> ";
58		printTime( str, seg.time.end );
59		str += "]  ";
60
61		for( uint32_t j = 0; j < seg.countTokens; j++ )
62		{
63			const sToken& tok = tokens[ seg.firstToken + j ];
64			if( !printSpecial && ( tok.flags & eTokenFlags::Special ) )
65				continue;
66			str += k_colors[ colorIndex( tok ) ];
67			str += tok.text;
68			str += "\033[0m";
69		}
70		logInfo( u8"%s", cstr( str ) );
71	}
72
73	return S_OK;
74}