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.
Files changed (122) hide show
  1. package/README.md +26 -23
  2. package/data/schemas/options.json +177 -36
  3. package/dist/alignment/SpeechAlignment.d.ts +1 -1
  4. package/dist/alignment/SpeechAlignment.js +1 -1
  5. package/dist/alignment/SpeechAlignment.js.map +1 -1
  6. package/dist/api/API.d.ts +1 -0
  7. package/dist/api/API.js +1 -0
  8. package/dist/api/API.js.map +1 -1
  9. package/dist/api/APIOptions.d.ts +1 -0
  10. package/dist/api/Alignment.d.ts +3 -3
  11. package/dist/api/Alignment.js +5 -10
  12. package/dist/api/Alignment.js.map +1 -1
  13. package/dist/api/LanguageDetection.d.ts +5 -7
  14. package/dist/api/LanguageDetection.js +3 -2
  15. package/dist/api/LanguageDetection.js.map +1 -1
  16. package/dist/api/Recognition.d.ts +4 -5
  17. package/dist/api/Recognition.js +5 -8
  18. package/dist/api/Recognition.js.map +1 -1
  19. package/dist/api/SourceSeparation.d.ts +2 -0
  20. package/dist/api/SourceSeparation.js +4 -2
  21. package/dist/api/SourceSeparation.js.map +1 -1
  22. package/dist/api/Synthesis.d.ts +3 -1
  23. package/dist/api/Synthesis.js +9 -10
  24. package/dist/api/Synthesis.js.map +1 -1
  25. package/dist/api/Translation.d.ts +1 -1
  26. package/dist/api/Translation.js +4 -8
  27. package/dist/api/Translation.js.map +1 -1
  28. package/dist/api/TranslationAlignment.d.ts +31 -0
  29. package/dist/api/TranslationAlignment.js +121 -0
  30. package/dist/api/TranslationAlignment.js.map +1 -0
  31. package/dist/api/VoiceActivityDetection.d.ts +5 -1
  32. package/dist/api/VoiceActivityDetection.js +38 -2
  33. package/dist/api/VoiceActivityDetection.js.map +1 -1
  34. package/dist/audio/AudioPlayer.js +6 -1
  35. package/dist/audio/AudioPlayer.js.map +1 -1
  36. package/dist/cli/CLI.js +85 -0
  37. package/dist/cli/CLI.js.map +1 -1
  38. package/dist/dsp/FFT.js.map +1 -1
  39. package/dist/math/MedianFilter.d.ts +5 -0
  40. package/dist/math/MedianFilter.js +102 -0
  41. package/dist/math/MedianFilter.js.map +1 -0
  42. package/dist/math/VectorMath.d.ts +0 -2
  43. package/dist/math/VectorMath.js +1 -25
  44. package/dist/math/VectorMath.js.map +1 -1
  45. package/dist/recognition/OpenAICloudSTT.d.ts +1 -1
  46. package/dist/recognition/OpenAICloudSTT.js.map +1 -1
  47. package/dist/recognition/SileroSTT.d.ts +22 -1
  48. package/dist/recognition/SileroSTT.js +122 -95
  49. package/dist/recognition/SileroSTT.js.map +1 -1
  50. package/dist/recognition/WhisperCppSTT.js +1 -1
  51. package/dist/recognition/WhisperCppSTT.js.map +1 -1
  52. package/dist/recognition/WhisperSTT.d.ts +52 -19
  53. package/dist/recognition/WhisperSTT.js +645 -494
  54. package/dist/recognition/WhisperSTT.js.map +1 -1
  55. package/dist/server/Server.js.map +1 -1
  56. package/dist/source-separation/MDXNetSourceSeparation.d.ts +5 -3
  57. package/dist/source-separation/MDXNetSourceSeparation.js +26 -19
  58. package/dist/source-separation/MDXNetSourceSeparation.js.map +1 -1
  59. package/dist/speech-language-detection/SileroLanguageDetection.d.ts +15 -9
  60. package/dist/speech-language-detection/SileroLanguageDetection.js +23 -16
  61. package/dist/speech-language-detection/SileroLanguageDetection.js.map +1 -1
  62. package/dist/synthesis/EspeakTTS.js +4 -0
  63. package/dist/synthesis/EspeakTTS.js.map +1 -1
  64. package/dist/synthesis/GoogleCloudTTS.js.map +1 -1
  65. package/dist/synthesis/VitsTTS.d.ts +8 -6
  66. package/dist/synthesis/VitsTTS.js +36 -31
  67. package/dist/synthesis/VitsTTS.js.map +1 -1
  68. package/dist/tests/Test.js.map +1 -1
  69. package/dist/utilities/OnnxUtilities.d.ts +14 -0
  70. package/dist/utilities/OnnxUtilities.js +43 -0
  71. package/dist/utilities/OnnxUtilities.js.map +1 -0
  72. package/dist/utilities/Utilities.d.ts +4 -8
  73. package/dist/utilities/Utilities.js +35 -58
  74. package/dist/utilities/Utilities.js.map +1 -1
  75. package/dist/voice-activity-detection/SileroVAD.d.ts +5 -3
  76. package/dist/voice-activity-detection/SileroVAD.js +9 -11
  77. package/dist/voice-activity-detection/SileroVAD.js.map +1 -1
  78. package/docs/API.md +54 -34
  79. package/docs/CLI.md +25 -13
  80. package/docs/Contributing.md +4 -2
  81. package/docs/Engines.md +43 -32
  82. package/docs/Licenses.md +3 -4
  83. package/docs/Options.md +47 -11
  84. package/docs/Releases.md +4 -0
  85. package/docs/Server.md +8 -6
  86. package/docs/Tasklist.md +39 -52
  87. package/docs/Technical.md +1 -1
  88. package/package.json +8 -12
  89. package/src/alignment/SpeechAlignment.ts +1 -1
  90. package/src/api/API.ts +1 -0
  91. package/src/api/APIOptions.ts +1 -0
  92. package/src/api/Alignment.ts +10 -14
  93. package/src/api/LanguageDetection.ts +14 -10
  94. package/src/api/Recognition.ts +17 -10
  95. package/src/api/SourceSeparation.ts +7 -2
  96. package/src/api/Synthesis.ts +26 -11
  97. package/src/api/Translation.ts +14 -8
  98. package/src/api/TranslationAlignment.ts +213 -0
  99. package/src/api/VoiceActivityDetection.ts +66 -3
  100. package/src/audio/AudioPlayer.ts +6 -2
  101. package/src/cli/CLI.ts +121 -2
  102. package/src/dsp/FFT.ts +3 -0
  103. package/src/math/MedianFilter.ts +124 -0
  104. package/src/math/VectorMath.ts +1 -36
  105. package/src/recognition/OpenAICloudSTT.ts +27 -27
  106. package/src/recognition/SileroSTT.ts +149 -102
  107. package/src/recognition/WhisperCppSTT.ts +1 -1
  108. package/src/recognition/WhisperSTT.ts +961 -684
  109. package/src/server/Server.ts +1 -1
  110. package/src/source-separation/MDXNetSourceSeparation.ts +35 -19
  111. package/src/speech-language-detection/SileroLanguageDetection.ts +53 -33
  112. package/src/synthesis/EspeakTTS.ts +8 -0
  113. package/src/synthesis/GoogleCloudTTS.ts +12 -1
  114. package/src/synthesis/VitsTTS.ts +57 -46
  115. package/src/tests/Test.ts +1 -1
  116. package/src/utilities/OnnxUtilities.ts +68 -0
  117. package/src/utilities/Utilities.ts +38 -66
  118. package/src/voice-activity-detection/SileroVAD.ts +15 -15
  119. package/dist/utilities/NdArrayUtilities.d.ts +0 -3
  120. package/dist/utilities/NdArrayUtilities.js +0 -23
  121. package/dist/utilities/NdArrayUtilities.js.map +0 -1
  122. 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, getIntegerRange, getRepetitionScoreRelativeToFirstSubstring, getUTF32Chars, logToStderr, splitFloat32Array, yieldToEventLoop } from '../utilities/Utilities.js'
6
- import { indexOfMax, logOfVector, logSumExp, meanOfVector, medianFilter, softmax, stdDeviationOfVector } from '../math/VectorMath.js'
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 whisper = new Whisper(modelName, modelDir, seed)
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(sourceRawAudio: RawAudio, referenceText: string, modelName: WhisperModelName, modelDir: string, sourceLanguage: string) {
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 whisper = new Whisper(modelName, modelDir)
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, referenceText, sourceLanguage)
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(sourceRawAudio: RawAudio, modelName: WhisperModelName, modelDir: string, temperature: number) {
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 whisper = new Whisper(modelName, modelDir)
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 class Whisper {
108
- modelName: WhisperModelName
109
- modelDir: string
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(modelName: WhisperModelName, modelDir: string, rngSeed = 461845907) {
143
- this.modelName = modelName
144
- this.modelDir = modelDir
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 initializeIfNeeded() {
188
- await this.initializeTokenizerIfNeeded()
189
- await this.initializeEncoderSessionIfNeeded()
190
- await this.initializeDecoderSessionIfNeeded()
191
- }
192
-
193
- async initializeTokenizerIfNeeded() {
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
- crossAttentionQKs: partCrossAttentionQKs,
342
- decodedTokensConfidence
343
- } = await this.decodeTokens(audioPartFeatures, initialTokens, audioPartDuration, isFirstPart, isFinalPart, options)
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
- audioOffset = audioEndOffset
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, referenceText: string, language: string) {
396
- await this.initializeIfNeeded()
477
+ async align(rawAudio: RawAudio, transcript: string, sourceLanguage: string, task: 'transcribe' | 'translate', whisperAlignmentOptions: WhisperAlignmentOptions) {
478
+ await this.initializeTokenizerIfNeeded()
397
479
 
398
- const logger = new Logger()
480
+ whisperAlignmentOptions = extendDeep(defaultWhisperAlignmentOptions, whisperAlignmentOptions)
481
+
482
+ const shouldSplitToSentences = false
399
483
 
400
- await logger.startAsync('Prepare for alignment')
484
+ let simplifiedTranscript = ''
401
485
 
402
- referenceText = referenceText.replaceAll(/\s+/g, ' ')
486
+ if (shouldSplitToSentences) {
487
+ const sentences = splitToSentences(transcript, 'en')
403
488
 
404
- const audioDuration = Math.min(getRawAudioDuration(rawAudio), 30)
405
- const audioFrameCount = this.secondsToFrame(audioDuration)
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
- const initialTokens = this.getTextStartTokens(language, 'transcribe', true)
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
- let tokens = [...initialTokens, ...this.textToTokens(referenceText), endOfTextToken]
510
+ const logitFilter: WhisperLogitFilter = (logits, decodedTokens, isFirstPart, isFinalPart) => {
511
+ const decodedTextTokens = decodedTokens.filter(token => this.isTextToken(token))
412
512
 
413
- logger.end()
414
- const audioFeatures = await this.encodeAudio(rawAudio)
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
- await logger.startAsync('Infer cross-attention QKs')
417
- let crossAttentionQKs = await this.inferCrossAttentionQKs(tokens, audioFeatures)
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
- tokens = tokens.slice(initialTokens.length, tokens.length - 1)
420
- crossAttentionQKs = crossAttentionQKs.slice(initialTokens.length, crossAttentionQKs.length - 1)
527
+ return -Infinity
528
+ })
421
529
 
422
- await logger.startAsync('Extract word timeline')
423
- const alignmentPath = await this.findAlignmentPathFromQKs(crossAttentionQKs, tokens, 0, audioFrameCount)//, this.getAlignmentHeadIndexes())
424
- const tokenTimeline = await this.getTokenTimelineFromAlignmentPath(alignmentPath, tokens, 0, audioDuration)
530
+ return newLogits
531
+ }
425
532
 
426
- const wordTimeline = this.tokenTimelineToWordTimeline(tokenTimeline, language)
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
- logger.end()
549
+ // Recognize
550
+ const { timeline } = await this.recognize(rawAudio, task, sourceLanguage, options, logitFilter)
429
551
 
430
- return wordTimeline
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 spaceToken = this.textToTokens(' ')[0]
679
+ const suppressedTextTokens = this.getSuppressedTextTokens()
680
+ const suppressedMetadataTokens = this.getSuppressedMetadataTokens()
681
+
682
+ const allowedPunctuationMarks = this.getAllowedPunctuationMarks()
518
683
 
519
- const maxDecodedTokenCount = options.maxTokensPerPart!
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[][] = [new Array(1501)]
526
-
527
- let lastTimestampTokenIndex = -1
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 < maxDecodedTokenCount; decodedTokenCount++) {
714
+ for (let decodedTokenCount = 0; decodedTokenCount < options.maxTokensPerPart!; decodedTokenCount++) {
540
715
  const isInitialState = decodedTokens.length == initialTokens.length
541
716
 
542
- const tokensToDecode = isInitialState ? decodedTokens : [decodedTokens[decodedTokens.length - 1]]
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 and run decoder
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 = { tokens: tokensTensor, audio_features: audioFeatures, kv_cache: kvCacheTensor, offset: offsetTensor }
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
- // Compute logits
573
- const resultLogits = splitFloat32Array(logitsBuffer, logitsBuffer.length / decoderOutputs['logits'].dims[1])
574
- const allTokenLogits = Array.from(resultLogits[resultLogits.length - 1])
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
- const shouldSuppressToken = isWrongTokenForInitialState || isInSuppressedList
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
- if (shouldSuppressToken) {
588
- allTokenLogits[logitIndex] = -Infinity
589
- }
762
+ // Suppress metadata tokens in the suppression set
763
+ for (const suppressedTokenIndex of suppressedMetadataTokens) {
764
+ allTokenLogits[suppressedTokenIndex] = -Infinity
590
765
  }
591
766
 
592
- // Derive token probabilities
593
- let bufferedTokensToPrint: number[] = []
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
- let shouldDecodeNonTimestampToken = true
772
+ const timestampTokenLogits = allTokenLogits.slice(timestampTokensStart)
604
773
 
605
- if (options.decodeTimestampTokens) {
606
- const probabilities = softmax(allTokenLogits as any, 1.0)
607
- const logProbabilities = logOfVector(probabilities)
774
+ const decodeTimestampTokenIfNeeded = () => {
775
+ // Try to decode a timestamp token, if needed
608
776
 
609
- const nonTimestampTokenLogProbs = logProbabilities.slice(0, timestampTokensStart)
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
- const shouldDecodeTimestampToken = logSumExpOfTimestampTokenLogProbs > valueOfMaxNonTimestampLogProb
620
-
621
- const previousTokenWasTimestamp = this.isTimestampToken(decodedTokens[decodedTokens.length - 1])
622
- const secondPreviousTokenWasTimestamp = decodedTokens.length < 2 || this.isTimestampToken(decodedTokens[decodedTokens.length - 2])
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
- if (shouldDecodeTimestampToken || (previousTokenWasTimestamp && !secondPreviousTokenWasTimestamp)) {
629
- if (previousTokenWasTimestamp) {
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
- lastTimestampTokenIndex = decodedTokens.length
637
-
638
- const previousTokenTimestamp = this.timestampTokenToSeconds(previousToken)
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
- if (previousTokenTimestamp >= audioDuration) {
641
- break
642
- }
643
- } else {
644
- const timestampToken = timestampTokensStart + indexOfMaxTimestampLogProb
645
- const confidence = probabilities[timestampToken]
824
+ addToken(previousToken, previousTokenTimestampLogits, previousTokenConfidence, crossAttentionQKsForToken)
646
825
 
647
- addToken(timestampToken, timestampTokenLogits, confidence)
648
- }
826
+ lastTimestampTokenIndex = decodedTokens.length
827
+ } else {
828
+ const timestampToken = timestampTokensStart + indexOfMaxTimestampLogProb
829
+ const confidence = probabilities[timestampToken]
649
830
 
650
- shouldDecodeNonTimestampToken = false
831
+ addToken(timestampToken, timestampTokenLogits, confidence, crossAttentionQKsForToken)
651
832
  }
652
- }
653
-
654
- if (shouldDecodeNonTimestampToken) {
655
- const topLogitCount = options.topCandidateCount!
656
833
 
657
- const nonTimestampTokenLogits = allTokenLogits.slice(0, timestampTokensStart)
834
+ return true
835
+ }
658
836
 
659
- const sortedNonTimestampTokenLogitsWithIndexes =
660
- Array.from(nonTimestampTokenLogits).map((logit, index) => ({ token: index, logit }))
837
+ const timestampTokenDecoded = decodeTimestampTokenIfNeeded()
661
838
 
662
- sortedNonTimestampTokenLogitsWithIndexes.sort((a, b) => b.logit - a.logit)
839
+ if (timestampTokenDecoded) {
840
+ await yieldToEventLoop()
663
841
 
664
- let topCandidates = sortedNonTimestampTokenLogitsWithIndexes.slice(0, topLogitCount)
665
- .map(entry => ({
666
- token: entry.token,
667
- logit: entry.logit,
668
- text: this.tokenToText(entry.token, true)
669
- }))
842
+ continue
843
+ }
670
844
 
671
- //// Repetition suppression code
672
- if (options.suppressRepetition) {
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
- return maxScore
678
- })
848
+ let shouldDecodeEndfOfTextToken = false
679
849
 
680
- const thresholdRepetitionScore = 4
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
- if (topCandidatesRepetitionScores.every(score => score >= thresholdRepetitionScore)) {
683
- const indexOfMaxScore = topCandidatesRepetitionScores.indexOf(Math.max(...topCandidatesRepetitionScores))
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 topCandidateProbabilities = softmax(topCandidates.map(a => a.logit), options.temperature)
858
+ const indexOfMaximumOtherTokenLogit = indexOfMax(otherTokensLogits)
859
+ const maximumOtherTokenLogit = nonTimestampTokenLogits[indexOfMaximumOtherTokenLogit]
692
860
 
693
- //// Remove end-of-text token from candidates if its probability isn't high enough
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
- return true
701
- })
863
+ if (endProbabilities[0] > options.endTokenThreshold!) {
864
+ shouldDecodeEndfOfTextToken = true
702
865
  }
703
- ////
866
+ }
704
867
 
705
- const rankOfPromisingPunctuationToken = topCandidates.findIndex((entry, index) => {
706
- const tokenText = this.tokenToText(entry.token).trim()
868
+ if (logitFilter) {
869
+ // Apply custom logit filter function if given
870
+ nonTimestampTokenLogits = logitFilter(nonTimestampTokenLogits, decodedTokens, isFirstPart, isFinalPart)
707
871
 
708
- const isPunctuationToken = allowedPunctuationMarks.includes(tokenText)
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
- if (!isPunctuationToken) {
711
- return false
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
- const tokenProb = topCandidateProbabilities[index]
891
+ // Suppress space token if at initial state
892
+ if (isInitialState) {
893
+ nonTimestampTokenLogits[spaceToken] = -Infinity
894
+ }
895
+ }
715
896
 
716
- return tokenProb >= options.punctuationThreshold!
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
- let rankOfSpaceToken = topCandidates.findIndex(candidate => candidate.token === spaceToken)
902
+ break
903
+ }
720
904
 
721
- if (rankOfSpaceToken < 0) {
722
- rankOfSpaceToken = Infinity
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
- let chosenCandidateRank: number
910
+ // Find top candidates
911
+ const sortedNonTimestampLogitsWithIndexes =
912
+ Array.from(nonTimestampTokenLogits).map((logit, index) => ({ token: index, logit }))
726
913
 
727
- if (rankOfPromisingPunctuationToken >= 0 &&
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
- const chosenToken = topCandidates[chosenCandidateRank].token
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
- if (this.isTextToken(chosenToken)) {
737
- bufferedTokensToPrint.push(chosenToken)
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
- let textToPrint = this.tokensToText(bufferedTokensToPrint)
930
+ const filteredCandidates: typeof topCandidates = []
740
931
 
741
- if (textToPrint.codePointAt(0) !== 65533) {
742
- if (isFirstPart && decodedTokens.every(token => this.isMetadataToken(token))) {
743
- textToPrint = textToPrint.trimStart()
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
- logger.write(textToPrint)
938
+ const { longestMatch, longestCycleRepetition } = getTokenRepetitionScore([candidate.token, ...lastDecodedTextTokens])
747
939
 
748
- bufferedTokensToPrint = []
940
+ if (longestMatch >= thresholdMatchLength || longestCycleRepetition >= thresholdCycleRepetition) {
941
+ continue
749
942
  }
750
- }
751
-
752
- const confidence = topCandidateProbabilities[chosenCandidateRank]
753
943
 
754
- addToken(chosenToken, timestampTokenLogits, confidence)
944
+ filteredCandidates.push(candidate)
945
+ }
755
946
 
756
- if (chosenToken === endOfTextToken) {
757
- break
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
- await yieldToEventLoop()
762
- }
959
+ // Compute top candidate probabilities
960
+ const topCandidateProbabilities = softmax(topCandidates.map(a => a.logit), options.temperature)
763
961
 
764
- if (timestampsSeenCount >= 2 && !isFinalPart) {
765
- decodedTokens = decodedTokens.slice(0, lastTimestampTokenIndex)
766
- decodedTokensTimestampLogits = decodedTokensTimestampLogits.slice(0, lastTimestampTokenIndex)
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
- logger.write('\n')
771
- logger.end()
966
+ const isPunctuationToken = allowedPunctuationMarks.includes(tokenText)
772
967
 
773
- // Return the tokens
774
- return {
775
- decodedTokens,
776
- decodedTokensTimestampLogits,
777
- crossAttentionQKs: decodedTokensCrossAttentionQKs,
778
- decodedTokensConfidence
779
- }
780
- }
968
+ if (!isPunctuationToken) {
969
+ return false
970
+ }
781
971
 
782
- async inferCrossAttentionQKs(tokens: number[], audioFeatures: Onnx.Tensor) {
783
- const offset = 0
972
+ const tokenProb = topCandidateProbabilities[index]
784
973
 
785
- const Onnx = await import('onnxruntime-node')
974
+ return tokenProb >= options.punctuationThreshold!
975
+ })
786
976
 
787
- const tokensTensor = new Onnx.Tensor('int64', new BigInt64Array(tokens.map(token => BigInt(token))), [1, tokens.length])
788
- const offsetTensor = new Onnx.Tensor('int64', new BigInt64Array([BigInt(offset)]), [])
977
+ // Find rank of space token
978
+ let rankOfSpaceToken = topCandidates.findIndex(candidate => candidate.token === spaceToken)
789
979
 
790
- const initialKvDimensions = this.getKvDimensions(1, tokens.length)
791
- const kvCacheTensor = new Onnx.Tensor('float32', new Float32Array(initialKvDimensions[0] * initialKvDimensions[1] * initialKvDimensions[2] * initialKvDimensions[3]), initialKvDimensions)
980
+ if (rankOfSpaceToken < 0) {
981
+ rankOfSpaceToken = Infinity
982
+ }
792
983
 
793
- const decoderInputs = { tokens: tokensTensor, audio_features: audioFeatures, kv_cache: kvCacheTensor, offset: offsetTensor }
984
+ // Choose token
985
+ let chosenCandidateRank: number
794
986
 
795
- const decoderOutputs = await this.textDecoder!.run(decoderInputs)
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
- const crossAttentionQKsTensor = decoderOutputs['cross_attention_qks']
992
+ chosenCandidateRank = rankOfPromisingPunctuationToken
993
+ } else {
994
+ // Otherwise, select randomly from top k distribution
995
+ chosenCandidateRank = this.randomGen.selectRandomIndexFromDistribution(topCandidateProbabilities)
996
+ }
798
997
 
799
- const tensorShape = crossAttentionQKsTensor.dims.slice()
998
+ // Add chosen token
999
+ const chosenToken = topCandidates[chosenCandidateRank].token
1000
+ const chosenTokenConfidence = topCandidateProbabilities[chosenCandidateRank]
800
1001
 
801
- const ndarray = (await import('ndarray')).default
1002
+ addToken(chosenToken, timestampTokenLogits, chosenTokenConfidence, crossAttentionQKsForToken)
802
1003
 
803
- let qkArray = ndarray(crossAttentionQKsTensor.data, crossAttentionQKsTensor.dims.slice())
804
- qkArray = qkArray.transpose(3, 0, 1, 2, 4)
1004
+ // If chosen token is the end-of-text token, break
1005
+ if (chosenToken === endOfTextToken) {
1006
+ break
1007
+ }
805
1008
 
806
- const tokenCrossAttentionQKsTensors: Onnx.Tensor[] = []
1009
+ // Print token if needed
1010
+ if (this.isTextToken(chosenToken)) {
1011
+ bufferedTokensToPrint.push(chosenToken)
807
1012
 
808
- for (let i0 = 0; i0 < qkArray.shape[0]; i0++) {
809
- const dataForToken: number[] = []
1013
+ let textToPrint = this.tokensToText(bufferedTokensToPrint)
810
1014
 
811
- for (let i1 = 0; i1 < qkArray.shape[1]; i1++) {
812
- for (let i2 = 0; i2 < qkArray.shape[2]; i2++) {
813
- for (let i3 = 0; i3 < qkArray.shape[3]; i3++) {
814
- for (let i4 = 0; i4 < qkArray.shape[4]; i4++) {
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
- const newTensorShape = tensorShape.slice()
822
- newTensorShape[3] = 1
1027
+ await yieldToEventLoop()
1028
+ }
823
1029
 
824
- const newTensor = new Onnx.Tensor('float32', dataForToken, newTensorShape)
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
- tokenCrossAttentionQKsTensors.push(newTensor)
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
- return tokenCrossAttentionQKsTensors
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.subarray(0, maxAudioSamples), 0)
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: TimelineEntry[][] = []
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
- const newGroups: TimelineEntry[][] = []
1167
+ {
1168
+ const splitGroups: Timeline[] = []
1060
1169
 
1061
- for (let groupIndex = 0; groupIndex < groups.length; groupIndex++) {
1062
- const group = groups[groupIndex]
1063
- const nextGroup = groups[groupIndex + 1]
1170
+ for (let groupIndex = 0; groupIndex < groups.length; groupIndex++) {
1171
+ const group = groups[groupIndex]
1172
+ const nextGroup = groups[groupIndex + 1]
1064
1173
 
1065
- if (
1066
- group.length > 1 &&
1067
- group[group.length - 1].text === '.' &&
1068
- (!nextGroup || [' ', '['].includes(nextGroup[0].text[0]))) {
1174
+ if (
1175
+ group.length > 1 &&
1176
+ group[group.length - 1].text === '.' &&
1177
+ (!nextGroup || [' ', '['].includes(nextGroup[0].text[0]))) {
1069
1178
 
1070
- newGroups.push(group.slice(0, group.length - 1))
1071
- newGroups.push(group.slice(group.length - 1))
1072
- } else {
1073
- newGroups.push(group)
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
- groups = newGroups
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: Onnx.Tensor[], tokens: number[], segmentStartFrame: number, segmentEndFrame: number, headIndexes?: number[]) {
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 fixateTimestampTokens = false
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 meanOfAllWeights = meanOfVector(allWeightsForHead)
1216
- const stdDeviationOfAllWeights = stdDeviationOfVector(allWeightsForHead) + 1e-10
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] - meanOfAllWeights) / stdDeviationOfAllWeights
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] = medianFilter(head[tokenIndex], medianFilterWidth)
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 (fixateTimestampTokens) {
1259
- // Fixate timestamp tokens to the original ones detected
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 = [...Array(tokenCount).keys()]
1277
- const frameIndexes = [...Array(segmentFrameCount).keys()]
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.getNonWordTokenData()
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
- getNonWordTokenData() {
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(tokenTextWithoutWhitespace)
1697
+ const isNonWordToken = /^[\s\p{Punctuation}\p{Symbol}]+$/u.test(tokenText)
1481
1698
 
1482
- const containsInvalidUTF8 = getUTF32Chars(tokenTextWithoutWhitespace).utf32chars.includes(invalidUTF8Char)
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
+ }