echogarden 1.0.4 → 1.1.0
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- package/README.md +26 -23
- package/data/schemas/options.json +177 -36
- package/dist/alignment/SpeechAlignment.d.ts +1 -1
- package/dist/alignment/SpeechAlignment.js +1 -1
- package/dist/alignment/SpeechAlignment.js.map +1 -1
- package/dist/api/API.d.ts +1 -0
- package/dist/api/API.js +1 -0
- package/dist/api/API.js.map +1 -1
- package/dist/api/APIOptions.d.ts +1 -0
- package/dist/api/Alignment.d.ts +3 -3
- package/dist/api/Alignment.js +5 -10
- package/dist/api/Alignment.js.map +1 -1
- package/dist/api/LanguageDetection.d.ts +5 -7
- package/dist/api/LanguageDetection.js +3 -2
- package/dist/api/LanguageDetection.js.map +1 -1
- package/dist/api/Recognition.d.ts +4 -5
- package/dist/api/Recognition.js +5 -8
- package/dist/api/Recognition.js.map +1 -1
- package/dist/api/SourceSeparation.d.ts +2 -0
- package/dist/api/SourceSeparation.js +4 -2
- package/dist/api/SourceSeparation.js.map +1 -1
- package/dist/api/Synthesis.d.ts +3 -1
- package/dist/api/Synthesis.js +9 -10
- package/dist/api/Synthesis.js.map +1 -1
- package/dist/api/Translation.d.ts +1 -1
- package/dist/api/Translation.js +4 -8
- package/dist/api/Translation.js.map +1 -1
- package/dist/api/TranslationAlignment.d.ts +31 -0
- package/dist/api/TranslationAlignment.js +121 -0
- package/dist/api/TranslationAlignment.js.map +1 -0
- package/dist/api/VoiceActivityDetection.d.ts +5 -1
- package/dist/api/VoiceActivityDetection.js +38 -2
- package/dist/api/VoiceActivityDetection.js.map +1 -1
- package/dist/audio/AudioPlayer.js +6 -1
- package/dist/audio/AudioPlayer.js.map +1 -1
- package/dist/cli/CLI.js +85 -0
- package/dist/cli/CLI.js.map +1 -1
- package/dist/dsp/FFT.js.map +1 -1
- package/dist/math/MedianFilter.d.ts +5 -0
- package/dist/math/MedianFilter.js +102 -0
- package/dist/math/MedianFilter.js.map +1 -0
- package/dist/math/VectorMath.d.ts +0 -2
- package/dist/math/VectorMath.js +1 -25
- package/dist/math/VectorMath.js.map +1 -1
- package/dist/recognition/OpenAICloudSTT.d.ts +1 -1
- package/dist/recognition/OpenAICloudSTT.js.map +1 -1
- package/dist/recognition/SileroSTT.d.ts +22 -1
- package/dist/recognition/SileroSTT.js +122 -95
- package/dist/recognition/SileroSTT.js.map +1 -1
- package/dist/recognition/WhisperCppSTT.js +1 -1
- package/dist/recognition/WhisperCppSTT.js.map +1 -1
- package/dist/recognition/WhisperSTT.d.ts +52 -19
- package/dist/recognition/WhisperSTT.js +645 -494
- package/dist/recognition/WhisperSTT.js.map +1 -1
- package/dist/server/Server.js.map +1 -1
- package/dist/source-separation/MDXNetSourceSeparation.d.ts +5 -3
- package/dist/source-separation/MDXNetSourceSeparation.js +26 -19
- package/dist/source-separation/MDXNetSourceSeparation.js.map +1 -1
- package/dist/speech-language-detection/SileroLanguageDetection.d.ts +15 -9
- package/dist/speech-language-detection/SileroLanguageDetection.js +23 -16
- package/dist/speech-language-detection/SileroLanguageDetection.js.map +1 -1
- package/dist/synthesis/EspeakTTS.js +4 -0
- package/dist/synthesis/EspeakTTS.js.map +1 -1
- package/dist/synthesis/GoogleCloudTTS.js.map +1 -1
- package/dist/synthesis/VitsTTS.d.ts +8 -6
- package/dist/synthesis/VitsTTS.js +36 -31
- package/dist/synthesis/VitsTTS.js.map +1 -1
- package/dist/tests/Test.js.map +1 -1
- package/dist/utilities/OnnxUtilities.d.ts +14 -0
- package/dist/utilities/OnnxUtilities.js +43 -0
- package/dist/utilities/OnnxUtilities.js.map +1 -0
- package/dist/utilities/Utilities.d.ts +4 -8
- package/dist/utilities/Utilities.js +35 -58
- package/dist/utilities/Utilities.js.map +1 -1
- package/dist/voice-activity-detection/SileroVAD.d.ts +5 -3
- package/dist/voice-activity-detection/SileroVAD.js +9 -11
- package/dist/voice-activity-detection/SileroVAD.js.map +1 -1
- package/docs/API.md +54 -34
- package/docs/CLI.md +25 -13
- package/docs/Contributing.md +4 -2
- package/docs/Engines.md +43 -32
- package/docs/Licenses.md +3 -4
- package/docs/Options.md +47 -11
- package/docs/Releases.md +4 -0
- package/docs/Server.md +8 -6
- package/docs/Tasklist.md +39 -52
- package/docs/Technical.md +1 -1
- package/package.json +8 -12
- package/src/alignment/SpeechAlignment.ts +1 -1
- package/src/api/API.ts +1 -0
- package/src/api/APIOptions.ts +1 -0
- package/src/api/Alignment.ts +10 -14
- package/src/api/LanguageDetection.ts +14 -10
- package/src/api/Recognition.ts +17 -10
- package/src/api/SourceSeparation.ts +7 -2
- package/src/api/Synthesis.ts +26 -11
- package/src/api/Translation.ts +14 -8
- package/src/api/TranslationAlignment.ts +213 -0
- package/src/api/VoiceActivityDetection.ts +66 -3
- package/src/audio/AudioPlayer.ts +6 -2
- package/src/cli/CLI.ts +121 -2
- package/src/dsp/FFT.ts +3 -0
- package/src/math/MedianFilter.ts +124 -0
- package/src/math/VectorMath.ts +1 -36
- package/src/recognition/OpenAICloudSTT.ts +27 -27
- package/src/recognition/SileroSTT.ts +149 -102
- package/src/recognition/WhisperCppSTT.ts +1 -1
- package/src/recognition/WhisperSTT.ts +961 -684
- package/src/server/Server.ts +1 -1
- package/src/source-separation/MDXNetSourceSeparation.ts +35 -19
- package/src/speech-language-detection/SileroLanguageDetection.ts +53 -33
- package/src/synthesis/EspeakTTS.ts +8 -0
- package/src/synthesis/GoogleCloudTTS.ts +12 -1
- package/src/synthesis/VitsTTS.ts +57 -46
- package/src/tests/Test.ts +1 -1
- package/src/utilities/OnnxUtilities.ts +68 -0
- package/src/utilities/Utilities.ts +38 -66
- package/src/voice-activity-detection/SileroVAD.ts +15 -15
- package/dist/utilities/NdArrayUtilities.d.ts +0 -3
- package/dist/utilities/NdArrayUtilities.js +0 -23
- package/dist/utilities/NdArrayUtilities.js.map +0 -1
- package/src/utilities/NdArrayUtilities.ts +0 -31
|
@@ -2,14 +2,14 @@ import type * as Onnx from 'onnxruntime-node'
|
|
|
2
2
|
|
|
3
3
|
import { Logger } from '../utilities/Logger.js'
|
|
4
4
|
import { computeMelSpectogramUsingFilterbanks, Filterbank } from '../dsp/MelSpectogram.js'
|
|
5
|
-
import { clip,
|
|
6
|
-
import { indexOfMax, logOfVector, logSumExp, meanOfVector,
|
|
5
|
+
import { clip, containsInvalidCodepoint, getIntegerRange, getTokenRepetitionScore, splitFloat32Array, yieldToEventLoop } from '../utilities/Utilities.js'
|
|
6
|
+
import { indexOfMax, logOfVector, logSumExp, meanOfVector, softmax, stdDeviationOfVector } from '../math/VectorMath.js'
|
|
7
7
|
|
|
8
8
|
import { alignDTWWindowed } from '../alignment/DTWSequenceAlignmentWindowed.js'
|
|
9
9
|
import { extendDeep } from '../utilities/ObjectUtilities.js'
|
|
10
10
|
import { Timeline, TimelineEntry } from '../utilities/Timeline.js'
|
|
11
11
|
import { AlignmentPath } from '../alignment/SpeechAlignment.js'
|
|
12
|
-
import { getRawAudioDuration, RawAudio } from '../audio/AudioUtilities.js'
|
|
12
|
+
import { getRawAudioDuration, RawAudio, sliceRawAudio } from '../audio/AudioUtilities.js'
|
|
13
13
|
import { readFile } from '../utilities/FileSystem.js'
|
|
14
14
|
import path from 'path'
|
|
15
15
|
import type { LanguageDetectionResults } from '../api/API.js'
|
|
@@ -19,9 +19,21 @@ import chalk from 'chalk'
|
|
|
19
19
|
import { XorShift32RNG } from '../utilities/RandomGenerator.js'
|
|
20
20
|
import { detectSpeechLanguageByParts } from '../api/LanguageDetection.js'
|
|
21
21
|
import { type Tiktoken } from 'tiktoken/lite'
|
|
22
|
-
import { isPunctuation, isWhitespace } from '../nlp/Segmentation.js'
|
|
22
|
+
import { isPunctuation, isWhitespace, isWord, splitToSentences, splitToWords } from '../nlp/Segmentation.js'
|
|
23
|
+
import { medianOf5Filter } from '../math/MedianFilter.js'
|
|
24
|
+
import { getDeflateCompressionMetricsForString } from '../utilities/Compression.js'
|
|
25
|
+
import { getOnnxSessionOptions, makeOnnxLikeFloat32Tensor, OnnxExecutionProvider, OnnxLikeFloat32Tensor } from '../utilities/OnnxUtilities.js'
|
|
26
|
+
|
|
27
|
+
export async function recognize(
|
|
28
|
+
sourceRawAudio: RawAudio,
|
|
29
|
+
modelName: WhisperModelName,
|
|
30
|
+
modelDir: string,
|
|
31
|
+
task: WhisperTask,
|
|
32
|
+
sourceLanguage: string,
|
|
33
|
+
options: WhisperOptions) {
|
|
34
|
+
|
|
35
|
+
options = extendDeep(defaultWhisperOptions, options)
|
|
23
36
|
|
|
24
|
-
export async function recognize(sourceRawAudio: RawAudio, modelName: WhisperModelName, modelDir: string, task: WhisperTask, sourceLanguage: string, options: WhisperOptions) {
|
|
25
37
|
if (sourceRawAudio.sampleRate != 16000) {
|
|
26
38
|
throw new Error('Source audio must have a sampling rate of 16000')
|
|
27
39
|
}
|
|
@@ -46,14 +58,34 @@ export async function recognize(sourceRawAudio: RawAudio, modelName: WhisperMode
|
|
|
46
58
|
seed = Math.max(Math.floor(seed), 1) | 0
|
|
47
59
|
}
|
|
48
60
|
|
|
49
|
-
const
|
|
61
|
+
const encoderProviders: OnnxExecutionProvider[] =
|
|
62
|
+
options.encoderProvider ? [options.encoderProvider] : ['dml', 'cpu']
|
|
63
|
+
|
|
64
|
+
const decoderProviders: OnnxExecutionProvider[] =
|
|
65
|
+
options.decoderProvider ? [options.decoderProvider] : []
|
|
66
|
+
|
|
67
|
+
const whisper = new Whisper(
|
|
68
|
+
modelName,
|
|
69
|
+
modelDir,
|
|
70
|
+
encoderProviders,
|
|
71
|
+
decoderProviders,
|
|
72
|
+
seed)
|
|
50
73
|
|
|
51
74
|
const result = await whisper.recognize(sourceRawAudio, task, sourceLanguage, options)
|
|
52
75
|
|
|
53
76
|
return result
|
|
54
77
|
}
|
|
55
78
|
|
|
56
|
-
export async function align(
|
|
79
|
+
export async function align(
|
|
80
|
+
sourceRawAudio: RawAudio,
|
|
81
|
+
transcript: string,
|
|
82
|
+
modelName: WhisperModelName,
|
|
83
|
+
modelDir: string,
|
|
84
|
+
sourceLanguage: string,
|
|
85
|
+
options: WhisperAlignmentOptions) {
|
|
86
|
+
|
|
87
|
+
options = extendDeep(defaultWhisperAlignmentOptions, options)
|
|
88
|
+
|
|
57
89
|
if (sourceRawAudio.sampleRate != 16000) {
|
|
58
90
|
throw new Error('Source audio must have a sampling rate of 16000')
|
|
59
91
|
}
|
|
@@ -68,14 +100,72 @@ export async function align(sourceRawAudio: RawAudio, referenceText: string, mod
|
|
|
68
100
|
throw new Error(`The model '${modelName}' can only be used with English inputs. However, the given source language was ${languageCodeToName(sourceLanguage)}.`)
|
|
69
101
|
}
|
|
70
102
|
|
|
71
|
-
const
|
|
103
|
+
const encoderProviders: OnnxExecutionProvider[] =
|
|
104
|
+
options.encoderProvider ? [options.encoderProvider] : ['dml', 'cpu']
|
|
105
|
+
|
|
106
|
+
const decoderProviders: OnnxExecutionProvider[] =
|
|
107
|
+
options.decoderProvider ? [options.decoderProvider] : []
|
|
108
|
+
|
|
109
|
+
const whisper = new Whisper(
|
|
110
|
+
modelName,
|
|
111
|
+
modelDir,
|
|
112
|
+
encoderProviders,
|
|
113
|
+
decoderProviders,)
|
|
114
|
+
|
|
115
|
+
const timeline = await whisper.align(sourceRawAudio, transcript, sourceLanguage, 'transcribe', options)
|
|
116
|
+
|
|
117
|
+
return timeline
|
|
118
|
+
}
|
|
119
|
+
|
|
120
|
+
export async function alignEnglishTranslation(
|
|
121
|
+
sourceRawAudio: RawAudio,
|
|
122
|
+
translatedTranscript: string,
|
|
123
|
+
modelName: WhisperModelName,
|
|
124
|
+
modelDir: string,
|
|
125
|
+
sourceLanguage: string,
|
|
126
|
+
options: WhisperAlignmentOptions) {
|
|
127
|
+
|
|
128
|
+
options = extendDeep(defaultWhisperAlignmentOptions, options)
|
|
129
|
+
|
|
130
|
+
if (sourceRawAudio.sampleRate != 16000) {
|
|
131
|
+
throw new Error('Source audio must have a sampling rate of 16000')
|
|
132
|
+
}
|
|
133
|
+
|
|
134
|
+
sourceLanguage = getShortLanguageCode(sourceLanguage)
|
|
135
|
+
|
|
136
|
+
if (!(sourceLanguage in languageIdLookup)) {
|
|
137
|
+
throw new Error(`The source language ${languageCodeToName(sourceLanguage)} is not supported by the Whisper engine.`)
|
|
138
|
+
}
|
|
139
|
+
|
|
140
|
+
if (isEnglishOnlyModel(modelName)) {
|
|
141
|
+
throw new Error(`Translation alignment can only be done with multilingual models.`)
|
|
142
|
+
}
|
|
143
|
+
|
|
144
|
+
const encoderProviders: OnnxExecutionProvider[] =
|
|
145
|
+
options.encoderProvider ? [options.encoderProvider] : ['dml', 'cpu']
|
|
146
|
+
|
|
147
|
+
const decoderProviders: OnnxExecutionProvider[] =
|
|
148
|
+
options.decoderProvider ? [options.decoderProvider] : []
|
|
149
|
+
|
|
150
|
+
const whisper = new Whisper(
|
|
151
|
+
modelName,
|
|
152
|
+
modelDir,
|
|
153
|
+
encoderProviders,
|
|
154
|
+
decoderProviders,)
|
|
72
155
|
|
|
73
|
-
const timeline = await whisper.align(sourceRawAudio,
|
|
156
|
+
const timeline = await whisper.align(sourceRawAudio, translatedTranscript, sourceLanguage, 'translate', options)
|
|
74
157
|
|
|
75
158
|
return timeline
|
|
76
159
|
}
|
|
77
160
|
|
|
78
|
-
export async function detectLanguage(
|
|
161
|
+
export async function detectLanguage(
|
|
162
|
+
sourceRawAudio: RawAudio,
|
|
163
|
+
modelName: WhisperModelName,
|
|
164
|
+
modelDir: string,
|
|
165
|
+
options: WhisperLanguageDetectionOptions) {
|
|
166
|
+
|
|
167
|
+
options = extendDeep(defaultWhisperLanguageDetectionOptions, options)
|
|
168
|
+
|
|
79
169
|
if (sourceRawAudio.sampleRate != 16000) {
|
|
80
170
|
throw new Error('Source audio must have a sampling rate of 16000')
|
|
81
171
|
}
|
|
@@ -84,15 +174,25 @@ export async function detectLanguage(sourceRawAudio: RawAudio, modelName: Whispe
|
|
|
84
174
|
throw new Error(`Language detection is only supported with multilingual models.`)
|
|
85
175
|
}
|
|
86
176
|
|
|
87
|
-
if (temperature < 0) {
|
|
177
|
+
if (options.temperature! < 0) {
|
|
88
178
|
throw new Error(`Temperature cannot be negative`)
|
|
89
179
|
}
|
|
90
180
|
|
|
91
|
-
const
|
|
181
|
+
const encoderProviders: OnnxExecutionProvider[] =
|
|
182
|
+
options.encoderProvider ? [options.encoderProvider] : ['dml', 'cpu']
|
|
183
|
+
|
|
184
|
+
const decoderProviders: OnnxExecutionProvider[] =
|
|
185
|
+
options.decoderProvider ? [options.decoderProvider] : []
|
|
186
|
+
|
|
187
|
+
const whisper = new Whisper(
|
|
188
|
+
modelName,
|
|
189
|
+
modelDir,
|
|
190
|
+
encoderProviders,
|
|
191
|
+
decoderProviders)
|
|
92
192
|
|
|
93
193
|
async function detectLanguageForPart(partAudio: RawAudio) {
|
|
94
194
|
const audioFeatures = await whisper.encodeAudio(partAudio)
|
|
95
|
-
const partResults = await whisper.detectLanguage(audioFeatures, temperature)
|
|
195
|
+
const partResults = await whisper.detectLanguage(audioFeatures, options.temperature!)
|
|
96
196
|
|
|
97
197
|
return partResults
|
|
98
198
|
}
|
|
@@ -104,10 +204,65 @@ export async function detectLanguage(sourceRawAudio: RawAudio, modelName: Whispe
|
|
|
104
204
|
return results
|
|
105
205
|
}
|
|
106
206
|
|
|
107
|
-
export
|
|
108
|
-
|
|
109
|
-
|
|
207
|
+
export async function detectVoiceActivity(
|
|
208
|
+
sourceRawAudio: RawAudio,
|
|
209
|
+
modelName: WhisperModelName,
|
|
210
|
+
modelDir: string,
|
|
211
|
+
options: WhisperVADOptions) {
|
|
212
|
+
|
|
213
|
+
options = extendDeep(defaultWhisperVADOptions, options)
|
|
214
|
+
|
|
215
|
+
if (sourceRawAudio.sampleRate != 16000) {
|
|
216
|
+
throw new Error('Source audio must have a sampling rate of 16000')
|
|
217
|
+
}
|
|
218
|
+
|
|
219
|
+
if (options.temperature! < 0) {
|
|
220
|
+
throw new Error(`Temperature cannot be negative`)
|
|
221
|
+
}
|
|
222
|
+
|
|
223
|
+
const audioSamples = sourceRawAudio.audioChannels[0]
|
|
224
|
+
|
|
225
|
+
const partDuration = 5
|
|
226
|
+
const maxSamplesCountForPart = sourceRawAudio.sampleRate * partDuration
|
|
227
|
+
|
|
228
|
+
const encoderProviders: OnnxExecutionProvider[] =
|
|
229
|
+
options.encoderProvider ? [options.encoderProvider] : ['dml', 'cpu']
|
|
230
|
+
|
|
231
|
+
const decoderProviders: OnnxExecutionProvider[] =
|
|
232
|
+
options.decoderProvider ? [options.decoderProvider] : []
|
|
233
|
+
|
|
234
|
+
const whisper = new Whisper(
|
|
235
|
+
modelName,
|
|
236
|
+
modelDir,
|
|
237
|
+
encoderProviders,
|
|
238
|
+
decoderProviders)
|
|
239
|
+
|
|
240
|
+
const partProbabilities: Timeline = []
|
|
241
|
+
|
|
242
|
+
for (let sampleOffset = 0; sampleOffset < audioSamples.length; sampleOffset += maxSamplesCountForPart) {
|
|
243
|
+
const partSamples = sliceRawAudio(sourceRawAudio, sampleOffset, sampleOffset + maxSamplesCountForPart)
|
|
110
244
|
|
|
245
|
+
const samplesCountForPart = partSamples.audioChannels[0].length
|
|
246
|
+
|
|
247
|
+
const startTime = sampleOffset / sourceRawAudio.sampleRate
|
|
248
|
+
const endTime = (sampleOffset + samplesCountForPart) / sourceRawAudio.sampleRate
|
|
249
|
+
|
|
250
|
+
const encodedPartSamples = await whisper.encodeAudio(partSamples)
|
|
251
|
+
const probabilityForPart = await whisper.detectVoiceActivity(encodedPartSamples, options.temperature!)
|
|
252
|
+
|
|
253
|
+
partProbabilities.push({
|
|
254
|
+
type: 'segment',
|
|
255
|
+
text: '',
|
|
256
|
+
startTime,
|
|
257
|
+
endTime,
|
|
258
|
+
confidence: probabilityForPart,
|
|
259
|
+
})
|
|
260
|
+
}
|
|
261
|
+
|
|
262
|
+
return { partProbabilities }
|
|
263
|
+
}
|
|
264
|
+
|
|
265
|
+
export class Whisper {
|
|
111
266
|
isMultiligualModel: boolean
|
|
112
267
|
|
|
113
268
|
audioEncoder?: Onnx.InferenceSession
|
|
@@ -115,11 +270,6 @@ export class Whisper {
|
|
|
115
270
|
|
|
116
271
|
tiktoken?: Tiktoken
|
|
117
272
|
|
|
118
|
-
onnxOptions: Onnx.InferenceSession.SessionOptions = {
|
|
119
|
-
logSeverityLevel: 2,
|
|
120
|
-
executionProviders: ['cpu']
|
|
121
|
-
}
|
|
122
|
-
|
|
123
273
|
tokenConfig: {
|
|
124
274
|
endOfTextToken: number
|
|
125
275
|
startOfTextToken: number
|
|
@@ -139,9 +289,12 @@ export class Whisper {
|
|
|
139
289
|
|
|
140
290
|
randomGen: XorShift32RNG
|
|
141
291
|
|
|
142
|
-
constructor(
|
|
143
|
-
|
|
144
|
-
|
|
292
|
+
constructor(
|
|
293
|
+
public readonly modelName: WhisperModelName,
|
|
294
|
+
public readonly modelDir: string,
|
|
295
|
+
public readonly encoderExecutionProviders: OnnxExecutionProvider[],
|
|
296
|
+
public readonly decoderExecutionProviders: OnnxExecutionProvider[],
|
|
297
|
+
rngSeed = 461845907) {
|
|
145
298
|
|
|
146
299
|
this.isMultiligualModel = isMultilingualModel(this.modelName)
|
|
147
300
|
|
|
@@ -184,128 +337,40 @@ export class Whisper {
|
|
|
184
337
|
this.randomGen = new XorShift32RNG(rngSeed)
|
|
185
338
|
}
|
|
186
339
|
|
|
187
|
-
async
|
|
188
|
-
|
|
189
|
-
|
|
190
|
-
|
|
191
|
-
|
|
192
|
-
|
|
193
|
-
|
|
194
|
-
if (this.tiktoken) {
|
|
195
|
-
return
|
|
196
|
-
}
|
|
197
|
-
|
|
198
|
-
const logger = new Logger()
|
|
199
|
-
await logger.startAsync('Load tokenizer data')
|
|
200
|
-
|
|
201
|
-
const tiktokenModulePackagePath = await loadPackage('whisper-tiktoken-data')
|
|
202
|
-
|
|
203
|
-
const tiktokenDataFilePath = path.join(tiktokenModulePackagePath, this.isMultiligualModel ? 'multilingual.tiktoken' : 'gpt2.tiktoken')
|
|
204
|
-
let tiktokenData = await readFile(tiktokenDataFilePath, { encoding: 'utf8' })
|
|
205
|
-
|
|
206
|
-
const tokenConfig = this.tokenConfig
|
|
207
|
-
|
|
208
|
-
const metadataTokens: Record<number, string> = {
|
|
209
|
-
[tokenConfig.endOfTextToken]: '[EndOfText]',
|
|
210
|
-
[tokenConfig.startOfTextToken]: '[StartOfText]',
|
|
211
|
-
[tokenConfig.translateTaskToken]: '[TranslateTask]',
|
|
212
|
-
[tokenConfig.transcribeTaskToken]: '[TranscribeTask]',
|
|
213
|
-
[tokenConfig.startOfPromptToken]: '[StartOfPrompt]',
|
|
214
|
-
[tokenConfig.nonSpeechToken]: '[NonSpeech]',
|
|
215
|
-
[tokenConfig.noTimestampsToken]: '[NoTimestamps]',
|
|
216
|
-
}
|
|
217
|
-
|
|
218
|
-
if (this.isMultiligualModel) {
|
|
219
|
-
metadataTokens[50256] = '[Unused_50256]'
|
|
220
|
-
metadataTokens[50360] = '[Unused_50360]'
|
|
221
|
-
}
|
|
222
|
-
|
|
223
|
-
const languageTokenCount = tokenConfig.languageTokensEnd - tokenConfig.languageTokensStart
|
|
224
|
-
|
|
225
|
-
for (let i = 0; i < languageTokenCount; i++) {
|
|
226
|
-
const tokenIndex = this.tokenConfig.languageTokensStart + i
|
|
227
|
-
|
|
228
|
-
metadataTokens[tokenIndex] = `[Language_${i}]`
|
|
229
|
-
}
|
|
230
|
-
|
|
231
|
-
const timestampTokensCount = 1501
|
|
232
|
-
|
|
233
|
-
for (let i = 0; i < timestampTokensCount; i++) {
|
|
234
|
-
const tokenIndex = this.tokenConfig.timestampTokensStart + i
|
|
235
|
-
const tokenTime = this.timestampTokenToSeconds(tokenIndex)
|
|
236
|
-
|
|
237
|
-
metadataTokens[tokenIndex] = `[Timestamp_${tokenTime.toFixed(2)}]`
|
|
238
|
-
}
|
|
239
|
-
|
|
240
|
-
const inverseMetadataTokensLookup: Record<string, number> = {}
|
|
241
|
-
|
|
242
|
-
for (const [key, value] of Object.entries(metadataTokens)) {
|
|
243
|
-
inverseMetadataTokensLookup[value] = parseInt(key)
|
|
244
|
-
}
|
|
245
|
-
|
|
246
|
-
const patternString = `'s|'t|'re|'ve|'m|'ll|'d| ?\\p{L}+| ?\\p{N}+| ?[^\\s\\p{L}\\p{N}]+|\\s+(?!\\S)|\\s+`
|
|
247
|
-
|
|
248
|
-
const { Tiktoken } = await import('tiktoken/lite')
|
|
249
|
-
|
|
250
|
-
this.tiktoken = new Tiktoken(tiktokenData, inverseMetadataTokensLookup, patternString)
|
|
251
|
-
|
|
252
|
-
logger.end()
|
|
253
|
-
}
|
|
254
|
-
|
|
255
|
-
async initializeEncoderSessionIfNeeded() {
|
|
256
|
-
if (this.audioEncoder) {
|
|
257
|
-
return
|
|
258
|
-
}
|
|
259
|
-
|
|
260
|
-
const logger = new Logger()
|
|
261
|
-
|
|
262
|
-
await logger.startAsync(`Create encoder model inference session for model '${this.modelName}'`)
|
|
263
|
-
|
|
264
|
-
const encoderFilePath = path.join(this.modelDir, 'encoder.onnx')
|
|
265
|
-
|
|
266
|
-
const Onnx = await import('onnxruntime-node')
|
|
267
|
-
|
|
268
|
-
this.audioEncoder = await Onnx.InferenceSession.create(encoderFilePath, this.onnxOptions)
|
|
269
|
-
|
|
270
|
-
logger.end()
|
|
271
|
-
}
|
|
272
|
-
|
|
273
|
-
async initializeDecoderSessionIfNeeded() {
|
|
274
|
-
if (this.textDecoder) {
|
|
275
|
-
return
|
|
276
|
-
}
|
|
277
|
-
|
|
278
|
-
const logger = new Logger()
|
|
279
|
-
|
|
280
|
-
await logger.startAsync(`Create decoder model inference session for model '${this.modelName}'`)
|
|
281
|
-
|
|
282
|
-
const decoderFilePath = path.join(this.modelDir, 'decoder.onnx')
|
|
283
|
-
|
|
284
|
-
const Onnx = await import('onnxruntime-node')
|
|
285
|
-
|
|
286
|
-
this.textDecoder = await Onnx.InferenceSession.create(decoderFilePath, this.onnxOptions)
|
|
287
|
-
|
|
288
|
-
logger.end()
|
|
289
|
-
}
|
|
290
|
-
|
|
291
|
-
async recognize(rawAudio: RawAudio, task: WhisperTask, language: string, options: WhisperOptions) {
|
|
340
|
+
async recognize(
|
|
341
|
+
rawAudio: RawAudio,
|
|
342
|
+
task: WhisperTask,
|
|
343
|
+
language: string,
|
|
344
|
+
options: WhisperOptions,
|
|
345
|
+
logitFilter?: WhisperLogitFilter,
|
|
346
|
+
) {
|
|
292
347
|
await this.initializeIfNeeded()
|
|
293
348
|
|
|
294
349
|
const logger = new Logger()
|
|
295
350
|
|
|
351
|
+
options = extendDeep(defaultWhisperOptions, options)
|
|
352
|
+
options.model = this.modelName
|
|
353
|
+
|
|
296
354
|
const audioSamples = rawAudio.audioChannels[0]
|
|
297
355
|
const sampleRate = rawAudio.sampleRate
|
|
298
356
|
const prompt = options.prompt
|
|
357
|
+
const decodeTimestampTokens = options.decodeTimestampTokens!
|
|
299
358
|
|
|
300
359
|
const maxAudioSamplesPerPart = sampleRate * 30
|
|
301
360
|
|
|
302
|
-
const decodeTimestampTokens = options.decodeTimestampTokens!
|
|
303
|
-
|
|
304
361
|
let previousPartTextTokens: number[] = []
|
|
305
362
|
|
|
306
363
|
let timeline: Timeline = []
|
|
307
364
|
let allDecodedTokens: number[] = []
|
|
308
365
|
|
|
366
|
+
let wrappedLogitFilter: WhisperLogitFilter | undefined
|
|
367
|
+
|
|
368
|
+
if (logitFilter) {
|
|
369
|
+
wrappedLogitFilter = (logits, partDecodedTokens, isFirstPart, isFinalPart) => {
|
|
370
|
+
return logitFilter(logits, [...allDecodedTokens, ...partDecodedTokens], isFirstPart, isFinalPart)
|
|
371
|
+
}
|
|
372
|
+
}
|
|
373
|
+
|
|
309
374
|
for (let audioOffset = 0; audioOffset < audioSamples.length;) {
|
|
310
375
|
const segmentStartTime = audioOffset / sampleRate
|
|
311
376
|
|
|
@@ -338,9 +403,17 @@ export class Whisper {
|
|
|
338
403
|
|
|
339
404
|
let {
|
|
340
405
|
decodedTokens: partTokens,
|
|
341
|
-
|
|
342
|
-
|
|
343
|
-
} = await this.decodeTokens(
|
|
406
|
+
decodedTokensConfidence: partTokensConfidence,
|
|
407
|
+
decodedTokensCrossAttentionQKs: partCrossAttentionQKs,
|
|
408
|
+
} = await this.decodeTokens(
|
|
409
|
+
audioPartFeatures,
|
|
410
|
+
initialTokens,
|
|
411
|
+
audioPartDuration,
|
|
412
|
+
isFirstPart,
|
|
413
|
+
isFinalPart,
|
|
414
|
+
options,
|
|
415
|
+
wrappedLogitFilter,
|
|
416
|
+
)
|
|
344
417
|
|
|
345
418
|
const lastToken = partTokens[partTokens.length - 1]
|
|
346
419
|
const lastTokenIsTimestamp = this.isTimestampToken(lastToken)
|
|
@@ -364,27 +437,36 @@ export class Whisper {
|
|
|
364
437
|
throw new Error('Unexpected: partTokens.length != partCrossAttentionQKs.length')
|
|
365
438
|
}
|
|
366
439
|
|
|
440
|
+
// Prepare tokens
|
|
367
441
|
partTokens = partTokens.slice(initialTokens.length)
|
|
368
|
-
|
|
369
|
-
//const compressionRatioForPart = (await getDeflateCompressionMetricsForString(this.tokensToText(partTokens))).ratio
|
|
370
|
-
|
|
442
|
+
partTokensConfidence = partTokensConfidence.slice(initialTokens.length)
|
|
371
443
|
partCrossAttentionQKs = partCrossAttentionQKs.slice(initialTokens.length)
|
|
372
444
|
|
|
445
|
+
// Compute compression ratio for part
|
|
446
|
+
if (false) {
|
|
447
|
+
const compressionRatioForPart = (await getDeflateCompressionMetricsForString(this.tokensToText(partTokens))).ratio
|
|
448
|
+
}
|
|
449
|
+
|
|
450
|
+
// Find alignment path
|
|
373
451
|
const alignmentPath = await this.findAlignmentPathFromQKs(partCrossAttentionQKs, partTokens, 0, segmentFrameCount) //, alignmentHeadsIndexes[this.modelName])
|
|
374
|
-
const partTimeline = await this.getTokenTimelineFromAlignmentPath(alignmentPath, partTokens, segmentStartTime, segmentEndTime, decodedTokensConfidence)
|
|
375
452
|
|
|
376
|
-
|
|
453
|
+
// Generate timeline from alignment path
|
|
454
|
+
const partTimeline = await this.getTokenTimelineFromAlignmentPath(alignmentPath, partTokens, segmentStartTime, segmentEndTime, partTokensConfidence)
|
|
377
455
|
|
|
378
456
|
allDecodedTokens.push(...partTokens)
|
|
379
457
|
timeline.push(...partTimeline)
|
|
380
458
|
|
|
381
459
|
previousPartTextTokens = partTokens.filter(token => this.isTextToken(token))
|
|
382
460
|
|
|
461
|
+
audioOffset = audioEndOffset
|
|
462
|
+
|
|
383
463
|
logger.end()
|
|
384
464
|
}
|
|
385
465
|
|
|
466
|
+
// Convert token timeline to word timeline
|
|
386
467
|
timeline = this.tokenTimelineToWordTimeline(timeline, language)
|
|
387
468
|
|
|
469
|
+
// Convert tokens to transcript
|
|
388
470
|
const transcript = this.tokensToText(allDecodedTokens).trim()
|
|
389
471
|
|
|
390
472
|
logger.end()
|
|
@@ -392,42 +474,82 @@ export class Whisper {
|
|
|
392
474
|
return { transcript, timeline }
|
|
393
475
|
}
|
|
394
476
|
|
|
395
|
-
async align(rawAudio: RawAudio,
|
|
396
|
-
await this.
|
|
477
|
+
async align(rawAudio: RawAudio, transcript: string, sourceLanguage: string, task: 'transcribe' | 'translate', whisperAlignmentOptions: WhisperAlignmentOptions) {
|
|
478
|
+
await this.initializeTokenizerIfNeeded()
|
|
397
479
|
|
|
398
|
-
|
|
480
|
+
whisperAlignmentOptions = extendDeep(defaultWhisperAlignmentOptions, whisperAlignmentOptions)
|
|
481
|
+
|
|
482
|
+
const shouldSplitToSentences = false
|
|
399
483
|
|
|
400
|
-
|
|
484
|
+
let simplifiedTranscript = ''
|
|
401
485
|
|
|
402
|
-
|
|
486
|
+
if (shouldSplitToSentences) {
|
|
487
|
+
const sentences = splitToSentences(transcript, 'en')
|
|
403
488
|
|
|
404
|
-
|
|
405
|
-
|
|
489
|
+
for (const sentence of sentences) {
|
|
490
|
+
let sentenceWords = await splitToWords(sentence, 'en')
|
|
491
|
+
sentenceWords = sentenceWords.filter(word => isWord(word))
|
|
492
|
+
|
|
493
|
+
simplifiedTranscript += sentenceWords.join(' ')
|
|
494
|
+
simplifiedTranscript += ' '
|
|
495
|
+
}
|
|
496
|
+
} else {
|
|
497
|
+
let words = await splitToWords(transcript, 'en')
|
|
498
|
+
words = words.map(word => word.trim())
|
|
499
|
+
words = words.filter(word => isWord(word))
|
|
500
|
+
simplifiedTranscript = words.join(' ')
|
|
501
|
+
}
|
|
406
502
|
|
|
407
|
-
|
|
503
|
+
// Tokenize the transcript
|
|
504
|
+
const simplifiedTranscriptTokens = this.textToTokens(simplifiedTranscript)
|
|
408
505
|
|
|
506
|
+
// Initialize custom logit filter that allows only the transcript tokens to be decoded
|
|
507
|
+
// in order.
|
|
409
508
|
const endOfTextToken = this.tokenConfig.endOfTextToken
|
|
410
509
|
|
|
411
|
-
|
|
510
|
+
const logitFilter: WhisperLogitFilter = (logits, decodedTokens, isFirstPart, isFinalPart) => {
|
|
511
|
+
const decodedTextTokens = decodedTokens.filter(token => this.isTextToken(token))
|
|
412
512
|
|
|
413
|
-
|
|
414
|
-
|
|
513
|
+
const nextTokenToDecode = simplifiedTranscriptTokens[decodedTextTokens.length] ?? endOfTextToken
|
|
514
|
+
|
|
515
|
+
const newLogits = logits.map((logit, index) => {
|
|
516
|
+
if (index === nextTokenToDecode) {
|
|
517
|
+
return logit
|
|
518
|
+
}
|
|
415
519
|
|
|
416
|
-
|
|
417
|
-
|
|
520
|
+
// If it's the final part, the ent-of-text token logit is set to -Infinity.
|
|
521
|
+
// This will force to force all transcript tokens to be decoded even if the model doesn't
|
|
522
|
+
// recognize them.
|
|
523
|
+
if (!isFinalPart && index === endOfTextToken) {
|
|
524
|
+
return logit
|
|
525
|
+
}
|
|
418
526
|
|
|
419
|
-
|
|
420
|
-
|
|
527
|
+
return -Infinity
|
|
528
|
+
})
|
|
421
529
|
|
|
422
|
-
|
|
423
|
-
|
|
424
|
-
const tokenTimeline = await this.getTokenTimelineFromAlignmentPath(alignmentPath, tokens, 0, audioDuration)
|
|
530
|
+
return newLogits
|
|
531
|
+
}
|
|
425
532
|
|
|
426
|
-
|
|
533
|
+
// Set options for alignment
|
|
534
|
+
const options: WhisperOptions = {
|
|
535
|
+
model: this.modelName,
|
|
536
|
+
temperature: 0.0,
|
|
537
|
+
prompt: undefined,
|
|
538
|
+
topCandidateCount: 1,
|
|
539
|
+
punctuationThreshold: Infinity,
|
|
540
|
+
autoPromptParts: false,
|
|
541
|
+
maxTokensPerPart: Infinity,
|
|
542
|
+
suppressRepetition: false,
|
|
543
|
+
decodeTimestampTokens: true,
|
|
544
|
+
endTokenThreshold: whisperAlignmentOptions!.endTokenThreshold!,
|
|
545
|
+
includeEndTokenInCandidates: false,
|
|
546
|
+
seed: undefined,
|
|
547
|
+
}
|
|
427
548
|
|
|
428
|
-
|
|
549
|
+
// Recognize
|
|
550
|
+
const { timeline } = await this.recognize(rawAudio, task, sourceLanguage, options, logitFilter)
|
|
429
551
|
|
|
430
|
-
return
|
|
552
|
+
return timeline
|
|
431
553
|
}
|
|
432
554
|
|
|
433
555
|
async detectLanguage(audioFeatures: Onnx.Tensor, temperature: number): Promise<LanguageDetectionResults> {
|
|
@@ -435,7 +557,6 @@ export class Whisper {
|
|
|
435
557
|
throw new Error('Language detection is only supported with multilingual models')
|
|
436
558
|
}
|
|
437
559
|
|
|
438
|
-
await this.initializeTokenizerIfNeeded()
|
|
439
560
|
await this.initializeDecoderSessionIfNeeded()
|
|
440
561
|
|
|
441
562
|
// Prepare and run decoder
|
|
@@ -488,62 +609,113 @@ export class Whisper {
|
|
|
488
609
|
return results
|
|
489
610
|
}
|
|
490
611
|
|
|
612
|
+
async detectVoiceActivity(audioFeatures: Onnx.Tensor, temperature: number): Promise<number> {
|
|
613
|
+
await this.initializeDecoderSessionIfNeeded()
|
|
614
|
+
|
|
615
|
+
// Prepare and run decoder
|
|
616
|
+
const logger = new Logger()
|
|
617
|
+
await logger.startAsync('Detect voice activity with Whisper model')
|
|
618
|
+
|
|
619
|
+
const sotToken = this.tokenConfig.startOfTextToken
|
|
620
|
+
|
|
621
|
+
const initialTokens = [sotToken]
|
|
622
|
+
const offset = 0
|
|
623
|
+
|
|
624
|
+
const Onnx = await import('onnxruntime-node')
|
|
625
|
+
|
|
626
|
+
const initialKvDimensions = this.getKvDimensions(1, initialTokens.length)
|
|
627
|
+
const kvCacheTensor = new Onnx.Tensor('float32', new Float32Array(initialKvDimensions[0] * initialKvDimensions[1] * initialKvDimensions[2] * initialKvDimensions[3]), initialKvDimensions)
|
|
628
|
+
|
|
629
|
+
const tokensTensor = new Onnx.Tensor('int64', new BigInt64Array(initialTokens.map(token => BigInt(token))), [1, initialTokens.length])
|
|
630
|
+
const offsetTensor = new Onnx.Tensor('int64', new BigInt64Array([BigInt(offset)]), [])
|
|
631
|
+
|
|
632
|
+
const decoderInputs = {
|
|
633
|
+
tokens: tokensTensor,
|
|
634
|
+
audio_features: audioFeatures,
|
|
635
|
+
kv_cache: kvCacheTensor,
|
|
636
|
+
offset: offsetTensor
|
|
637
|
+
}
|
|
638
|
+
|
|
639
|
+
const decoderOutputs = await this.textDecoder!.run(decoderInputs)
|
|
640
|
+
const logitsBuffer = decoderOutputs['logits'].data as Float32Array
|
|
641
|
+
|
|
642
|
+
const tokenConfig = this.tokenConfig
|
|
643
|
+
|
|
644
|
+
const logits = Array.from(logitsBuffer)
|
|
645
|
+
|
|
646
|
+
const probabilities = softmax(logits, temperature)
|
|
647
|
+
|
|
648
|
+
const noSpeechProbability = probabilities[tokenConfig.nonSpeechToken]
|
|
649
|
+
|
|
650
|
+
return 1.0 - noSpeechProbability
|
|
651
|
+
}
|
|
652
|
+
|
|
653
|
+
// Decode tokens using the decoder model
|
|
491
654
|
async decodeTokens(
|
|
492
655
|
audioFeatures: Onnx.Tensor,
|
|
493
656
|
initialTokens: number[],
|
|
494
657
|
audioDuration: number,
|
|
495
658
|
isFirstPart: boolean,
|
|
496
659
|
isFinalPart: boolean,
|
|
497
|
-
options: WhisperOptions
|
|
660
|
+
options: WhisperOptions,
|
|
661
|
+
logitFilter?: WhisperLogitFilter) {
|
|
498
662
|
|
|
663
|
+
// Initialize
|
|
499
664
|
await this.initializeTokenizerIfNeeded()
|
|
500
665
|
await this.initializeDecoderSessionIfNeeded()
|
|
501
666
|
|
|
502
667
|
const logger = new Logger()
|
|
503
668
|
|
|
504
|
-
const allowedPunctuationMarks = this.getAllowedPunctuationMarks()
|
|
505
|
-
|
|
506
669
|
await logger.startAsync('Decode text tokens with Whisper decoder model')
|
|
507
670
|
|
|
508
671
|
options = extendDeep(defaultWhisperOptions, options)
|
|
509
672
|
|
|
510
673
|
const Onnx = await import('onnxruntime-node')
|
|
511
674
|
|
|
675
|
+
// Get token information
|
|
512
676
|
const endOfTextToken = this.tokenConfig.endOfTextToken
|
|
513
|
-
|
|
514
677
|
const timestampTokensStart = this.tokenConfig.timestampTokensStart
|
|
515
|
-
const suppressedTokens = new Set(this.getSuppressedTokens())
|
|
516
678
|
|
|
517
|
-
const
|
|
679
|
+
const suppressedTextTokens = this.getSuppressedTextTokens()
|
|
680
|
+
const suppressedMetadataTokens = this.getSuppressedMetadataTokens()
|
|
681
|
+
|
|
682
|
+
const allowedPunctuationMarks = this.getAllowedPunctuationMarks()
|
|
518
683
|
|
|
519
|
-
const
|
|
684
|
+
const spaceToken = this.textToTokens(' ')[0]
|
|
520
685
|
|
|
686
|
+
// Initialize variables for decoding loop
|
|
521
687
|
let decodedTokens = initialTokens.slice()
|
|
522
688
|
const initialKvDimensions = this.getKvDimensions(1, decodedTokens.length)
|
|
523
689
|
let kvCacheTensor = new Onnx.Tensor('float32', new Float32Array(initialKvDimensions[0] * initialKvDimensions[1] * initialKvDimensions[2] * initialKvDimensions[3]), initialKvDimensions)
|
|
524
690
|
|
|
525
|
-
let decodedTokensTimestampLogits: number[][] = [
|
|
526
|
-
|
|
527
|
-
let
|
|
528
|
-
|
|
529
|
-
let timestampsSeenCount = 0
|
|
530
|
-
|
|
531
|
-
const decodedTokensConfidence: number[] = []
|
|
532
|
-
let decodedTokensCrossAttentionQKs: Onnx.Tensor[] = []
|
|
691
|
+
let decodedTokensTimestampLogits: number[][] = []
|
|
692
|
+
let decodedTokensConfidence: number[] = []
|
|
693
|
+
let decodedTokensCrossAttentionQKs: OnnxLikeFloat32Tensor[] = []
|
|
533
694
|
|
|
534
695
|
for (let i = 0; i < decodedTokens.length; i++) {
|
|
696
|
+
decodedTokensTimestampLogits.push(new Array(1501))
|
|
697
|
+
decodedTokensConfidence.push(1.0)
|
|
535
698
|
decodedTokensCrossAttentionQKs.push(undefined as any)
|
|
536
699
|
}
|
|
537
700
|
|
|
701
|
+
let lastTimestampTokenIndex = -1
|
|
702
|
+
let timestampTokenSeenCount = 0
|
|
703
|
+
let bufferedTokensToPrint: number[] = []
|
|
704
|
+
|
|
705
|
+
// Define method to add a token to output
|
|
706
|
+
function addToken(tokenToAdd: number, timestampLogits: number[], confidence: number, crossAttentionQKs: OnnxLikeFloat32Tensor) {
|
|
707
|
+
decodedTokens.push(tokenToAdd)
|
|
708
|
+
decodedTokensTimestampLogits.push(timestampLogits)
|
|
709
|
+
decodedTokensConfidence.push(confidence)
|
|
710
|
+
decodedTokensCrossAttentionQKs.push(crossAttentionQKs)
|
|
711
|
+
}
|
|
712
|
+
|
|
538
713
|
// Start decoding loop
|
|
539
|
-
for (let decodedTokenCount = 0; decodedTokenCount <
|
|
714
|
+
for (let decodedTokenCount = 0; decodedTokenCount < options.maxTokensPerPart!; decodedTokenCount++) {
|
|
540
715
|
const isInitialState = decodedTokens.length == initialTokens.length
|
|
541
716
|
|
|
542
|
-
|
|
543
|
-
const offset = isInitialState ? 0 : decodedTokens.length
|
|
544
|
-
|
|
717
|
+
// If not in initial state, reshape KV Cache tensor to accomodate a new output token
|
|
545
718
|
if (!isInitialState) {
|
|
546
|
-
// Reshape KV Cache tensor
|
|
547
719
|
const dims = kvCacheTensor.dims
|
|
548
720
|
|
|
549
721
|
const currentKvCacheGroups = splitFloat32Array(kvCacheTensor.data as Float32Array, dims[2] * dims[3])
|
|
@@ -558,277 +730,327 @@ export class Whisper {
|
|
|
558
730
|
kvCacheTensor = reshapedKvCacheTensor
|
|
559
731
|
}
|
|
560
732
|
|
|
561
|
-
// Prepare
|
|
733
|
+
// Prepare values for decoder
|
|
734
|
+
const tokensToDecode = isInitialState ? decodedTokens : [decodedTokens[decodedTokens.length - 1]]
|
|
735
|
+
const offset = isInitialState ? 0 : decodedTokens.length
|
|
736
|
+
|
|
562
737
|
const tokensTensor = new Onnx.Tensor('int64', new BigInt64Array(tokensToDecode.map(token => BigInt(token))), [1, tokensToDecode.length])
|
|
563
738
|
const offsetTensor = new Onnx.Tensor('int64', new BigInt64Array([BigInt(offset)]), [])
|
|
564
739
|
|
|
565
|
-
const decoderInputs = {
|
|
740
|
+
const decoderInputs = {
|
|
741
|
+
tokens: tokensTensor,
|
|
742
|
+
audio_features: audioFeatures,
|
|
743
|
+
kv_cache: kvCacheTensor,
|
|
744
|
+
offset: offsetTensor
|
|
745
|
+
}
|
|
566
746
|
|
|
747
|
+
// Run decoder
|
|
567
748
|
const decoderOutputs = await this.textDecoder!.run(decoderInputs)
|
|
568
749
|
|
|
750
|
+
// Store results
|
|
569
751
|
const logitsBuffer = decoderOutputs['logits'].data as Float32Array
|
|
570
752
|
kvCacheTensor = decoderOutputs['output_kv_cache'] as any
|
|
571
753
|
|
|
572
|
-
|
|
573
|
-
const
|
|
574
|
-
|
|
575
|
-
const timestampTokenLogits = allTokenLogits.slice(timestampTokensStart)
|
|
576
|
-
|
|
577
|
-
// Suppress tokens
|
|
578
|
-
for (let logitIndex = 0; logitIndex < allTokenLogits.length; logitIndex++) {
|
|
579
|
-
const isWrongTokenForInitialState =
|
|
580
|
-
isInitialState &&
|
|
581
|
-
(logitIndex === spaceToken || logitIndex === endOfTextToken)
|
|
582
|
-
|
|
583
|
-
const isInSuppressedList = suppressedTokens.has(logitIndex)
|
|
754
|
+
const crossAttentionQKsForTokenOnnx = decoderOutputs['cross_attention_qks']
|
|
755
|
+
const crossAttentionQKsForToken = makeOnnxLikeFloat32Tensor(crossAttentionQKsForTokenOnnx)
|
|
756
|
+
crossAttentionQKsForTokenOnnx.dispose()
|
|
584
757
|
|
|
585
|
-
|
|
758
|
+
// Get logits
|
|
759
|
+
const resultLogitsFloatArrays = splitFloat32Array(logitsBuffer, logitsBuffer.length / decoderOutputs['logits'].dims[1])
|
|
760
|
+
const allTokenLogits = Array.from(resultLogitsFloatArrays[resultLogitsFloatArrays.length - 1])
|
|
586
761
|
|
|
587
|
-
|
|
588
|
-
|
|
589
|
-
|
|
762
|
+
// Suppress metadata tokens in the suppression set
|
|
763
|
+
for (const suppressedTokenIndex of suppressedMetadataTokens) {
|
|
764
|
+
allTokenLogits[suppressedTokenIndex] = -Infinity
|
|
590
765
|
}
|
|
591
766
|
|
|
592
|
-
|
|
593
|
-
|
|
594
|
-
|
|
595
|
-
// Add best token
|
|
596
|
-
function addToken(tokenToAdd: number, timestampLogits: number[], confidence: number) {
|
|
597
|
-
decodedTokens.push(tokenToAdd)
|
|
598
|
-
decodedTokensTimestampLogits.push(timestampLogits)
|
|
599
|
-
decodedTokensCrossAttentionQKs.push(decoderOutputs['cross_attention_qks'])
|
|
600
|
-
decodedTokensConfidence.push(confidence)
|
|
767
|
+
if (isInitialState) {
|
|
768
|
+
// If in initial state, suppress end-of-text token
|
|
769
|
+
allTokenLogits[endOfTextToken] = -Infinity
|
|
601
770
|
}
|
|
602
771
|
|
|
603
|
-
|
|
772
|
+
const timestampTokenLogits = allTokenLogits.slice(timestampTokensStart)
|
|
604
773
|
|
|
605
|
-
|
|
606
|
-
|
|
607
|
-
const logProbabilities = logOfVector(probabilities)
|
|
774
|
+
const decodeTimestampTokenIfNeeded = () => {
|
|
775
|
+
// Try to decode a timestamp token, if needed
|
|
608
776
|
|
|
609
|
-
|
|
777
|
+
// If timestamp tokens is disabled in options, don't decode a timestamp
|
|
778
|
+
if (!options.decodeTimestampTokens) {
|
|
779
|
+
return false
|
|
780
|
+
}
|
|
781
|
+
|
|
782
|
+
const previousTokenWasTimestamp = this.isTimestampToken(decodedTokens[decodedTokens.length - 1])
|
|
783
|
+
const secondPreviousTokenWasTimestamp = this.isTimestampToken(decodedTokens[decodedTokens.length - 2])
|
|
784
|
+
|
|
785
|
+
// If there are two successive timestamp tokens decoded, or the previous timestamp was the first token,
|
|
786
|
+
// don't decode a timestamp
|
|
787
|
+
if (previousTokenWasTimestamp &&
|
|
788
|
+
(decodedTokens.length === initialTokens.length + 1) || secondPreviousTokenWasTimestamp) {
|
|
789
|
+
return false
|
|
790
|
+
}
|
|
791
|
+
|
|
792
|
+
// Derive token probabilities
|
|
793
|
+
const probabilities = softmax(allTokenLogits as any, 1.0)
|
|
794
|
+
const logProbabilities = logOfVector(probabilities)
|
|
610
795
|
|
|
796
|
+
const nonTimestampTokenLogProbs = logProbabilities.slice(0, timestampTokensStart)
|
|
797
|
+
|
|
798
|
+
// Find highest non-timestamp token
|
|
611
799
|
const indexOfMaxNonTimestampLogProb = indexOfMax(nonTimestampTokenLogProbs)
|
|
612
800
|
const valueOfMaxNonTimestampLogProb = nonTimestampTokenLogProbs[indexOfMaxNonTimestampLogProb]
|
|
613
801
|
|
|
802
|
+
// Find highest timestamp token
|
|
614
803
|
const timestampTokenLogProbs = logProbabilities.slice(timestampTokensStart)
|
|
615
804
|
const indexOfMaxTimestampLogProb = indexOfMax(timestampTokenLogProbs)
|
|
616
805
|
|
|
806
|
+
// Compute the log of the sum of exponentials of the log probabilities
|
|
807
|
+
// of the timestamp tokens
|
|
617
808
|
const logSumExpOfTimestampTokenLogProbs = logSumExp(timestampTokenLogProbs)
|
|
618
809
|
|
|
619
|
-
|
|
620
|
-
|
|
621
|
-
|
|
622
|
-
|
|
623
|
-
|
|
624
|
-
if (shouldDecodeTimestampToken && !previousTokenWasTimestamp) {
|
|
625
|
-
timestampsSeenCount += 1
|
|
810
|
+
// If the sum isn't greater than the log probability of the highest non-timestamp token,
|
|
811
|
+
// don't decode a timestamp
|
|
812
|
+
if (logSumExpOfTimestampTokenLogProbs <= valueOfMaxNonTimestampLogProb) {
|
|
813
|
+
return false
|
|
626
814
|
}
|
|
627
815
|
|
|
628
|
-
|
|
629
|
-
|
|
630
|
-
const previousToken = decodedTokens[decodedTokens.length - 1]
|
|
631
|
-
const previousTokenTimestampLogits = decodedTokensTimestampLogits[decodedTokensTimestampLogits.length - 1]
|
|
632
|
-
const previousTokenConfidence = decodedTokensConfidence[decodedTokensConfidence.length - 1]
|
|
633
|
-
|
|
634
|
-
addToken(previousToken, previousTokenTimestampLogits, previousTokenConfidence)
|
|
816
|
+
// Decode a timestamp token
|
|
817
|
+
timestampTokenSeenCount += 1
|
|
635
818
|
|
|
636
|
-
|
|
637
|
-
|
|
638
|
-
|
|
819
|
+
if (previousTokenWasTimestamp) {
|
|
820
|
+
const previousToken = decodedTokens[decodedTokens.length - 1]
|
|
821
|
+
const previousTokenTimestampLogits = decodedTokensTimestampLogits[decodedTokensTimestampLogits.length - 1]
|
|
822
|
+
const previousTokenConfidence = decodedTokensConfidence[decodedTokensConfidence.length - 1]
|
|
639
823
|
|
|
640
|
-
|
|
641
|
-
break
|
|
642
|
-
}
|
|
643
|
-
} else {
|
|
644
|
-
const timestampToken = timestampTokensStart + indexOfMaxTimestampLogProb
|
|
645
|
-
const confidence = probabilities[timestampToken]
|
|
824
|
+
addToken(previousToken, previousTokenTimestampLogits, previousTokenConfidence, crossAttentionQKsForToken)
|
|
646
825
|
|
|
647
|
-
|
|
648
|
-
|
|
826
|
+
lastTimestampTokenIndex = decodedTokens.length
|
|
827
|
+
} else {
|
|
828
|
+
const timestampToken = timestampTokensStart + indexOfMaxTimestampLogProb
|
|
829
|
+
const confidence = probabilities[timestampToken]
|
|
649
830
|
|
|
650
|
-
|
|
831
|
+
addToken(timestampToken, timestampTokenLogits, confidence, crossAttentionQKsForToken)
|
|
651
832
|
}
|
|
652
|
-
}
|
|
653
|
-
|
|
654
|
-
if (shouldDecodeNonTimestampToken) {
|
|
655
|
-
const topLogitCount = options.topCandidateCount!
|
|
656
833
|
|
|
657
|
-
|
|
834
|
+
return true
|
|
835
|
+
}
|
|
658
836
|
|
|
659
|
-
|
|
660
|
-
Array.from(nonTimestampTokenLogits).map((logit, index) => ({ token: index, logit }))
|
|
837
|
+
const timestampTokenDecoded = decodeTimestampTokenIfNeeded()
|
|
661
838
|
|
|
662
|
-
|
|
839
|
+
if (timestampTokenDecoded) {
|
|
840
|
+
await yieldToEventLoop()
|
|
663
841
|
|
|
664
|
-
|
|
665
|
-
|
|
666
|
-
token: entry.token,
|
|
667
|
-
logit: entry.logit,
|
|
668
|
-
text: this.tokenToText(entry.token, true)
|
|
669
|
-
}))
|
|
842
|
+
continue
|
|
843
|
+
}
|
|
670
844
|
|
|
671
|
-
|
|
672
|
-
|
|
673
|
-
const topCandidatesRepetitionScores = topCandidates.map(entry => {
|
|
674
|
-
const lastDecodedTextTokens = decodedTokens.filter(token => this.isTextToken(token)).reverse().slice(0, 20)
|
|
675
|
-
const { maxScore } = getRepetitionScoreRelativeToFirstSubstring([entry.token, ...lastDecodedTextTokens])
|
|
845
|
+
// Decode a non-timestamp token
|
|
846
|
+
let nonTimestampTokenLogits = allTokenLogits.slice(0, timestampTokensStart)
|
|
676
847
|
|
|
677
|
-
|
|
678
|
-
})
|
|
848
|
+
let shouldDecodeEndfOfTextToken = false
|
|
679
849
|
|
|
680
|
-
|
|
850
|
+
// If not in initial state, and the end-of-text token's probability is sufficiently higher than
|
|
851
|
+
// the second highest ranked token, then set to accept it
|
|
852
|
+
if (!isInitialState) {
|
|
853
|
+
const endOfTextTokenLogit = nonTimestampTokenLogits[endOfTextToken]
|
|
681
854
|
|
|
682
|
-
|
|
683
|
-
|
|
684
|
-
topCandidates = [topCandidates[indexOfMaxScore]]
|
|
685
|
-
} else {
|
|
686
|
-
topCandidates = topCandidates.filter((candidate, index) => topCandidatesRepetitionScores[index] < thresholdRepetitionScore)
|
|
687
|
-
}
|
|
688
|
-
}
|
|
689
|
-
////
|
|
855
|
+
const otherTokensLogits = nonTimestampTokenLogits.slice()
|
|
856
|
+
otherTokensLogits[endOfTextToken] = -Infinity
|
|
690
857
|
|
|
691
|
-
const
|
|
858
|
+
const indexOfMaximumOtherTokenLogit = indexOfMax(otherTokensLogits)
|
|
859
|
+
const maximumOtherTokenLogit = nonTimestampTokenLogits[indexOfMaximumOtherTokenLogit]
|
|
692
860
|
|
|
693
|
-
|
|
694
|
-
if (options.decodeTimestampTokens === false) {
|
|
695
|
-
topCandidates = topCandidates.filter((candidate, index) => {
|
|
696
|
-
if (candidate.token === endOfTextToken) {
|
|
697
|
-
return topCandidateProbabilities[index] >= 0.9
|
|
698
|
-
}
|
|
861
|
+
const endProbabilities = softmax([endOfTextTokenLogit, maximumOtherTokenLogit], 1.0)
|
|
699
862
|
|
|
700
|
-
|
|
701
|
-
|
|
863
|
+
if (endProbabilities[0] > options.endTokenThreshold!) {
|
|
864
|
+
shouldDecodeEndfOfTextToken = true
|
|
702
865
|
}
|
|
703
|
-
|
|
866
|
+
}
|
|
704
867
|
|
|
705
|
-
|
|
706
|
-
|
|
868
|
+
if (logitFilter) {
|
|
869
|
+
// Apply custom logit filter function if given
|
|
870
|
+
nonTimestampTokenLogits = logitFilter(nonTimestampTokenLogits, decodedTokens, isFirstPart, isFinalPart)
|
|
707
871
|
|
|
708
|
-
|
|
872
|
+
// If the custom filter set the end-of-text token to be Infinity, or -Infinity,
|
|
873
|
+
// then override any previous decision and accept or reject it, respectively
|
|
874
|
+
if (nonTimestampTokenLogits[endOfTextToken] === Infinity) {
|
|
875
|
+
shouldDecodeEndfOfTextToken = true
|
|
876
|
+
} else if (nonTimestampTokenLogits[endOfTextToken] === -Infinity) {
|
|
877
|
+
shouldDecodeEndfOfTextToken = false
|
|
878
|
+
}
|
|
709
879
|
|
|
710
|
-
|
|
711
|
-
|
|
712
|
-
|
|
880
|
+
// If filter caused all word token logits to be -Infinity, then there is no
|
|
881
|
+
// other token to decode. Fall back to accept end-of-text
|
|
882
|
+
if (nonTimestampTokenLogits.slice(0, endOfTextToken).every(logit => logit === -Infinity)) {
|
|
883
|
+
shouldDecodeEndfOfTextToken = true
|
|
884
|
+
}
|
|
885
|
+
} else {
|
|
886
|
+
// Otherwise, suppress text tokens in the suppression set
|
|
887
|
+
for (const suppressedTokenIndex of suppressedTextTokens) {
|
|
888
|
+
nonTimestampTokenLogits[suppressedTokenIndex] = -Infinity
|
|
889
|
+
}
|
|
713
890
|
|
|
714
|
-
|
|
891
|
+
// Suppress space token if at initial state
|
|
892
|
+
if (isInitialState) {
|
|
893
|
+
nonTimestampTokenLogits[spaceToken] = -Infinity
|
|
894
|
+
}
|
|
895
|
+
}
|
|
715
896
|
|
|
716
|
-
|
|
717
|
-
|
|
897
|
+
// If end-of-text token should be decoded, then add it and break
|
|
898
|
+
// out of the loop
|
|
899
|
+
if (shouldDecodeEndfOfTextToken) {
|
|
900
|
+
addToken(endOfTextToken, timestampTokenLogits, 1.0, crossAttentionQKsForToken)
|
|
718
901
|
|
|
719
|
-
|
|
902
|
+
break
|
|
903
|
+
}
|
|
720
904
|
|
|
721
|
-
|
|
722
|
-
|
|
723
|
-
|
|
905
|
+
// Suppress end-of-text if it shouldn't be included in candidates
|
|
906
|
+
if (!options.includeEndTokenInCandidates) {
|
|
907
|
+
nonTimestampTokenLogits[endOfTextToken] = -Infinity
|
|
908
|
+
}
|
|
724
909
|
|
|
725
|
-
|
|
910
|
+
// Find top candidates
|
|
911
|
+
const sortedNonTimestampLogitsWithIndexes =
|
|
912
|
+
Array.from(nonTimestampTokenLogits).map((logit, index) => ({ token: index, logit }))
|
|
726
913
|
|
|
727
|
-
|
|
728
|
-
rankOfPromisingPunctuationToken < rankOfSpaceToken) {
|
|
729
|
-
chosenCandidateRank = rankOfPromisingPunctuationToken
|
|
730
|
-
} else {
|
|
731
|
-
chosenCandidateRank = this.randomGen.selectRandomIndexFromDistribution(topCandidateProbabilities)
|
|
732
|
-
}
|
|
914
|
+
sortedNonTimestampLogitsWithIndexes.sort((a, b) => b.logit - a.logit)
|
|
733
915
|
|
|
734
|
-
|
|
916
|
+
let topCandidates = sortedNonTimestampLogitsWithIndexes.slice(0, options.topCandidateCount!)
|
|
917
|
+
.map(entry => ({
|
|
918
|
+
token: entry.token,
|
|
919
|
+
logit: entry.logit,
|
|
920
|
+
text: this.tokenToText(entry.token, true)
|
|
921
|
+
}))
|
|
735
922
|
|
|
736
|
-
|
|
737
|
-
|
|
923
|
+
// Apply repetition suppression if enabled
|
|
924
|
+
if (options.suppressRepetition) {
|
|
925
|
+
// Using some hardcoded constants, for now
|
|
926
|
+
const tokenWindowSize = 30
|
|
927
|
+
const thresholdMatchLength = 6
|
|
928
|
+
const thresholdCycleRepetition = 2.0
|
|
738
929
|
|
|
739
|
-
|
|
930
|
+
const filteredCandidates: typeof topCandidates = []
|
|
740
931
|
|
|
741
|
-
|
|
742
|
-
|
|
743
|
-
|
|
744
|
-
|
|
932
|
+
for (const candidate of topCandidates) {
|
|
933
|
+
const lastDecodedTextTokens = decodedTokens
|
|
934
|
+
.filter(token => this.isTextToken(token))
|
|
935
|
+
.reverse()
|
|
936
|
+
.slice(0, tokenWindowSize)
|
|
745
937
|
|
|
746
|
-
|
|
938
|
+
const { longestMatch, longestCycleRepetition } = getTokenRepetitionScore([candidate.token, ...lastDecodedTextTokens])
|
|
747
939
|
|
|
748
|
-
|
|
940
|
+
if (longestMatch >= thresholdMatchLength || longestCycleRepetition >= thresholdCycleRepetition) {
|
|
941
|
+
continue
|
|
749
942
|
}
|
|
750
|
-
}
|
|
751
|
-
|
|
752
|
-
const confidence = topCandidateProbabilities[chosenCandidateRank]
|
|
753
943
|
|
|
754
|
-
|
|
944
|
+
filteredCandidates.push(candidate)
|
|
945
|
+
}
|
|
755
946
|
|
|
756
|
-
|
|
757
|
-
|
|
947
|
+
// If all candidates have been filtered out, accept an end-of-text token
|
|
948
|
+
if (filteredCandidates.length === 0) {
|
|
949
|
+
filteredCandidates.push({
|
|
950
|
+
token: endOfTextToken,
|
|
951
|
+
logit: Infinity,
|
|
952
|
+
text: this.tokenToText(endOfTextToken, true)
|
|
953
|
+
})
|
|
758
954
|
}
|
|
955
|
+
|
|
956
|
+
topCandidates = filteredCandidates
|
|
759
957
|
}
|
|
760
958
|
|
|
761
|
-
|
|
762
|
-
|
|
959
|
+
// Compute top candidate probabilities
|
|
960
|
+
const topCandidateProbabilities = softmax(topCandidates.map(a => a.logit), options.temperature)
|
|
763
961
|
|
|
764
|
-
|
|
765
|
-
|
|
766
|
-
|
|
767
|
-
decodedTokensCrossAttentionQKs = decodedTokensCrossAttentionQKs.slice(0, lastTimestampTokenIndex)
|
|
768
|
-
}
|
|
962
|
+
// Find highest ranking punctuation token
|
|
963
|
+
const rankOfPromisingPunctuationToken = topCandidates.findIndex((entry, index) => {
|
|
964
|
+
const tokenText = this.tokenToText(entry.token).trim()
|
|
769
965
|
|
|
770
|
-
|
|
771
|
-
logger.end()
|
|
966
|
+
const isPunctuationToken = allowedPunctuationMarks.includes(tokenText)
|
|
772
967
|
|
|
773
|
-
|
|
774
|
-
|
|
775
|
-
|
|
776
|
-
decodedTokensTimestampLogits,
|
|
777
|
-
crossAttentionQKs: decodedTokensCrossAttentionQKs,
|
|
778
|
-
decodedTokensConfidence
|
|
779
|
-
}
|
|
780
|
-
}
|
|
968
|
+
if (!isPunctuationToken) {
|
|
969
|
+
return false
|
|
970
|
+
}
|
|
781
971
|
|
|
782
|
-
|
|
783
|
-
const offset = 0
|
|
972
|
+
const tokenProb = topCandidateProbabilities[index]
|
|
784
973
|
|
|
785
|
-
|
|
974
|
+
return tokenProb >= options.punctuationThreshold!
|
|
975
|
+
})
|
|
786
976
|
|
|
787
|
-
|
|
788
|
-
|
|
977
|
+
// Find rank of space token
|
|
978
|
+
let rankOfSpaceToken = topCandidates.findIndex(candidate => candidate.token === spaceToken)
|
|
789
979
|
|
|
790
|
-
|
|
791
|
-
|
|
980
|
+
if (rankOfSpaceToken < 0) {
|
|
981
|
+
rankOfSpaceToken = Infinity
|
|
982
|
+
}
|
|
792
983
|
|
|
793
|
-
|
|
984
|
+
// Choose token
|
|
985
|
+
let chosenCandidateRank: number
|
|
794
986
|
|
|
795
|
-
|
|
987
|
+
// Select a high-ranking punctuation token if found, and it has
|
|
988
|
+
// a rank higher than the space token,
|
|
989
|
+
if (rankOfPromisingPunctuationToken >= 0 &&
|
|
990
|
+
rankOfPromisingPunctuationToken < rankOfSpaceToken) {
|
|
796
991
|
|
|
797
|
-
|
|
992
|
+
chosenCandidateRank = rankOfPromisingPunctuationToken
|
|
993
|
+
} else {
|
|
994
|
+
// Otherwise, select randomly from top k distribution
|
|
995
|
+
chosenCandidateRank = this.randomGen.selectRandomIndexFromDistribution(topCandidateProbabilities)
|
|
996
|
+
}
|
|
798
997
|
|
|
799
|
-
|
|
998
|
+
// Add chosen token
|
|
999
|
+
const chosenToken = topCandidates[chosenCandidateRank].token
|
|
1000
|
+
const chosenTokenConfidence = topCandidateProbabilities[chosenCandidateRank]
|
|
800
1001
|
|
|
801
|
-
|
|
1002
|
+
addToken(chosenToken, timestampTokenLogits, chosenTokenConfidence, crossAttentionQKsForToken)
|
|
802
1003
|
|
|
803
|
-
|
|
804
|
-
|
|
1004
|
+
// If chosen token is the end-of-text token, break
|
|
1005
|
+
if (chosenToken === endOfTextToken) {
|
|
1006
|
+
break
|
|
1007
|
+
}
|
|
805
1008
|
|
|
806
|
-
|
|
1009
|
+
// Print token if needed
|
|
1010
|
+
if (this.isTextToken(chosenToken)) {
|
|
1011
|
+
bufferedTokensToPrint.push(chosenToken)
|
|
807
1012
|
|
|
808
|
-
|
|
809
|
-
const dataForToken: number[] = []
|
|
1013
|
+
let textToPrint = this.tokensToText(bufferedTokensToPrint)
|
|
810
1014
|
|
|
811
|
-
|
|
812
|
-
|
|
813
|
-
|
|
814
|
-
|
|
815
|
-
dataForToken.push(qkArray.get(i0, i1, i2, i3, i4) as number)
|
|
816
|
-
}
|
|
1015
|
+
// If the decoded text is valid, print it
|
|
1016
|
+
if (!containsInvalidCodepoint(textToPrint)) {
|
|
1017
|
+
if (isFirstPart && decodedTokens.every(token => this.isMetadataToken(token))) {
|
|
1018
|
+
textToPrint = textToPrint.trimStart()
|
|
817
1019
|
}
|
|
1020
|
+
|
|
1021
|
+
logger.write(textToPrint)
|
|
1022
|
+
|
|
1023
|
+
bufferedTokensToPrint = []
|
|
818
1024
|
}
|
|
819
1025
|
}
|
|
820
1026
|
|
|
821
|
-
|
|
822
|
-
|
|
1027
|
+
await yieldToEventLoop()
|
|
1028
|
+
}
|
|
823
1029
|
|
|
824
|
-
|
|
1030
|
+
// If at least two timestamp tokens were decoded and it's not the final part,
|
|
1031
|
+
// truncate up to the last timestamp token
|
|
1032
|
+
if (timestampTokenSeenCount >= 2 && !isFinalPart) {
|
|
1033
|
+
const sliceEndTokenIndex = lastTimestampTokenIndex
|
|
825
1034
|
|
|
826
|
-
|
|
1035
|
+
decodedTokens = decodedTokens.slice(0, sliceEndTokenIndex)
|
|
1036
|
+
decodedTokensTimestampLogits = decodedTokensTimestampLogits.slice(0, sliceEndTokenIndex)
|
|
1037
|
+
decodedTokensCrossAttentionQKs = decodedTokensCrossAttentionQKs.slice(0, sliceEndTokenIndex)
|
|
1038
|
+
decodedTokensConfidence = decodedTokensConfidence.slice(0, sliceEndTokenIndex)
|
|
827
1039
|
}
|
|
828
1040
|
|
|
829
|
-
|
|
1041
|
+
logger.write('\n')
|
|
1042
|
+
logger.end()
|
|
1043
|
+
|
|
1044
|
+
// Return the decoded tokens
|
|
1045
|
+
return {
|
|
1046
|
+
decodedTokens,
|
|
1047
|
+
decodedTokensTimestampLogits,
|
|
1048
|
+
decodedTokensConfidence,
|
|
1049
|
+
decodedTokensCrossAttentionQKs,
|
|
1050
|
+
}
|
|
830
1051
|
}
|
|
831
1052
|
|
|
1053
|
+
// Encode audio using the encoder model
|
|
832
1054
|
async encodeAudio(rawAudio: RawAudio) {
|
|
833
1055
|
await this.initializeEncoderSessionIfNeeded()
|
|
834
1056
|
|
|
@@ -846,10 +1068,14 @@ export class Whisper {
|
|
|
846
1068
|
const maxAudioSamples = sampleRate * 30
|
|
847
1069
|
const maxAudioFrames = 3000
|
|
848
1070
|
|
|
1071
|
+
if (audioSamples.length > maxAudioSamples) {
|
|
1072
|
+
throw new Error(`Audio part is longer than 30 seconds`)
|
|
1073
|
+
}
|
|
1074
|
+
|
|
849
1075
|
await logger.startAsync('Extract mel spectogram from audio part')
|
|
850
1076
|
|
|
851
1077
|
const paddedAudioSamples = new Float32Array(maxAudioSamples)
|
|
852
|
-
paddedAudioSamples.set(audioSamples
|
|
1078
|
+
paddedAudioSamples.set(audioSamples, 0)
|
|
853
1079
|
|
|
854
1080
|
const rawAudioPart: RawAudio = { audioChannels: [paddedAudioSamples], sampleRate }
|
|
855
1081
|
|
|
@@ -893,128 +1119,6 @@ export class Whisper {
|
|
|
893
1119
|
return encodedAudioFeatures
|
|
894
1120
|
}
|
|
895
1121
|
|
|
896
|
-
addSegmentsToTimeline(timeline: Timeline, tokens: number[], initialTimeOffset: number, audioDuration: number) {
|
|
897
|
-
const timestampTokensStart = this.tokenConfig.timestampTokensStart
|
|
898
|
-
|
|
899
|
-
for (let i = 0; i < tokens.length; i++) {
|
|
900
|
-
const token = tokens[i]
|
|
901
|
-
|
|
902
|
-
if (token == this.tokenConfig.startOfTextToken || token == this.tokenConfig.endOfTextToken) {
|
|
903
|
-
continue
|
|
904
|
-
}
|
|
905
|
-
|
|
906
|
-
const tokenIsTimestamp = token >= timestampTokensStart
|
|
907
|
-
const previousTokenWasTimestamp = tokens.length > 1 && tokens[i - 1] >= timestampTokensStart
|
|
908
|
-
|
|
909
|
-
if (tokenIsTimestamp) {
|
|
910
|
-
if (previousTokenWasTimestamp) {
|
|
911
|
-
continue
|
|
912
|
-
}
|
|
913
|
-
|
|
914
|
-
let startTime = initialTimeOffset + this.timestampTokenToSeconds(token)
|
|
915
|
-
|
|
916
|
-
startTime = Math.min(startTime, audioDuration)
|
|
917
|
-
|
|
918
|
-
if (timeline.length > 0) {
|
|
919
|
-
timeline[timeline.length - 1].endTime = startTime
|
|
920
|
-
}
|
|
921
|
-
|
|
922
|
-
timeline.push({
|
|
923
|
-
type: 'segment',
|
|
924
|
-
text: '',
|
|
925
|
-
startTime,
|
|
926
|
-
endTime: -1,
|
|
927
|
-
})
|
|
928
|
-
} else {
|
|
929
|
-
if (timeline.length == 0) {
|
|
930
|
-
timeline.push({
|
|
931
|
-
type: 'segment',
|
|
932
|
-
text: '',
|
|
933
|
-
startTime: initialTimeOffset,
|
|
934
|
-
endTime: -1,
|
|
935
|
-
})
|
|
936
|
-
}
|
|
937
|
-
|
|
938
|
-
const tokenText = this.tokenToText(token)
|
|
939
|
-
|
|
940
|
-
timeline[timeline.length - 1].text += tokenText
|
|
941
|
-
}
|
|
942
|
-
}
|
|
943
|
-
}
|
|
944
|
-
|
|
945
|
-
async addWordsToTimeline(timeline: Timeline, tokens: number[], rawAudio: RawAudio, crossAttentionQKs: Onnx.Tensor[], initialAudioTimeOffset: number, duration: number) {
|
|
946
|
-
let segmentStartTime = 0
|
|
947
|
-
let segmentTokens: number[] = []
|
|
948
|
-
let segmentCrossAttentionQKs: Onnx.Tensor[] = []
|
|
949
|
-
|
|
950
|
-
for (let tokenIndex = 0; tokenIndex < tokens.length; tokenIndex++) {
|
|
951
|
-
const token = tokens[tokenIndex]
|
|
952
|
-
const tokenCrossAttentionQKs = crossAttentionQKs[tokenIndex]
|
|
953
|
-
|
|
954
|
-
const segmentTokensWithoutTimestamps = segmentTokens.filter(token => this.isNonTimestampToken(token))
|
|
955
|
-
|
|
956
|
-
const isTimestamp = this.isTimestampToken(token)
|
|
957
|
-
|
|
958
|
-
if (isTimestamp || tokenIndex == tokens.length - 1) {
|
|
959
|
-
let tokenTime: number
|
|
960
|
-
|
|
961
|
-
if (isTimestamp) {
|
|
962
|
-
tokenTime = this.timestampTokenToSeconds(token)
|
|
963
|
-
} else {
|
|
964
|
-
tokenTime = duration
|
|
965
|
-
}
|
|
966
|
-
|
|
967
|
-
if (segmentTokensWithoutTimestamps.length > 0) {
|
|
968
|
-
const segmentEndTime = tokenTime
|
|
969
|
-
|
|
970
|
-
const segmentStartFrame = this.secondsToFrame(segmentStartTime)
|
|
971
|
-
let segmentEndFrame = this.secondsToFrame(segmentEndTime)
|
|
972
|
-
|
|
973
|
-
if (segmentStartFrame == segmentEndFrame) {
|
|
974
|
-
segmentEndFrame += 1
|
|
975
|
-
}
|
|
976
|
-
|
|
977
|
-
const segmentFrameCount = segmentEndFrame - segmentStartFrame
|
|
978
|
-
|
|
979
|
-
const reinferCrossAttentionQKs = true
|
|
980
|
-
|
|
981
|
-
if (reinferCrossAttentionQKs) {
|
|
982
|
-
const initialTokens = this.getTextStartTokens('en', 'transcribe')
|
|
983
|
-
const tokensToDecode = [...initialTokens, ...segmentTokensWithoutTimestamps]
|
|
984
|
-
|
|
985
|
-
//const segmentAudioFeaturesBuffer = audioFeatures.data.slice(segmentStartFrame * audioFeatures.dims[2], segmentEndFrame * audioFeatures.dims[2])
|
|
986
|
-
//const segmentAudioFeatures = new Onnx.Tensor('float32', segmentAudioFeaturesBuffer, [1, segmentFrameCount, audioFeatures.dims[2]])
|
|
987
|
-
|
|
988
|
-
const segmentAudioSamples = rawAudio.audioChannels[0].slice(Math.floor(segmentStartTime * rawAudio.sampleRate), Math.floor(segmentEndTime * rawAudio.sampleRate))
|
|
989
|
-
const segmentRawAudio: RawAudio = { audioChannels: [segmentAudioSamples], sampleRate: rawAudio.sampleRate }
|
|
990
|
-
|
|
991
|
-
const segmentAudioFeatures = await this.encodeAudio(segmentRawAudio)
|
|
992
|
-
|
|
993
|
-
const reinferredCrossAttentionQKs = await this.inferCrossAttentionQKs(tokensToDecode, segmentAudioFeatures)
|
|
994
|
-
reinferredCrossAttentionQKs.slice(initialTokens.length)
|
|
995
|
-
|
|
996
|
-
const alignmentPath = await this.findAlignmentPathFromQKs(reinferredCrossAttentionQKs, tokensToDecode, 0, segmentFrameCount)//, alignmentHeadsIndexes[modelName])
|
|
997
|
-
const tokenTimeline = await this.getTokenTimelineFromAlignmentPath(alignmentPath, segmentTokensWithoutTimestamps, initialAudioTimeOffset + segmentStartTime, initialAudioTimeOffset + segmentEndTime)
|
|
998
|
-
|
|
999
|
-
timeline.push(...tokenTimeline)
|
|
1000
|
-
} else {
|
|
1001
|
-
const alignmentPath = await this.findAlignmentPathFromQKs(segmentCrossAttentionQKs, segmentTokens, segmentStartFrame, segmentEndFrame)//, alignmentHeadsIndexes[modelName])
|
|
1002
|
-
const tokenTimeline = await this.getTokenTimelineFromAlignmentPath(alignmentPath, segmentTokens, initialAudioTimeOffset, initialAudioTimeOffset + segmentEndTime)
|
|
1003
|
-
|
|
1004
|
-
timeline.push(...tokenTimeline)
|
|
1005
|
-
}
|
|
1006
|
-
}
|
|
1007
|
-
|
|
1008
|
-
segmentStartTime = tokenTime
|
|
1009
|
-
segmentTokens = []
|
|
1010
|
-
segmentCrossAttentionQKs = []
|
|
1011
|
-
}
|
|
1012
|
-
|
|
1013
|
-
segmentTokens.push(token)
|
|
1014
|
-
segmentCrossAttentionQKs.push(tokenCrossAttentionQKs)
|
|
1015
|
-
}
|
|
1016
|
-
}
|
|
1017
|
-
|
|
1018
1122
|
tokenTimelineToWordTimeline(tokenTimeline: Timeline, language: string): Timeline {
|
|
1019
1123
|
function isSeparatorCharacter(char: string) {
|
|
1020
1124
|
const nonSeparatingPunctuation = [`'`, `-`, `.`, `·`, `•`]
|
|
@@ -1034,9 +1138,13 @@ export class Whisper {
|
|
|
1034
1138
|
return isSeparatorCharacter(text[text.length - 1])
|
|
1035
1139
|
}
|
|
1036
1140
|
|
|
1141
|
+
if (language != 'zh' && language != 'ja') {
|
|
1142
|
+
tokenTimeline = tokenTimeline.filter(entry => this.isTextToken(entry.id!))
|
|
1143
|
+
}
|
|
1144
|
+
|
|
1037
1145
|
const resultTimeline: Timeline = []
|
|
1038
1146
|
|
|
1039
|
-
let groups:
|
|
1147
|
+
let groups: Timeline[] = []
|
|
1040
1148
|
|
|
1041
1149
|
for (let tokenIndex = 0; tokenIndex < tokenTimeline.length; tokenIndex++) {
|
|
1042
1150
|
const entry = tokenTimeline[tokenIndex]
|
|
@@ -1056,25 +1164,27 @@ export class Whisper {
|
|
|
1056
1164
|
}
|
|
1057
1165
|
}
|
|
1058
1166
|
|
|
1059
|
-
|
|
1167
|
+
{
|
|
1168
|
+
const splitGroups: Timeline[] = []
|
|
1060
1169
|
|
|
1061
|
-
|
|
1062
|
-
|
|
1063
|
-
|
|
1170
|
+
for (let groupIndex = 0; groupIndex < groups.length; groupIndex++) {
|
|
1171
|
+
const group = groups[groupIndex]
|
|
1172
|
+
const nextGroup = groups[groupIndex + 1]
|
|
1064
1173
|
|
|
1065
|
-
|
|
1066
|
-
|
|
1067
|
-
|
|
1068
|
-
|
|
1174
|
+
if (
|
|
1175
|
+
group.length > 1 &&
|
|
1176
|
+
group[group.length - 1].text === '.' &&
|
|
1177
|
+
(!nextGroup || [' ', '['].includes(nextGroup[0].text[0]))) {
|
|
1069
1178
|
|
|
1070
|
-
|
|
1071
|
-
|
|
1072
|
-
|
|
1073
|
-
|
|
1179
|
+
splitGroups.push(group.slice(0, group.length - 1))
|
|
1180
|
+
splitGroups.push(group.slice(group.length - 1))
|
|
1181
|
+
} else {
|
|
1182
|
+
splitGroups.push(group)
|
|
1183
|
+
}
|
|
1074
1184
|
}
|
|
1075
|
-
}
|
|
1076
1185
|
|
|
1077
|
-
|
|
1186
|
+
groups = splitGroups
|
|
1187
|
+
}
|
|
1078
1188
|
|
|
1079
1189
|
for (const group of groups) {
|
|
1080
1190
|
let groupText = this.tokensToText(group.map(entry => entry.id!))
|
|
@@ -1150,7 +1260,7 @@ export class Whisper {
|
|
|
1150
1260
|
return tokenTimeline
|
|
1151
1261
|
}
|
|
1152
1262
|
|
|
1153
|
-
async findAlignmentPathFromQKs(qksTensors:
|
|
1263
|
+
async findAlignmentPathFromQKs(qksTensors: OnnxLikeFloat32Tensor[], tokens: number[], segmentStartFrame: number, segmentEndFrame: number, headIndexes?: number[]) {
|
|
1154
1264
|
const segmentFrameCount = segmentEndFrame - segmentStartFrame
|
|
1155
1265
|
|
|
1156
1266
|
if (segmentFrameCount === 0 || tokens.length === 0 || qksTensors.length === 0) {
|
|
@@ -1193,13 +1303,12 @@ export class Whisper {
|
|
|
1193
1303
|
const applySoftmax = true
|
|
1194
1304
|
const normalize = true
|
|
1195
1305
|
const applyMedianFilter = true
|
|
1196
|
-
const
|
|
1306
|
+
const anchorTimestampTokens = false
|
|
1197
1307
|
|
|
1198
1308
|
const softmaxTemperature = 1.0
|
|
1199
|
-
const medianFilterWidth = 7
|
|
1200
1309
|
|
|
1310
|
+
// Apply softmax to each token's frames, if enabled
|
|
1201
1311
|
if (applySoftmax) {
|
|
1202
|
-
// Apply softmax to each token's frames
|
|
1203
1312
|
for (const head of attentionHeads) {
|
|
1204
1313
|
for (let tokenIndex = 0; tokenIndex < tokenCount; tokenIndex++) {
|
|
1205
1314
|
head[tokenIndex] = softmax(head[tokenIndex], softmaxTemperature)
|
|
@@ -1207,27 +1316,29 @@ export class Whisper {
|
|
|
1207
1316
|
}
|
|
1208
1317
|
}
|
|
1209
1318
|
|
|
1319
|
+
// Normalize all weights in each individual head, if enabled
|
|
1210
1320
|
if (normalize) {
|
|
1211
|
-
// Normalize all weights in each individual head
|
|
1212
1321
|
for (const head of attentionHeads) {
|
|
1213
1322
|
const allWeightsForHead = head.flatMap(tokenFrames => tokenFrames)
|
|
1214
1323
|
|
|
1215
|
-
const
|
|
1216
|
-
const
|
|
1324
|
+
const meanOfAllWeightsForHead = meanOfVector(allWeightsForHead)
|
|
1325
|
+
const stdDeviationOfAllWeightsForHead = stdDeviationOfVector(allWeightsForHead, 'population', meanOfAllWeightsForHead) + 1e-10
|
|
1326
|
+
|
|
1327
|
+
const stdDeviationReciprocal = 1.0 / (stdDeviationOfAllWeightsForHead + 1e-10)
|
|
1217
1328
|
|
|
1218
1329
|
for (const tokenFrames of head) {
|
|
1219
1330
|
for (let frameIndex = 0; frameIndex < tokenFrames.length; frameIndex++) {
|
|
1220
|
-
tokenFrames[frameIndex] = (tokenFrames[frameIndex] -
|
|
1331
|
+
tokenFrames[frameIndex] = (tokenFrames[frameIndex] - meanOfAllWeightsForHead) * stdDeviationReciprocal
|
|
1221
1332
|
}
|
|
1222
1333
|
}
|
|
1223
1334
|
}
|
|
1224
1335
|
}
|
|
1225
1336
|
|
|
1337
|
+
// Apply median filter to each token's frames, if enabled
|
|
1226
1338
|
if (applyMedianFilter) {
|
|
1227
|
-
// Apply median filter to each token's frames
|
|
1228
1339
|
for (const head of attentionHeads) {
|
|
1229
1340
|
for (let tokenIndex = 0; tokenIndex < tokenCount; tokenIndex++) {
|
|
1230
|
-
head[tokenIndex] =
|
|
1341
|
+
head[tokenIndex] = medianOf5Filter(head[tokenIndex])
|
|
1231
1342
|
}
|
|
1232
1343
|
}
|
|
1233
1344
|
}
|
|
@@ -1255,8 +1366,8 @@ export class Whisper {
|
|
|
1255
1366
|
}
|
|
1256
1367
|
}
|
|
1257
1368
|
|
|
1258
|
-
if
|
|
1259
|
-
|
|
1369
|
+
// Anchor timestamp tokens timestamps to their original values, if enabled
|
|
1370
|
+
if (anchorTimestampTokens) {
|
|
1260
1371
|
const timestampTokensStart = this.tokenConfig.timestampTokensStart
|
|
1261
1372
|
|
|
1262
1373
|
for (let tokenIndex = 0; tokenIndex < tokens.length; tokenIndex++) {
|
|
@@ -1273,8 +1384,8 @@ export class Whisper {
|
|
|
1273
1384
|
}
|
|
1274
1385
|
|
|
1275
1386
|
// Perform DTW
|
|
1276
|
-
const tokenIndexes =
|
|
1277
|
-
const frameIndexes =
|
|
1387
|
+
const tokenIndexes = getIntegerRange(0, tokenCount)
|
|
1388
|
+
const frameIndexes = getIntegerRange(0, segmentFrameCount)
|
|
1278
1389
|
|
|
1279
1390
|
let { path } = alignDTWWindowed(tokenIndexes, frameIndexes, (tokenIndex, frameIndex) => {
|
|
1280
1391
|
return -frameMeansForToken[tokenIndex][frameIndex]
|
|
@@ -1285,6 +1396,114 @@ export class Whisper {
|
|
|
1285
1396
|
return path
|
|
1286
1397
|
}
|
|
1287
1398
|
|
|
1399
|
+
async initializeIfNeeded() {
|
|
1400
|
+
await this.initializeTokenizerIfNeeded()
|
|
1401
|
+
await this.initializeEncoderSessionIfNeeded()
|
|
1402
|
+
await this.initializeDecoderSessionIfNeeded()
|
|
1403
|
+
}
|
|
1404
|
+
|
|
1405
|
+
async initializeTokenizerIfNeeded() {
|
|
1406
|
+
if (this.tiktoken) {
|
|
1407
|
+
return
|
|
1408
|
+
}
|
|
1409
|
+
|
|
1410
|
+
const logger = new Logger()
|
|
1411
|
+
await logger.startAsync('Load tokenizer data')
|
|
1412
|
+
|
|
1413
|
+
const tiktokenModulePackagePath = await loadPackage('whisper-tiktoken-data')
|
|
1414
|
+
|
|
1415
|
+
const tiktokenDataFilePath = path.join(tiktokenModulePackagePath, this.isMultiligualModel ? 'multilingual.tiktoken' : 'gpt2.tiktoken')
|
|
1416
|
+
let tiktokenData = await readFile(tiktokenDataFilePath, { encoding: 'utf8' })
|
|
1417
|
+
|
|
1418
|
+
const tokenConfig = this.tokenConfig
|
|
1419
|
+
|
|
1420
|
+
const metadataTokens: Record<number, string> = {
|
|
1421
|
+
[tokenConfig.endOfTextToken]: '[EndOfText]',
|
|
1422
|
+
[tokenConfig.startOfTextToken]: '[StartOfText]',
|
|
1423
|
+
[tokenConfig.translateTaskToken]: '[TranslateTask]',
|
|
1424
|
+
[tokenConfig.transcribeTaskToken]: '[TranscribeTask]',
|
|
1425
|
+
[tokenConfig.startOfPromptToken]: '[StartOfPrompt]',
|
|
1426
|
+
[tokenConfig.nonSpeechToken]: '[NonSpeech]',
|
|
1427
|
+
[tokenConfig.noTimestampsToken]: '[NoTimestamps]',
|
|
1428
|
+
}
|
|
1429
|
+
|
|
1430
|
+
if (this.isMultiligualModel) {
|
|
1431
|
+
metadataTokens[50256] = '[Unused_50256]'
|
|
1432
|
+
metadataTokens[50360] = '[Unused_50360]'
|
|
1433
|
+
}
|
|
1434
|
+
|
|
1435
|
+
const languageTokenCount = tokenConfig.languageTokensEnd - tokenConfig.languageTokensStart
|
|
1436
|
+
|
|
1437
|
+
for (let i = 0; i < languageTokenCount; i++) {
|
|
1438
|
+
const tokenIndex = this.tokenConfig.languageTokensStart + i
|
|
1439
|
+
|
|
1440
|
+
metadataTokens[tokenIndex] = `[Language_${i}]`
|
|
1441
|
+
}
|
|
1442
|
+
|
|
1443
|
+
const timestampTokensCount = 1501
|
|
1444
|
+
|
|
1445
|
+
for (let i = 0; i < timestampTokensCount; i++) {
|
|
1446
|
+
const tokenIndex = this.tokenConfig.timestampTokensStart + i
|
|
1447
|
+
const tokenTime = this.timestampTokenToSeconds(tokenIndex)
|
|
1448
|
+
|
|
1449
|
+
metadataTokens[tokenIndex] = `[Timestamp_${tokenTime.toFixed(2)}]`
|
|
1450
|
+
}
|
|
1451
|
+
|
|
1452
|
+
const inverseMetadataTokensLookup: Record<string, number> = {}
|
|
1453
|
+
|
|
1454
|
+
for (const [key, value] of Object.entries(metadataTokens)) {
|
|
1455
|
+
inverseMetadataTokensLookup[value] = parseInt(key)
|
|
1456
|
+
}
|
|
1457
|
+
|
|
1458
|
+
const patternString = `'s|'t|'re|'ve|'m|'ll|'d| ?\\p{L}+| ?\\p{N}+| ?[^\\s\\p{L}\\p{N}]+|\\s+(?!\\S)|\\s+`
|
|
1459
|
+
|
|
1460
|
+
const { Tiktoken } = await import('tiktoken/lite')
|
|
1461
|
+
|
|
1462
|
+
this.tiktoken = new Tiktoken(tiktokenData, inverseMetadataTokensLookup, patternString)
|
|
1463
|
+
|
|
1464
|
+
logger.end()
|
|
1465
|
+
}
|
|
1466
|
+
|
|
1467
|
+
async initializeEncoderSessionIfNeeded() {
|
|
1468
|
+
if (this.audioEncoder) {
|
|
1469
|
+
return
|
|
1470
|
+
}
|
|
1471
|
+
|
|
1472
|
+
const logger = new Logger()
|
|
1473
|
+
|
|
1474
|
+
await logger.startAsync(`Create encoder inference session for model '${this.modelName}'`)
|
|
1475
|
+
|
|
1476
|
+
const encoderFilePath = path.join(this.modelDir, 'encoder.onnx')
|
|
1477
|
+
|
|
1478
|
+
const Onnx = await import('onnxruntime-node')
|
|
1479
|
+
|
|
1480
|
+
const onnxSessionOptions = getOnnxSessionOptions({ executionProviders: this.encoderExecutionProviders })
|
|
1481
|
+
|
|
1482
|
+
this.audioEncoder = await Onnx.InferenceSession.create(encoderFilePath, onnxSessionOptions)
|
|
1483
|
+
|
|
1484
|
+
logger.end()
|
|
1485
|
+
}
|
|
1486
|
+
|
|
1487
|
+
async initializeDecoderSessionIfNeeded() {
|
|
1488
|
+
if (this.textDecoder) {
|
|
1489
|
+
return
|
|
1490
|
+
}
|
|
1491
|
+
|
|
1492
|
+
const logger = new Logger()
|
|
1493
|
+
|
|
1494
|
+
await logger.startAsync(`Create decoder inference session for model '${this.modelName}'`)
|
|
1495
|
+
|
|
1496
|
+
const decoderFilePath = path.join(this.modelDir, 'decoder.onnx')
|
|
1497
|
+
|
|
1498
|
+
const Onnx = await import('onnxruntime-node')
|
|
1499
|
+
|
|
1500
|
+
const onnxSessionOptions = getOnnxSessionOptions({ executionProviders: this.decoderExecutionProviders })
|
|
1501
|
+
|
|
1502
|
+
this.textDecoder = await Onnx.InferenceSession.create(decoderFilePath, onnxSessionOptions)
|
|
1503
|
+
|
|
1504
|
+
logger.end()
|
|
1505
|
+
}
|
|
1506
|
+
|
|
1288
1507
|
getKvDimensions(groupCount: number, length: number) {
|
|
1289
1508
|
const modelName = this.modelName
|
|
1290
1509
|
|
|
@@ -1432,7 +1651,7 @@ export class Whisper {
|
|
|
1432
1651
|
getSuppressedTextTokens() {
|
|
1433
1652
|
const allowedPunctuationMarks = this.getAllowedPunctuationMarks()
|
|
1434
1653
|
|
|
1435
|
-
const nonWordTokensData = this.
|
|
1654
|
+
const nonWordTokensData = this.getWordTokenData().nonWordTokenData
|
|
1436
1655
|
|
|
1437
1656
|
const suppressedTextTokens = nonWordTokensData
|
|
1438
1657
|
.filter(entry => !allowedPunctuationMarks.includes(entry.text))
|
|
@@ -1468,31 +1687,33 @@ export class Whisper {
|
|
|
1468
1687
|
return allowedPunctuation
|
|
1469
1688
|
}
|
|
1470
1689
|
|
|
1471
|
-
|
|
1690
|
+
getWordTokenData() {
|
|
1691
|
+
const wordTokenData: WhisperTokenData[] = []
|
|
1472
1692
|
const nonWordTokenData: WhisperTokenData[] = []
|
|
1473
1693
|
|
|
1474
|
-
const invalidUTF8Char = String.fromCharCode(65533)
|
|
1475
|
-
|
|
1476
1694
|
for (let i = 0; i < this.tokenConfig.endOfTextToken; i++) {
|
|
1477
1695
|
const tokenText = this.tokenToText(i, false)
|
|
1478
|
-
const tokenTextWithoutWhitespace = tokenText.replaceAll(/\s/g, '')
|
|
1479
1696
|
|
|
1480
|
-
const isNonWordToken = /^[\p{Punctuation}\p{Symbol}]+$/u.test(
|
|
1697
|
+
const isNonWordToken = /^[\s\p{Punctuation}\p{Symbol}]+$/u.test(tokenText)
|
|
1481
1698
|
|
|
1482
|
-
const containsInvalidUTF8 =
|
|
1699
|
+
const containsInvalidUTF8 = containsInvalidCodepoint(tokenText)
|
|
1483
1700
|
|
|
1484
|
-
if (isNonWordToken && !containsInvalidUTF8) {
|
|
1701
|
+
if (isNonWordToken && (this.isEnglishOnlyModel || !containsInvalidUTF8)) {
|
|
1485
1702
|
nonWordTokenData.push({
|
|
1486
1703
|
id: i,
|
|
1487
1704
|
text: tokenText,
|
|
1488
1705
|
})
|
|
1706
|
+
} else {
|
|
1707
|
+
wordTokenData.push({
|
|
1708
|
+
id: i,
|
|
1709
|
+
text: tokenText,
|
|
1710
|
+
})
|
|
1489
1711
|
}
|
|
1490
1712
|
}
|
|
1491
1713
|
|
|
1492
|
-
return nonWordTokenData
|
|
1714
|
+
return { wordTokenData, nonWordTokenData }
|
|
1493
1715
|
}
|
|
1494
1716
|
|
|
1495
|
-
|
|
1496
1717
|
getTokensData(tokens: number[]) {
|
|
1497
1718
|
const tokensData: WhisperTokenData[] = []
|
|
1498
1719
|
|
|
@@ -1507,168 +1728,6 @@ export class Whisper {
|
|
|
1507
1728
|
}
|
|
1508
1729
|
}
|
|
1509
1730
|
|
|
1510
|
-
const filterbanks: Filterbank[] = [
|
|
1511
|
-
/* 0 */ { startIndex: 1, weights: [0.02486259490251541,] },
|
|
1512
|
-
|
|
1513
|
-
/* 1 */ { startIndex: 1, weights: [0.001990821911022067, 0.022871771827340126,] },
|
|
1514
|
-
|
|
1515
|
-
/* 2 */ { startIndex: 2, weights: [0.003981643822044134, 0.02088095061480999,] },
|
|
1516
|
-
|
|
1517
|
-
/* 3 */ { startIndex: 3, weights: [0.0059724655002355576, 0.018890129402279854,] },
|
|
1518
|
-
|
|
1519
|
-
/* 4 */ { startIndex: 4, weights: [0.007963287644088268, 0.01689930632710457,] },
|
|
1520
|
-
|
|
1521
|
-
/* 5 */ { startIndex: 5, weights: [0.009954108856618404, 0.014908484183251858,] },
|
|
1522
|
-
|
|
1523
|
-
/* 6 */ { startIndex: 6, weights: [0.011944931000471115, 0.012917662039399147,] },
|
|
1524
|
-
|
|
1525
|
-
/* 7 */ { startIndex: 7, weights: [0.013935752213001251, 0.010926840826869011,] },
|
|
1526
|
-
|
|
1527
|
-
/* 8 */ { startIndex: 8, weights: [0.015926575288176537, 0.0089360186830163,] },
|
|
1528
|
-
|
|
1529
|
-
/* 9 */ { startIndex: 9, weights: [0.017917396500706673, 0.006945197004824877,] },
|
|
1530
|
-
|
|
1531
|
-
/* 10 */ { startIndex: 10, weights: [0.01990821771323681, 0.004954374860972166,] },
|
|
1532
|
-
|
|
1533
|
-
/* 11 */ { startIndex: 11, weights: [0.021899040788412094, 0.0029635531827807426,] },
|
|
1534
|
-
|
|
1535
|
-
/* 12 */ { startIndex: 12, weights: [0.02388986200094223, 0.0009727313299663365,] },
|
|
1536
|
-
|
|
1537
|
-
/* 13 */ { startIndex: 13, weights: [0.025880683213472366,] },
|
|
1538
|
-
|
|
1539
|
-
/* 14 */ { startIndex: 14, weights: [0.025835324078798294,] },
|
|
1540
|
-
|
|
1541
|
-
/* 15 */ { startIndex: 14, weights: [0.0010180906392633915, 0.023844502866268158,] },
|
|
1542
|
-
|
|
1543
|
-
/* 16 */ { startIndex: 15, weights: [0.003008912317454815, 0.021853681653738022,] },
|
|
1544
|
-
|
|
1545
|
-
/* 17 */ { startIndex: 16, weights: [0.004999734461307526, 0.019862858578562737,] },
|
|
1546
|
-
|
|
1547
|
-
/* 18 */ { startIndex: 17, weights: [0.006990555673837662, 0.0178720373660326,] },
|
|
1548
|
-
|
|
1549
|
-
/* 19 */ { startIndex: 18, weights: [0.008981377817690372, 0.015881216153502464,] },
|
|
1550
|
-
|
|
1551
|
-
/* 20 */ { startIndex: 19, weights: [0.010972199961543083, 0.013890394009649754,] },
|
|
1552
|
-
|
|
1553
|
-
/* 21 */ { startIndex: 20, weights: [0.01296302117407322, 0.011899571865797043,] },
|
|
1554
|
-
|
|
1555
|
-
/* 22 */ { startIndex: 21, weights: [0.01495384331792593, 0.009908749721944332,] },
|
|
1556
|
-
|
|
1557
|
-
/* 23 */ { startIndex: 22, weights: [0.01694466546177864, 0.007917927578091621,] },
|
|
1558
|
-
|
|
1559
|
-
/* 24 */ { startIndex: 23, weights: [0.018935488536953926, 0.005927106365561485,] },
|
|
1560
|
-
|
|
1561
|
-
/* 25 */ { startIndex: 24, weights: [0.020874010398983955, 0.004040425643324852,] },
|
|
1562
|
-
|
|
1563
|
-
/* 26 */ { startIndex: 25, weights: [0.022114217281341553, 0.0033186059445142746,] },
|
|
1564
|
-
|
|
1565
|
-
/* 27 */ { startIndex: 26, weights: [0.02173672430217266, 0.0036109676584601402,] },
|
|
1566
|
-
|
|
1567
|
-
/* 28 */ { startIndex: 27, weights: [0.020497702062129974, 0.004762193653732538,] },
|
|
1568
|
-
|
|
1569
|
-
/* 29 */ { startIndex: 28, weights: [0.018486659973859787, 0.006592618301510811,] },
|
|
1570
|
-
|
|
1571
|
-
/* 30 */ { startIndex: 29, weights: [0.01585603691637516, 0.00896277092397213,] },
|
|
1572
|
-
|
|
1573
|
-
/* 31 */ { startIndex: 30, weights: [0.012738768011331558, 0.011751330457627773,] },
|
|
1574
|
-
|
|
1575
|
-
/* 32 */ { startIndex: 31, weights: [0.009250369854271412, 0.014853144995868206,] },
|
|
1576
|
-
|
|
1577
|
-
/* 33 */ { startIndex: 32, weights: [0.005490840878337622, 0.018177473917603493, 0.0028155462350696325,] },
|
|
1578
|
-
|
|
1579
|
-
/* 34 */ { startIndex: 33, weights: [0.0015463664894923568, 0.01632951945066452, 0.007420188747346401,] },
|
|
1580
|
-
|
|
1581
|
-
/* 35 */ { startIndex: 35, weights: [0.011181050911545753, 0.012018864043056965,] },
|
|
1582
|
-
|
|
1583
|
-
/* 36 */ { startIndex: 36, weights: [0.006065350491553545, 0.016561277210712433, 0.004360878840088844,] },
|
|
1584
|
-
|
|
1585
|
-
/* 37 */ { startIndex: 37, weights: [0.0010297985281795263, 0.012770536355674267, 0.009707189165055752,] },
|
|
1586
|
-
|
|
1587
|
-
/* 38 */ { startIndex: 39, weights: [0.006986402906477451, 0.01485429983586073, 0.004391219466924667,] },
|
|
1588
|
-
|
|
1589
|
-
/* 39 */ { startIndex: 40, weights: [0.001418047584593296, 0.011486922390758991, 0.010089744813740253, 0.00040022286702878773,] },
|
|
1590
|
-
|
|
1591
|
-
/* 40 */ { startIndex: 42, weights: [0.005411104764789343, 0.014735566452145576, 0.006518189795315266,] },
|
|
1592
|
-
|
|
1593
|
-
/* 41 */ { startIndex: 44, weights: [0.00827841367572546, 0.012277561239898205, 0.00396781275048852,] },
|
|
1594
|
-
|
|
1595
|
-
/* 42 */ { startIndex: 45, weights: [0.002187808509916067, 0.010184479877352715, 0.00998187530785799, 0.0022864851634949446,] },
|
|
1596
|
-
|
|
1597
|
-
/* 43 */ { startIndex: 47, weights: [0.00386943481862545, 0.011274894699454308, 0.008466221392154694, 0.0013397691072896123,] },
|
|
1598
|
-
|
|
1599
|
-
/* 44 */ { startIndex: 49, weights: [0.004820294212549925, 0.011678251437842846, 0.007608682848513126, 0.0010091039584949613,] },
|
|
1600
|
-
|
|
1601
|
-
/* 45 */ { startIndex: 51, weights: [0.005156961735337973, 0.011507894843816757, 0.007301822770386934, 0.0011901655234396458,] },
|
|
1602
|
-
|
|
1603
|
-
/* 46 */ { startIndex: 53, weights: [0.004982104524970055, 0.010863498784601688, 0.007451189681887627, 0.001791381393559277,] },
|
|
1604
|
-
|
|
1605
|
-
/* 47 */ { startIndex: 55, weights: [0.004385921638458967, 0.009832492098212242, 0.007973956875503063, 0.002732589840888977,] },
|
|
1606
|
-
|
|
1607
|
-
/* 48 */ { startIndex: 57, weights: [0.0034474546555429697, 0.008491347543895245, 0.008797688409686089, 0.00394382793456316,] },
|
|
1608
|
-
|
|
1609
|
-
/* 49 */ { startIndex: 59, weights: [0.0022357646375894547, 0.0069067515432834625, 0.009859241545200348, 0.005364237818866968, 0.0008692338014952838,] },
|
|
1610
|
-
|
|
1611
|
-
/* 50 */ { startIndex: 61, weights: [0.0008110002381727099, 0.005136650986969471, 0.00946230161935091, 0.0069410777650773525, 0.0027783995028585196,] },
|
|
1612
|
-
|
|
1613
|
-
/* 51 */ { startIndex: 64, weights: [0.003231203882023692, 0.007237049750983715, 0.00862883497029543, 0.004773912951350212, 0.0009189908159896731,] },
|
|
1614
|
-
|
|
1615
|
-
/* 52 */ { startIndex: 66, weights: [0.001233637798577547, 0.0049433219246566296, 0.008653006516397, 0.006818502210080624, 0.003248583758249879,] },
|
|
1616
|
-
|
|
1617
|
-
/* 53 */ { startIndex: 69, weights: [0.0026164355222135782, 0.006051854696124792, 0.008880467154085636, 0.005574479699134827, 0.002268492942675948,] },
|
|
1618
|
-
|
|
1619
|
-
/* 54 */ { startIndex: 71, weights: [0.0002863667905330658, 0.003467798000201583, 0.0066492292098701, 0.00787146482616663, 0.004809896927326918, 0.0017483289120718837,] },
|
|
1620
|
-
|
|
1621
|
-
/* 55 */ { startIndex: 74, weights: [0.0009245910914614797, 0.0038708120118826628, 0.00681703258305788, 0.007283343467861414, 0.004448124207556248, 0.0016129047144204378,] },
|
|
1622
|
-
|
|
1623
|
-
/* 56 */ { startIndex: 77, weights: [0.0011703289346769452, 0.003898728871718049, 0.006627128925174475, 0.0070473202504217625, 0.004421714693307877, 0.0017961094854399562,] },
|
|
1624
|
-
|
|
1625
|
-
/* 57 */ { startIndex: 80, weights: [0.0010892992140725255, 0.003615982597693801, 0.006142666097730398, 0.007102936040610075, 0.004671447444707155, 0.002239959081634879,] },
|
|
1626
|
-
|
|
1627
|
-
/* 58 */ { startIndex: 83, weights: [0.0007392280967906117, 0.0030791081953793764, 0.005418988410383463, 0.007397185545414686, 0.005145462695509195, 0.002893739379942417, 0.0006420162972062826,] },
|
|
1628
|
-
|
|
1629
|
-
/* 59 */ { startIndex: 86, weights: [0.00017068670422304422, 0.0023375742603093386, 0.004504461772739887, 0.0066713495180010796, 0.005798479542136192, 0.003713231300935149, 0.0016279831761494279,] },
|
|
1630
|
-
|
|
1631
|
-
/* 60 */ { startIndex: 90, weights: [0.0014345343224704266, 0.0034412189852446318, 0.005447904113680124, 0.006591092795133591, 0.004660011734813452, 0.002728930441662669, 0.0007978491485118866,] },
|
|
1632
|
-
|
|
1633
|
-
/* 61 */ { startIndex: 93, weights: [0.0004075043834745884, 0.002265830524265766, 0.004124156199395657, 0.005982482805848122, 0.005700822453945875, 0.003912510350346565, 0.0021241982467472553, 0.0003358862304594368,] },
|
|
1634
|
-
|
|
1635
|
-
/* 62 */ { startIndex: 97, weights: [0.0010099108330905437, 0.002730846870690584, 0.004451782442629337, 0.006172718480229378, 0.005150905344635248, 0.0034948070533573627, 0.0018387088784947991, 0.0001826105872169137,] },
|
|
1636
|
-
|
|
1637
|
-
/* 63 */ { startIndex: 101, weights: [0.0012943691108375788, 0.002888072282075882, 0.004481775686144829, 0.006075479090213776, 0.0048866597935557365, 0.003353001084178686, 0.0018193417927250266, 0.00028568264679051936,] },
|
|
1638
|
-
|
|
1639
|
-
/* 64 */ { startIndex: 105, weights: [0.0013131388695910573, 0.0027890161145478487, 0.004264893010258675, 0.0057407706044614315, 0.004859979264438152, 0.0034397069830447435, 0.0020194342359900475, 0.0005991620710119605,] },
|
|
1640
|
-
|
|
1641
|
-
/* 65 */ { startIndex: 109, weights: [0.0011121684219688177, 0.002478930866345763, 0.0038456933107227087, 0.0052124555222690105, 0.005028639920055866, 0.0037133716978132725, 0.002398103242740035, 0.0010828346712514758,] },
|
|
1642
|
-
|
|
1643
|
-
/* 66 */ { startIndex: 113, weights: [0.0007317548734135926, 0.0019974694587290287, 0.003263183869421482, 0.004528898745775223, 0.005355686880648136, 0.004137659445405006, 0.0029196315445005894, 0.0017016039928421378, 0.0004835762665607035,] },
|
|
1644
|
-
|
|
1645
|
-
/* 67 */ { startIndex: 117, weights: [0.00020713974663522094, 0.0013792773243039846, 0.0025514145381748676, 0.003723552217707038, 0.004895689897239208, 0.004680895246565342, 0.0035529187880456448, 0.0024249425623565912, 0.0012969663366675377, 0.00016899015463422984,] },
|
|
1646
|
-
|
|
1647
|
-
/* 68 */ { startIndex: 122, weights: [0.0006545265205204487, 0.0017400053329765797, 0.0028254841454327106, 0.003910962492227554, 0.004996441304683685, 0.0042709787376224995, 0.003226396394893527, 0.002181813819333911, 0.0011372314766049385, 9.264905384043232e-05,] },
|
|
1648
|
-
|
|
1649
|
-
/* 69 */ { startIndex: 127, weights: [0.000854626705404371, 0.001859853626228869, 0.002865080488845706, 0.003870307235047221, 0.00487553421407938, 0.00408313749358058, 0.003115783678367734, 0.0021484296303242445, 0.001181075582280755, 0.0002137213887181133,] },
|
|
1650
|
-
|
|
1651
|
-
/* 70 */ { startIndex: 132, weights: [0.0008483415003865957, 0.0017792496364563704, 0.0027101580053567886, 0.0036410661414265633, 0.004571974277496338, 0.004079728852957487, 0.003183893393725157, 0.002288057701662183, 0.0013922222424298525, 0.0004963868414051831,] },
|
|
1652
|
-
|
|
1653
|
-
/* 71 */ { startIndex: 137, weights: [0.0006716204807162285, 0.0015337044605985284, 0.002395788673311472, 0.0032578727696090937, 0.004119956865906715, 0.004227725323289633, 0.0033981208689510822, 0.0025685166474431753, 0.0017389123095199466, 0.0009093079133890569, 7.970355363795534e-05,] },
|
|
1654
|
-
|
|
1655
|
-
/* 72 */ { startIndex: 142, weights: [0.0003559796023182571, 0.0011543278815224767, 0.0019526762189343572, 0.002751024439930916, 0.0035493727773427963, 0.004347721114754677, 0.0037299629766494036, 0.002961693098768592, 0.00219342322088778, 0.0014251532265916467, 0.0006568834069184959,] },
|
|
1656
|
-
|
|
1657
|
-
/* 73 */ { startIndex: 148, weights: [0.0006682946113869548, 0.0014076193328946829, 0.0021469437051564455, 0.002886268775910139, 0.0036255933810025454, 0.004154576454311609, 0.0034431067761033773, 0.0027316368650645018, 0.0020201667211949825, 0.0013086966937407851, 0.0005972267827019095,] },
|
|
1658
|
-
|
|
1659
|
-
/* 74 */ { startIndex: 153, weights: [9.926508937496692e-05, 0.0007839298341423273, 0.001468594535253942, 0.0021532592363655567, 0.0028379240538924932, 0.0035225888714194298, 0.0039915177039802074, 0.0033326479606330395, 0.002673778682947159, 0.002014909405261278, 0.0013560398947447538, 0.0006971705006435513, 3.8301113818306476e-05,] },
|
|
1660
|
-
|
|
1661
|
-
/* 75 */ { startIndex: 159, weights: [0.00010181095422012731, 0.0007358568836934865, 0.0013699028640985489, 0.0020039486698806286, 0.002637994708493352, 0.0032720407471060753, 0.003906086552888155, 0.0033682563807815313, 0.0027580985333770514, 0.002147940918803215, 0.0015377833042293787, 0.0009276255150325596, 0.000317467754939571,] },
|
|
1662
|
-
|
|
1663
|
-
/* 76 */ { startIndex: 166, weights: [0.0005530364578589797, 0.0011402058880776167, 0.0017273754347115755, 0.0023145449813455343, 0.002901714527979493, 0.003488884074613452, 0.003523340215906501, 0.002958292607218027, 0.002393245231360197, 0.0018281979719176888, 0.001263150479644537, 0.0006981031037867069, 0.0001330557424807921,] },
|
|
1664
|
-
|
|
1665
|
-
/* 77 */ { startIndex: 172, weights: [0.0002608386566862464, 0.0008045974536798894, 0.0013483562506735325, 0.0018921148730441928, 0.0024358737282454967, 0.002979632467031479, 0.003523391205817461, 0.003251380519941449, 0.0027281083166599274, 0.002204835880547762, 0.001681563793681562, 0.001158291706815362, 0.0006350195035338402, 0.00011174729297636077,] },
|
|
1666
|
-
|
|
1667
|
-
/* 78 */ { startIndex: 179, weights: [0.0003849811910185963, 0.0008885387214832008, 0.001392096164636314, 0.0018956535495817661, 0.00239921105094254, 0.002902768552303314, 0.0034063260536640882, 0.003132763085886836, 0.0026481777895241976, 0.0021635922603309155, 0.0016790067311376333, 0.0011944210855290294, 0.0007098356145434082, 0.00022525011445395648,] },
|
|
1668
|
-
|
|
1669
|
-
/* 79 */ { startIndex: 186, weights: [0.000366741674952209, 0.0008330700220540166, 0.0012993983691558242, 0.0017657268326729536, 0.0022320549469441175, 0.002698383294045925, 0.0031647118739783764, 0.003141313325613737, 0.002692554146051407, 0.0022437951993197203, 0.00179503601975739, 0.0013462770730257034, 0.000897518009878695, 0.0004487590049393475,] },
|
|
1670
|
-
]
|
|
1671
|
-
|
|
1672
1731
|
export async function loadPackagesAndGetPaths(modelName: WhisperModelName | undefined, languageCode: string | undefined) {
|
|
1673
1732
|
if (modelName) {
|
|
1674
1733
|
modelName = normalizeWhisperModelName(modelName, languageCode)
|
|
@@ -1718,6 +1777,8 @@ export type WhisperTokenData = {
|
|
|
1718
1777
|
text: string
|
|
1719
1778
|
}
|
|
1720
1779
|
|
|
1780
|
+
export type WhisperLogitFilter = (logits: number[], decodedTokens: number[], isFirstPart: boolean, isFinalPart: boolean) => number[]
|
|
1781
|
+
|
|
1721
1782
|
export type WhisperModelName = 'tiny' | 'tiny.en' | 'base' | 'base.en' | 'small' | 'small.en' | 'medium' | 'medium.en' | 'large' | 'large-v1' | 'large-v2' | 'large-v3'
|
|
1722
1783
|
export type WhisperTask = 'transcribe' | 'translate' | 'detect-language'
|
|
1723
1784
|
|
|
@@ -1855,6 +1916,169 @@ const alignmentHeadsIndexes: { [name in WhisperModelName]: number[] } = {
|
|
|
1855
1916
|
'large': [212, 277, 331, 332, 333, 355, 356, 364, 371, 379, 391, 422, 423, 443, 449, 452, 465, 467, 473, 505, 521, 532, 555],
|
|
1856
1917
|
}
|
|
1857
1918
|
|
|
1919
|
+
const filterbanks: Filterbank[] = [
|
|
1920
|
+
/* 0 */ { startIndex: 1, weights: [0.02486259490251541,] },
|
|
1921
|
+
|
|
1922
|
+
/* 1 */ { startIndex: 1, weights: [0.001990821911022067, 0.022871771827340126,] },
|
|
1923
|
+
|
|
1924
|
+
/* 2 */ { startIndex: 2, weights: [0.003981643822044134, 0.02088095061480999,] },
|
|
1925
|
+
|
|
1926
|
+
/* 3 */ { startIndex: 3, weights: [0.0059724655002355576, 0.018890129402279854,] },
|
|
1927
|
+
|
|
1928
|
+
/* 4 */ { startIndex: 4, weights: [0.007963287644088268, 0.01689930632710457,] },
|
|
1929
|
+
|
|
1930
|
+
/* 5 */ { startIndex: 5, weights: [0.009954108856618404, 0.014908484183251858,] },
|
|
1931
|
+
|
|
1932
|
+
/* 6 */ { startIndex: 6, weights: [0.011944931000471115, 0.012917662039399147,] },
|
|
1933
|
+
|
|
1934
|
+
/* 7 */ { startIndex: 7, weights: [0.013935752213001251, 0.010926840826869011,] },
|
|
1935
|
+
|
|
1936
|
+
/* 8 */ { startIndex: 8, weights: [0.015926575288176537, 0.0089360186830163,] },
|
|
1937
|
+
|
|
1938
|
+
/* 9 */ { startIndex: 9, weights: [0.017917396500706673, 0.006945197004824877,] },
|
|
1939
|
+
|
|
1940
|
+
/* 10 */ { startIndex: 10, weights: [0.01990821771323681, 0.004954374860972166,] },
|
|
1941
|
+
|
|
1942
|
+
/* 11 */ { startIndex: 11, weights: [0.021899040788412094, 0.0029635531827807426,] },
|
|
1943
|
+
|
|
1944
|
+
/* 12 */ { startIndex: 12, weights: [0.02388986200094223, 0.0009727313299663365,] },
|
|
1945
|
+
|
|
1946
|
+
/* 13 */ { startIndex: 13, weights: [0.025880683213472366,] },
|
|
1947
|
+
|
|
1948
|
+
/* 14 */ { startIndex: 14, weights: [0.025835324078798294,] },
|
|
1949
|
+
|
|
1950
|
+
/* 15 */ { startIndex: 14, weights: [0.0010180906392633915, 0.023844502866268158,] },
|
|
1951
|
+
|
|
1952
|
+
/* 16 */ { startIndex: 15, weights: [0.003008912317454815, 0.021853681653738022,] },
|
|
1953
|
+
|
|
1954
|
+
/* 17 */ { startIndex: 16, weights: [0.004999734461307526, 0.019862858578562737,] },
|
|
1955
|
+
|
|
1956
|
+
/* 18 */ { startIndex: 17, weights: [0.006990555673837662, 0.0178720373660326,] },
|
|
1957
|
+
|
|
1958
|
+
/* 19 */ { startIndex: 18, weights: [0.008981377817690372, 0.015881216153502464,] },
|
|
1959
|
+
|
|
1960
|
+
/* 20 */ { startIndex: 19, weights: [0.010972199961543083, 0.013890394009649754,] },
|
|
1961
|
+
|
|
1962
|
+
/* 21 */ { startIndex: 20, weights: [0.01296302117407322, 0.011899571865797043,] },
|
|
1963
|
+
|
|
1964
|
+
/* 22 */ { startIndex: 21, weights: [0.01495384331792593, 0.009908749721944332,] },
|
|
1965
|
+
|
|
1966
|
+
/* 23 */ { startIndex: 22, weights: [0.01694466546177864, 0.007917927578091621,] },
|
|
1967
|
+
|
|
1968
|
+
/* 24 */ { startIndex: 23, weights: [0.018935488536953926, 0.005927106365561485,] },
|
|
1969
|
+
|
|
1970
|
+
/* 25 */ { startIndex: 24, weights: [0.020874010398983955, 0.004040425643324852,] },
|
|
1971
|
+
|
|
1972
|
+
/* 26 */ { startIndex: 25, weights: [0.022114217281341553, 0.0033186059445142746,] },
|
|
1973
|
+
|
|
1974
|
+
/* 27 */ { startIndex: 26, weights: [0.02173672430217266, 0.0036109676584601402,] },
|
|
1975
|
+
|
|
1976
|
+
/* 28 */ { startIndex: 27, weights: [0.020497702062129974, 0.004762193653732538,] },
|
|
1977
|
+
|
|
1978
|
+
/* 29 */ { startIndex: 28, weights: [0.018486659973859787, 0.006592618301510811,] },
|
|
1979
|
+
|
|
1980
|
+
/* 30 */ { startIndex: 29, weights: [0.01585603691637516, 0.00896277092397213,] },
|
|
1981
|
+
|
|
1982
|
+
/* 31 */ { startIndex: 30, weights: [0.012738768011331558, 0.011751330457627773,] },
|
|
1983
|
+
|
|
1984
|
+
/* 32 */ { startIndex: 31, weights: [0.009250369854271412, 0.014853144995868206,] },
|
|
1985
|
+
|
|
1986
|
+
/* 33 */ { startIndex: 32, weights: [0.005490840878337622, 0.018177473917603493, 0.0028155462350696325,] },
|
|
1987
|
+
|
|
1988
|
+
/* 34 */ { startIndex: 33, weights: [0.0015463664894923568, 0.01632951945066452, 0.007420188747346401,] },
|
|
1989
|
+
|
|
1990
|
+
/* 35 */ { startIndex: 35, weights: [0.011181050911545753, 0.012018864043056965,] },
|
|
1991
|
+
|
|
1992
|
+
/* 36 */ { startIndex: 36, weights: [0.006065350491553545, 0.016561277210712433, 0.004360878840088844,] },
|
|
1993
|
+
|
|
1994
|
+
/* 37 */ { startIndex: 37, weights: [0.0010297985281795263, 0.012770536355674267, 0.009707189165055752,] },
|
|
1995
|
+
|
|
1996
|
+
/* 38 */ { startIndex: 39, weights: [0.006986402906477451, 0.01485429983586073, 0.004391219466924667,] },
|
|
1997
|
+
|
|
1998
|
+
/* 39 */ { startIndex: 40, weights: [0.001418047584593296, 0.011486922390758991, 0.010089744813740253, 0.00040022286702878773,] },
|
|
1999
|
+
|
|
2000
|
+
/* 40 */ { startIndex: 42, weights: [0.005411104764789343, 0.014735566452145576, 0.006518189795315266,] },
|
|
2001
|
+
|
|
2002
|
+
/* 41 */ { startIndex: 44, weights: [0.00827841367572546, 0.012277561239898205, 0.00396781275048852,] },
|
|
2003
|
+
|
|
2004
|
+
/* 42 */ { startIndex: 45, weights: [0.002187808509916067, 0.010184479877352715, 0.00998187530785799, 0.0022864851634949446,] },
|
|
2005
|
+
|
|
2006
|
+
/* 43 */ { startIndex: 47, weights: [0.00386943481862545, 0.011274894699454308, 0.008466221392154694, 0.0013397691072896123,] },
|
|
2007
|
+
|
|
2008
|
+
/* 44 */ { startIndex: 49, weights: [0.004820294212549925, 0.011678251437842846, 0.007608682848513126, 0.0010091039584949613,] },
|
|
2009
|
+
|
|
2010
|
+
/* 45 */ { startIndex: 51, weights: [0.005156961735337973, 0.011507894843816757, 0.007301822770386934, 0.0011901655234396458,] },
|
|
2011
|
+
|
|
2012
|
+
/* 46 */ { startIndex: 53, weights: [0.004982104524970055, 0.010863498784601688, 0.007451189681887627, 0.001791381393559277,] },
|
|
2013
|
+
|
|
2014
|
+
/* 47 */ { startIndex: 55, weights: [0.004385921638458967, 0.009832492098212242, 0.007973956875503063, 0.002732589840888977,] },
|
|
2015
|
+
|
|
2016
|
+
/* 48 */ { startIndex: 57, weights: [0.0034474546555429697, 0.008491347543895245, 0.008797688409686089, 0.00394382793456316,] },
|
|
2017
|
+
|
|
2018
|
+
/* 49 */ { startIndex: 59, weights: [0.0022357646375894547, 0.0069067515432834625, 0.009859241545200348, 0.005364237818866968, 0.0008692338014952838,] },
|
|
2019
|
+
|
|
2020
|
+
/* 50 */ { startIndex: 61, weights: [0.0008110002381727099, 0.005136650986969471, 0.00946230161935091, 0.0069410777650773525, 0.0027783995028585196,] },
|
|
2021
|
+
|
|
2022
|
+
/* 51 */ { startIndex: 64, weights: [0.003231203882023692, 0.007237049750983715, 0.00862883497029543, 0.004773912951350212, 0.0009189908159896731,] },
|
|
2023
|
+
|
|
2024
|
+
/* 52 */ { startIndex: 66, weights: [0.001233637798577547, 0.0049433219246566296, 0.008653006516397, 0.006818502210080624, 0.003248583758249879,] },
|
|
2025
|
+
|
|
2026
|
+
/* 53 */ { startIndex: 69, weights: [0.0026164355222135782, 0.006051854696124792, 0.008880467154085636, 0.005574479699134827, 0.002268492942675948,] },
|
|
2027
|
+
|
|
2028
|
+
/* 54 */ { startIndex: 71, weights: [0.0002863667905330658, 0.003467798000201583, 0.0066492292098701, 0.00787146482616663, 0.004809896927326918, 0.0017483289120718837,] },
|
|
2029
|
+
|
|
2030
|
+
/* 55 */ { startIndex: 74, weights: [0.0009245910914614797, 0.0038708120118826628, 0.00681703258305788, 0.007283343467861414, 0.004448124207556248, 0.0016129047144204378,] },
|
|
2031
|
+
|
|
2032
|
+
/* 56 */ { startIndex: 77, weights: [0.0011703289346769452, 0.003898728871718049, 0.006627128925174475, 0.0070473202504217625, 0.004421714693307877, 0.0017961094854399562,] },
|
|
2033
|
+
|
|
2034
|
+
/* 57 */ { startIndex: 80, weights: [0.0010892992140725255, 0.003615982597693801, 0.006142666097730398, 0.007102936040610075, 0.004671447444707155, 0.002239959081634879,] },
|
|
2035
|
+
|
|
2036
|
+
/* 58 */ { startIndex: 83, weights: [0.0007392280967906117, 0.0030791081953793764, 0.005418988410383463, 0.007397185545414686, 0.005145462695509195, 0.002893739379942417, 0.0006420162972062826,] },
|
|
2037
|
+
|
|
2038
|
+
/* 59 */ { startIndex: 86, weights: [0.00017068670422304422, 0.0023375742603093386, 0.004504461772739887, 0.0066713495180010796, 0.005798479542136192, 0.003713231300935149, 0.0016279831761494279,] },
|
|
2039
|
+
|
|
2040
|
+
/* 60 */ { startIndex: 90, weights: [0.0014345343224704266, 0.0034412189852446318, 0.005447904113680124, 0.006591092795133591, 0.004660011734813452, 0.002728930441662669, 0.0007978491485118866,] },
|
|
2041
|
+
|
|
2042
|
+
/* 61 */ { startIndex: 93, weights: [0.0004075043834745884, 0.002265830524265766, 0.004124156199395657, 0.005982482805848122, 0.005700822453945875, 0.003912510350346565, 0.0021241982467472553, 0.0003358862304594368,] },
|
|
2043
|
+
|
|
2044
|
+
/* 62 */ { startIndex: 97, weights: [0.0010099108330905437, 0.002730846870690584, 0.004451782442629337, 0.006172718480229378, 0.005150905344635248, 0.0034948070533573627, 0.0018387088784947991, 0.0001826105872169137,] },
|
|
2045
|
+
|
|
2046
|
+
/* 63 */ { startIndex: 101, weights: [0.0012943691108375788, 0.002888072282075882, 0.004481775686144829, 0.006075479090213776, 0.0048866597935557365, 0.003353001084178686, 0.0018193417927250266, 0.00028568264679051936,] },
|
|
2047
|
+
|
|
2048
|
+
/* 64 */ { startIndex: 105, weights: [0.0013131388695910573, 0.0027890161145478487, 0.004264893010258675, 0.0057407706044614315, 0.004859979264438152, 0.0034397069830447435, 0.0020194342359900475, 0.0005991620710119605,] },
|
|
2049
|
+
|
|
2050
|
+
/* 65 */ { startIndex: 109, weights: [0.0011121684219688177, 0.002478930866345763, 0.0038456933107227087, 0.0052124555222690105, 0.005028639920055866, 0.0037133716978132725, 0.002398103242740035, 0.0010828346712514758,] },
|
|
2051
|
+
|
|
2052
|
+
/* 66 */ { startIndex: 113, weights: [0.0007317548734135926, 0.0019974694587290287, 0.003263183869421482, 0.004528898745775223, 0.005355686880648136, 0.004137659445405006, 0.0029196315445005894, 0.0017016039928421378, 0.0004835762665607035,] },
|
|
2053
|
+
|
|
2054
|
+
/* 67 */ { startIndex: 117, weights: [0.00020713974663522094, 0.0013792773243039846, 0.0025514145381748676, 0.003723552217707038, 0.004895689897239208, 0.004680895246565342, 0.0035529187880456448, 0.0024249425623565912, 0.0012969663366675377, 0.00016899015463422984,] },
|
|
2055
|
+
|
|
2056
|
+
/* 68 */ { startIndex: 122, weights: [0.0006545265205204487, 0.0017400053329765797, 0.0028254841454327106, 0.003910962492227554, 0.004996441304683685, 0.0042709787376224995, 0.003226396394893527, 0.002181813819333911, 0.0011372314766049385, 9.264905384043232e-05,] },
|
|
2057
|
+
|
|
2058
|
+
/* 69 */ { startIndex: 127, weights: [0.000854626705404371, 0.001859853626228869, 0.002865080488845706, 0.003870307235047221, 0.00487553421407938, 0.00408313749358058, 0.003115783678367734, 0.0021484296303242445, 0.001181075582280755, 0.0002137213887181133,] },
|
|
2059
|
+
|
|
2060
|
+
/* 70 */ { startIndex: 132, weights: [0.0008483415003865957, 0.0017792496364563704, 0.0027101580053567886, 0.0036410661414265633, 0.004571974277496338, 0.004079728852957487, 0.003183893393725157, 0.002288057701662183, 0.0013922222424298525, 0.0004963868414051831,] },
|
|
2061
|
+
|
|
2062
|
+
/* 71 */ { startIndex: 137, weights: [0.0006716204807162285, 0.0015337044605985284, 0.002395788673311472, 0.0032578727696090937, 0.004119956865906715, 0.004227725323289633, 0.0033981208689510822, 0.0025685166474431753, 0.0017389123095199466, 0.0009093079133890569, 7.970355363795534e-05,] },
|
|
2063
|
+
|
|
2064
|
+
/* 72 */ { startIndex: 142, weights: [0.0003559796023182571, 0.0011543278815224767, 0.0019526762189343572, 0.002751024439930916, 0.0035493727773427963, 0.004347721114754677, 0.0037299629766494036, 0.002961693098768592, 0.00219342322088778, 0.0014251532265916467, 0.0006568834069184959,] },
|
|
2065
|
+
|
|
2066
|
+
/* 73 */ { startIndex: 148, weights: [0.0006682946113869548, 0.0014076193328946829, 0.0021469437051564455, 0.002886268775910139, 0.0036255933810025454, 0.004154576454311609, 0.0034431067761033773, 0.0027316368650645018, 0.0020201667211949825, 0.0013086966937407851, 0.0005972267827019095,] },
|
|
2067
|
+
|
|
2068
|
+
/* 74 */ { startIndex: 153, weights: [9.926508937496692e-05, 0.0007839298341423273, 0.001468594535253942, 0.0021532592363655567, 0.0028379240538924932, 0.0035225888714194298, 0.0039915177039802074, 0.0033326479606330395, 0.002673778682947159, 0.002014909405261278, 0.0013560398947447538, 0.0006971705006435513, 3.8301113818306476e-05,] },
|
|
2069
|
+
|
|
2070
|
+
/* 75 */ { startIndex: 159, weights: [0.00010181095422012731, 0.0007358568836934865, 0.0013699028640985489, 0.0020039486698806286, 0.002637994708493352, 0.0032720407471060753, 0.003906086552888155, 0.0033682563807815313, 0.0027580985333770514, 0.002147940918803215, 0.0015377833042293787, 0.0009276255150325596, 0.000317467754939571,] },
|
|
2071
|
+
|
|
2072
|
+
/* 76 */ { startIndex: 166, weights: [0.0005530364578589797, 0.0011402058880776167, 0.0017273754347115755, 0.0023145449813455343, 0.002901714527979493, 0.003488884074613452, 0.003523340215906501, 0.002958292607218027, 0.002393245231360197, 0.0018281979719176888, 0.001263150479644537, 0.0006981031037867069, 0.0001330557424807921,] },
|
|
2073
|
+
|
|
2074
|
+
/* 77 */ { startIndex: 172, weights: [0.0002608386566862464, 0.0008045974536798894, 0.0013483562506735325, 0.0018921148730441928, 0.0024358737282454967, 0.002979632467031479, 0.003523391205817461, 0.003251380519941449, 0.0027281083166599274, 0.002204835880547762, 0.001681563793681562, 0.001158291706815362, 0.0006350195035338402, 0.00011174729297636077,] },
|
|
2075
|
+
|
|
2076
|
+
/* 78 */ { startIndex: 179, weights: [0.0003849811910185963, 0.0008885387214832008, 0.001392096164636314, 0.0018956535495817661, 0.00239921105094254, 0.002902768552303314, 0.0034063260536640882, 0.003132763085886836, 0.0026481777895241976, 0.0021635922603309155, 0.0016790067311376333, 0.0011944210855290294, 0.0007098356145434082, 0.00022525011445395648,] },
|
|
2077
|
+
|
|
2078
|
+
/* 79 */ { startIndex: 186, weights: [0.000366741674952209, 0.0008330700220540166, 0.0012993983691558242, 0.0017657268326729536, 0.0022320549469441175, 0.002698383294045925, 0.0031647118739783764, 0.003141313325613737, 0.002692554146051407, 0.0022437951993197203, 0.00179503601975739, 0.0013462770730257034, 0.000897518009878695, 0.0004487590049393475,] },
|
|
2079
|
+
]
|
|
2080
|
+
|
|
2081
|
+
// Recognition options
|
|
1858
2082
|
export interface WhisperOptions {
|
|
1859
2083
|
model?: WhisperModelName
|
|
1860
2084
|
temperature?: number
|
|
@@ -1865,6 +2089,10 @@ export interface WhisperOptions {
|
|
|
1865
2089
|
maxTokensPerPart?: number
|
|
1866
2090
|
suppressRepetition?: boolean
|
|
1867
2091
|
decodeTimestampTokens?: boolean
|
|
2092
|
+
endTokenThreshold?: number
|
|
2093
|
+
includeEndTokenInCandidates?: boolean
|
|
2094
|
+
encoderProvider?: OnnxExecutionProvider
|
|
2095
|
+
decoderProvider?: OnnxExecutionProvider
|
|
1868
2096
|
seed?: number
|
|
1869
2097
|
}
|
|
1870
2098
|
|
|
@@ -1878,5 +2106,54 @@ export const defaultWhisperOptions: WhisperOptions = {
|
|
|
1878
2106
|
maxTokensPerPart: 250,
|
|
1879
2107
|
suppressRepetition: true,
|
|
1880
2108
|
decodeTimestampTokens: true,
|
|
2109
|
+
endTokenThreshold: 0.9,
|
|
2110
|
+
includeEndTokenInCandidates: true,
|
|
2111
|
+
encoderProvider: undefined,
|
|
2112
|
+
decoderProvider: undefined,
|
|
1881
2113
|
seed: undefined,
|
|
1882
2114
|
}
|
|
2115
|
+
|
|
2116
|
+
// Alignment options
|
|
2117
|
+
export interface WhisperAlignmentOptions {
|
|
2118
|
+
model?: WhisperModelName
|
|
2119
|
+
endTokenThreshold?: number
|
|
2120
|
+
encoderProvider?: OnnxExecutionProvider
|
|
2121
|
+
decoderProvider?: OnnxExecutionProvider
|
|
2122
|
+
}
|
|
2123
|
+
|
|
2124
|
+
export const defaultWhisperAlignmentOptions: WhisperAlignmentOptions = {
|
|
2125
|
+
model: undefined,
|
|
2126
|
+
endTokenThreshold: 0.9,
|
|
2127
|
+
encoderProvider: undefined,
|
|
2128
|
+
decoderProvider: undefined
|
|
2129
|
+
}
|
|
2130
|
+
|
|
2131
|
+
// Language detection options
|
|
2132
|
+
export interface WhisperLanguageDetectionOptions {
|
|
2133
|
+
model?: WhisperModelName
|
|
2134
|
+
temperature?: number
|
|
2135
|
+
encoderProvider?: OnnxExecutionProvider
|
|
2136
|
+
decoderProvider?: OnnxExecutionProvider
|
|
2137
|
+
}
|
|
2138
|
+
|
|
2139
|
+
export const defaultWhisperLanguageDetectionOptions: WhisperLanguageDetectionOptions = {
|
|
2140
|
+
model: undefined,
|
|
2141
|
+
temperature: 1.0,
|
|
2142
|
+
encoderProvider: undefined,
|
|
2143
|
+
decoderProvider: undefined,
|
|
2144
|
+
}
|
|
2145
|
+
|
|
2146
|
+
// Voice activity detection options
|
|
2147
|
+
export interface WhisperVADOptions {
|
|
2148
|
+
model?: WhisperModelName
|
|
2149
|
+
temperature?: number
|
|
2150
|
+
encoderProvider?: OnnxExecutionProvider
|
|
2151
|
+
decoderProvider?: OnnxExecutionProvider
|
|
2152
|
+
}
|
|
2153
|
+
|
|
2154
|
+
export const defaultWhisperVADOptions: WhisperVADOptions = {
|
|
2155
|
+
model: undefined,
|
|
2156
|
+
temperature: 1.0,
|
|
2157
|
+
encoderProvider: undefined,
|
|
2158
|
+
decoderProvider: undefined,
|
|
2159
|
+
}
|