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
8c4603c
master
1#include "stdafx.h" 2#if BUILD_BOTH_VERSIONS 3#include "../API/iContext.cl.h" 4#include "convertThings.h" 5using namespace Whisper ; 6 7sFullParams makeNewParams (const whisper_full_params & wfp ) 8{ 9assert (nullptr == wfp .encoder_begin_callback ); 10assert (nullptr == wfp .new_segment_callback ); 11 12sFullParams res ; 13memset (& res ,0 ,sizeof (res ) ); 14 15res .strategy = (eSamplingStrategy )wfp .strategy ; 16res .cpuThreads = wfp .n_threads ; 17res .n_max_text_ctx = wfp .n_max_text_ctx ; 18res .offset_ms = wfp .offset_ms ; 19res .duration_ms = wfp .duration_ms ; 20 21// flags 22uint32_t flags = 0 ; 23if (wfp .translate )flags |= (uint32_t )eFullParamsFlags::Translate ; 24if (wfp .no_context )flags |= (uint32_t )eFullParamsFlags::NoContext ; 25if (wfp .single_segment )flags |= (uint32_t )eFullParamsFlags::SingleSegment ; 26if (wfp .print_special )flags |= (uint32_t )eFullParamsFlags::PrintSpecial ; 27if (wfp .print_progress )flags |= (uint32_t )eFullParamsFlags::PrintProgress ; 28if (wfp .print_realtime )flags |= (uint32_t )eFullParamsFlags::PrintRealtime ; 29if (wfp .print_timestamps )flags |= (uint32_t )eFullParamsFlags::PrintTimestamps ; 30if (wfp .token_timestamps )flags |= (uint32_t )eFullParamsFlags::TokenTimestamps ; 31if (wfp .speed_up )flags |= (uint32_t )eFullParamsFlags::SpeedupAudio ; 32res .flags = (eFullParamsFlags )flags ; 33 34res .language = findLanguageKeyA (wfp .language ); 35res .thold_pt = wfp .thold_pt ; 36res .thold_ptsum = wfp .thold_ptsum ; 37res .max_len = wfp .max_len ; 38res .greedy .n_past = wfp .greedy .n_past ; 39res .beam_search .n_past = wfp .beam_search .n_past ; 40res .beam_search .beam_width = wfp .beam_search .beam_width ; 41res .beam_search .n_best = wfp .beam_search .n_best ; 42res .audio_ctx = wfp .audio_ctx ; 43res .prompt_tokens = wfp .prompt_tokens ; 44res .prompt_n_tokens = wfp .prompt_n_tokens ; 45 46return res ; 47} 48 49namespace 50{ 51class NewParamsTemp 52 { 53char language [5 ]; 54iContext * newContext ; 55pfnNewSegment newSegment ; 56pfnEncoderBegin encoderBegin ; 57 58static bool encBegin (struct whisper_context * ctx ,void * user_data ); 59static void newSeg (struct whisper_context * ctx ,int n_new ,void * user_data ); 60 61public : 62 63void initialize (whisper_full_params & res ,const Whisper ::sFullParams & rsi ,Whisper ::iContext * context ) 64 { 65* (uint32_t * )(& language [0 ] )= rsi .language ; 66language [4 ]= '\0' ; 67res .language = language ; 68 69newContext = context ; 70 71if (nullptr != rsi .encoder_begin_callback ) 72 { 73encoderBegin = rsi .encoder_begin_callback ; 74res .encoder_begin_callback = & encBegin ; 75res .encoder_begin_callback_user_data = rsi .encoder_begin_callback_user_data ; 76 } 77else 78 { 79encoderBegin = nullptr ; 80res .encoder_begin_callback = nullptr ; 81res .encoder_begin_callback_user_data = nullptr ; 82 } 83 84if (nullptr != rsi .new_segment_callback ) 85 { 86newSegment = rsi .new_segment_callback ; 87res .new_segment_callback = & newSeg ; 88res .new_segment_callback_user_data = rsi .new_segment_callback_user_data ; 89 } 90else 91 { 92newSegment = nullptr ; 93res .new_segment_callback = nullptr ; 94res .new_segment_callback_user_data = nullptr ; 95 } 96 } 97 }; 98 99static thread_localNewParamsTemp npTemp ; 100 101bool NewParamsTemp ::encBegin (struct whisper_context * ctx ,void * user_data ) 102 { 103const NewParamsTemp & tmp = npTemp ; 104HRESULT hr = tmp .encoderBegin (tmp .newContext ,user_data ); 105if (SUCCEEDED (hr ) ) 106return S_OK == hr ; 107throw hr ; 108 } 109 110void NewParamsTemp ::newSeg (struct whisper_context * ctx ,int n_new ,void * user_data ) 111 { 112assert (n_new >=0 ); 113const NewParamsTemp & tmp = npTemp ; 114HRESULT hr = tmp .newSegment (tmp .newContext , (uint32_t )n_new ,user_data ); 115if (SUCCEEDED (hr ) ) 116return ; 117throw hr ; 118 } 119} 120 121whisper_full_params makeOldParams (const Whisper ::sFullParams & rsi ,Whisper ::iContext * context ) 122{ 123whisper_full_params res ; 124memset (& res ,0 ,sizeof (res ) ); 125 126res .strategy = (whisper_sampling_strategy )rsi .strategy ; 127res .n_threads = rsi .cpuThreads ; 128res .n_max_text_ctx = rsi .n_max_text_ctx ; 129res .offset_ms = rsi .offset_ms ; 130res .duration_ms = rsi .duration_ms ; 131 132// flags 133const uint32_t flags = (uint32_t )rsi .flags ; 134auto hasFlag = [= ](eFullParamsFlags bit ) {return 0 != (flags & (uint32_t )bit ); }; 135 136res .translate = hasFlag ( eFullParamsFlags::Translate ); 137res .no_context = hasFlag ( eFullParamsFlags::NoContext ); 138res .single_segment = hasFlag ( eFullParamsFlags::SingleSegment ); 139res .print_special = hasFlag ( eFullParamsFlags::PrintSpecial ); 140res .print_progress = hasFlag ( eFullParamsFlags::PrintProgress ); 141res .print_realtime = hasFlag ( eFullParamsFlags::PrintRealtime ); 142res .print_timestamps = hasFlag ( eFullParamsFlags::PrintTimestamps ); 143res .token_timestamps = hasFlag ( eFullParamsFlags::TokenTimestamps ); 144res .speed_up = hasFlag ( eFullParamsFlags::SpeedupAudio ); 145 146res .thold_pt = rsi .thold_pt ; 147res .thold_ptsum = rsi .thold_ptsum ; 148res .max_len = rsi .max_len ; 149res .greedy .n_past = rsi .greedy .n_past ; 150res .beam_search .n_past = rsi .beam_search .n_past ; 151res .beam_search .beam_width = rsi .beam_search .beam_width ; 152res .beam_search .n_best = rsi .beam_search .n_best ; 153res .audio_ctx = rsi .audio_ctx ; 154res .prompt_tokens = rsi .prompt_tokens ; 155res .prompt_n_tokens = rsi .prompt_n_tokens ; 156 157NewParamsTemp & tmp = npTemp ; 158tmp .initialize (res ,rsi ,context ); 159return res ; 160} 161 162#include "../Whisper/TranscribeResult.h" 163#include <mfapi.h> 164 165namespace 166{ 167inline sTimeSpan time (int64_t wt ) 168 { 169int64_t ticks = MFllMulDiv (wt ,10'000'000 ,100 ,0 ); 170return sTimeSpan { (uint64_t )ticks }; 171 } 172 173void makeNewResults (whisper_context * ctx ,Whisper ::eResultFlags flags ,TranscribeResult & res ) 174 { 175const bool makeTokens = 0 != (flags & eResultFlags::Tokens ); 176res .segments .clear (); 177res .tokens .clear (); 178 179const int countSegments = whisper_full_n_segments (ctx ); 180res .segments .resize (countSegments ); 181const int tokenEot = whisper_token_eot (ctx ); 182for (int i = 0 ;i < countSegments ;i ++ ) 183 { 184sSegment & seg = res .segments [i ]; 185seg .text = whisper_full_get_segment_text (ctx ,i ); 186seg .time .begin = time (whisper_full_get_segment_t0 (ctx ,i ) ); 187seg .time .end = time (whisper_full_get_segment_t1 (ctx ,i ) ); 188 189seg .firstToken = (uint32_t )res .tokens .size (); 190seg .countTokens = 0 ; 191if ( !makeTokens ) 192continue ; 193 194const int countTokens = whisper_full_n_tokens (ctx ,i ); 195seg .countTokens = countTokens ; 196res .tokens .resize (res .tokens .size ()+ countTokens ); 197for (int t = 0 ;t < countTokens ;t ++ ) 198 { 199sToken & tok = res .tokens [seg .firstToken + t ]; 200tok .text = whisper_full_get_token_text (ctx ,i ,t ); 201 202const whisper_token_data src = whisper_full_get_token_data (ctx ,i ,t ); 203tok .time .begin = time (src .t0 ); 204tok .time .end = time (src .t1 ); 205tok .probability = src .p ; 206tok .probabilityTimestamp = src .pt ; 207tok .ptsum = src .ptsum ; 208tok .vlen = src .vlen ; 209tok .id = src .id ; 210uint32_t flags = 0 ; 211if (src .id >=tokenEot ) 212flags |= eTokenFlags::Special ; 213tok .flags = (eTokenFlags )flags ; 214 } 215 } 216 } 217} 218 219HRESULT makeNewResults (whisper_context * ctx ,Whisper ::eResultFlags flags ,Whisper ::iTranscribeResult ** pp ) 220{ 221static TranscribeResultStatic trs ; 222if (flags & eResultFlags::NewObject ) 223 { 224return E_NOTIMPL ; 225 } 226else 227 { 228makeNewResults (ctx ,flags ,trs ); 229* pp = & trs ; 230 (* pp )-> AddRef (); 231return S_OK ; 232 } 233} 234#endif