echogarden 0.11.12 → 0.11.13

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/data/schemas/options.json +16 -0
  2. package/dist/api/Alignment.js +2 -2
  3. package/dist/api/Alignment.js.map +1 -1
  4. package/dist/api/Recognition.js +2 -2
  5. package/dist/api/Recognition.js.map +1 -1
  6. package/dist/api/Synthesis.js +5 -4
  7. package/dist/api/Synthesis.js.map +1 -1
  8. package/dist/api/Translation.js +2 -2
  9. package/dist/api/Translation.js.map +1 -1
  10. package/dist/audio/AudioUtilities.d.ts +1 -0
  11. package/dist/audio/AudioUtilities.js +25 -7
  12. package/dist/audio/AudioUtilities.js.map +1 -1
  13. package/dist/cli/CLI.js +2 -2
  14. package/dist/cli/CLI.js.map +1 -1
  15. package/dist/recognition/WhisperSTT.js +2 -2
  16. package/dist/recognition/WhisperSTT.js.map +1 -1
  17. package/dist/subtitles/Subtitles.d.ts +10 -7
  18. package/dist/subtitles/Subtitles.js +268 -207
  19. package/dist/subtitles/Subtitles.js.map +1 -1
  20. package/docs/Options.md +4 -2
  21. package/package.json +7 -6
  22. package/src/alignment/DTWMfccSequenceAlignment.ts +43 -0
  23. package/src/alignment/DTWSequenceAlignment.ts +121 -0
  24. package/src/alignment/DTWSequenceAlignmentWindowed.ts +210 -0
  25. package/src/alignment/LevenshteinSequenceAlignment.ts +126 -0
  26. package/src/alignment/SpeechAlignment.ts +488 -0
  27. package/src/api/API.ts +12 -0
  28. package/src/api/APIOptions.ts +15 -0
  29. package/src/api/Alignment.ts +329 -0
  30. package/src/api/Common.ts +16 -0
  31. package/src/api/Denoising.ts +120 -0
  32. package/src/api/LanguageDetection.ts +286 -0
  33. package/src/api/Recognition.ts +344 -0
  34. package/src/api/Synthesis.ts +1735 -0
  35. package/src/api/Translation.ts +143 -0
  36. package/src/api/Vad.ts +172 -0
  37. package/src/audio/AudioBufferConversion.ts +248 -0
  38. package/src/audio/AudioPlayer.ts +358 -0
  39. package/src/audio/AudioRecorder.ts +91 -0
  40. package/src/audio/AudioUtilities.ts +392 -0
  41. package/src/audio/SoxPath.ts +24 -0
  42. package/src/cli/CLI.ts +1360 -0
  43. package/src/cli/CLIConfigFile.ts +91 -0
  44. package/src/cli/CLILauncher.ts +26 -0
  45. package/src/cli/CLIOptionsSchema.ts +54 -0
  46. package/src/cli/CLIParser.ts +41 -0
  47. package/src/cli/CLIStarter.ts +40 -0
  48. package/src/codecs/FFMpegTranscoder.ts +214 -0
  49. package/src/codecs/TIMITCodec.ts +17 -0
  50. package/src/codecs/WaveCodec.ts +260 -0
  51. package/src/denoising/RNNoise.ts +95 -0
  52. package/src/dsp/BiquadFilter.ts +488 -0
  53. package/src/dsp/FFT.ts +187 -0
  54. package/src/dsp/MFCC.ts +227 -0
  55. package/src/dsp/MelSpectogram.ts +145 -0
  56. package/src/dsp/Rubberband.ts +249 -0
  57. package/src/dsp/Sonic.ts +59 -0
  58. package/src/dsp/SpeexResampler.ts +79 -0
  59. package/src/math/VectorMath.ts +812 -0
  60. package/src/nlp/ChineseSegmentation.ts +68 -0
  61. package/src/nlp/CompromiseNLP.ts +113 -0
  62. package/src/nlp/EspeakPhonemizer.ts +168 -0
  63. package/src/nlp/IPA.ts +139 -0
  64. package/src/nlp/JapaneseSegmentation.ts +53 -0
  65. package/src/nlp/Lexicon.ts +119 -0
  66. package/src/nlp/PhoneConversion.ts +508 -0
  67. package/src/nlp/Segmentation.ts +237 -0
  68. package/src/nlp/TextNormalizer.ts +160 -0
  69. package/src/recognition/AmazonTranscribeSTT.ts +112 -0
  70. package/src/recognition/AzureCognitiveServicesSTT.ts +76 -0
  71. package/src/recognition/GoogleCloudSTT.ts +92 -0
  72. package/src/recognition/SileroSTT.ts +173 -0
  73. package/src/recognition/VoskSTT.ts +112 -0
  74. package/src/recognition/WhisperSTT.ts +1518 -0
  75. package/src/server/Client.ts +297 -0
  76. package/src/server/Server.ts +178 -0
  77. package/src/server/ServerStarter.ts +12 -0
  78. package/src/server/Worker.ts +400 -0
  79. package/src/server/WorkerStarter.ts +38 -0
  80. package/src/speech-language-detection/SileroLanguageDetection.ts +105 -0
  81. package/src/subtitles/Subtitles.ts +478 -0
  82. package/src/synthesis/AwsPollyTTS.ts +78 -0
  83. package/src/synthesis/AzureCognitiveServicesTTS.ts +146 -0
  84. package/src/synthesis/CoquiServerTTS.ts +29 -0
  85. package/src/synthesis/ElevenLabsTTS.ts +104 -0
  86. package/src/synthesis/EspeakTTS.ts +552 -0
  87. package/src/synthesis/FliteTTS.ts +387 -0
  88. package/src/synthesis/GoogleCloudTTS.ts +112 -0
  89. package/src/synthesis/GoogleTranslateTTS.ts +210 -0
  90. package/src/synthesis/MicrosoftEdgeTTS.ts +298 -0
  91. package/src/synthesis/SamTTS.ts +30 -0
  92. package/src/synthesis/SapiTTS.ts +222 -0
  93. package/src/synthesis/StreamlabsPollyTTS.ts +114 -0
  94. package/src/synthesis/SvoxPicoTTS.ts +318 -0
  95. package/src/synthesis/VitsTTS.ts +734 -0
  96. package/src/tests/Test.ts +24 -0
  97. package/src/text-language-detection/FastTextLanguageDetection.ts +53 -0
  98. package/src/text-language-detection/TinyLDLanguageDetection.ts +16 -0
  99. package/src/typings/Fillers.d.ts +41 -0
  100. package/src/utilities/BinaryArrayConversion.ts +159 -0
  101. package/src/utilities/Compression.ts +91 -0
  102. package/src/utilities/FileDownloader.ts +201 -0
  103. package/src/utilities/FileSystem.ts +265 -0
  104. package/src/utilities/Hashing.ts +230 -0
  105. package/src/utilities/Locale.ts +119 -0
  106. package/src/utilities/Logger.ts +72 -0
  107. package/src/utilities/NdArrayUtilities.ts +31 -0
  108. package/src/utilities/ObjectUtilities.ts +169 -0
  109. package/src/utilities/OpenPromise.ts +13 -0
  110. package/src/utilities/PackageManager.ts +97 -0
  111. package/src/utilities/Queue.ts +17 -0
  112. package/src/utilities/RandomGenerator.ts +237 -0
  113. package/src/utilities/SignalChannel.ts +22 -0
  114. package/src/utilities/TarballMaker.ts +68 -0
  115. package/src/utilities/Timeline.ts +231 -0
  116. package/src/utilities/Timer.ts +93 -0
  117. package/src/utilities/Utilities.ts +574 -0
  118. package/src/utilities/WasmMemoryManager.ts +516 -0
  119. package/src/utilities/WebReader.ts +55 -0
  120. package/src/utilities/WikipediaReader.ts +41 -0
  121. package/src/voice-activity-detection/SileroVAD.ts +86 -0
  122. package/src/voice-activity-detection/WebRtcVAD.ts +76 -0
@@ -0,0 +1,1518 @@
1
+ import Onnx from 'onnxruntime-node'
2
+
3
+ import { Logger } from '../utilities/Logger.js'
4
+ import { computeMelSpectogramUsingFilterbanks, Filterbank } from "../dsp/MelSpectogram.js"
5
+ import { clip, splitFloat32Array, writeToStderr, yieldToEventLoop } from '../utilities/Utilities.js'
6
+ import { indexOfMax, logOfVector, logSumExp, meanOfVector, medianFilter, softmax, stdDeviationOfVector } from '../math/VectorMath.js'
7
+ import { isWordOrSymbolWord, splitToWords } from '../nlp/Segmentation.js'
8
+
9
+ import { alignDTWWindowed } from '../alignment/DTWSequenceAlignmentWindowed.js'
10
+ import { deepClone, extendDeep } from '../utilities/ObjectUtilities.js'
11
+ import { Timeline, TimelineEntry } from '../utilities/Timeline.js'
12
+ import { AlignmentPath } from '../alignment/SpeechAlignment.js'
13
+ import { getRawAudioDuration, RawAudio } from '../audio/AudioUtilities.js'
14
+ import { readAndParseJsonFile, readFile } from '../utilities/FileSystem.js'
15
+ import path from 'path'
16
+ import type { LanguageDetectionResults } from '../api/API.js'
17
+ import { getShortLanguageCode, languageCodeToName } from '../utilities/Locale.js'
18
+ import { loadPackage } from '../utilities/PackageManager.js'
19
+ import chalk from 'chalk'
20
+ import { XorShift32RNG } from '../utilities/RandomGenerator.js'
21
+ import { detectSpeechLanguageByParts } from '../api/LanguageDetection.js'
22
+
23
+ export async function recognize(sourceRawAudio: RawAudio, modelName: WhisperModelName, modelDir: string, tokenizerDir: string, task: WhisperTask, sourceLanguage: string, options: WhisperOptions) {
24
+ if (sourceRawAudio.sampleRate != 16000) {
25
+ throw new Error("Source audio must have a sampling rate of 16000")
26
+ }
27
+
28
+ sourceLanguage = getShortLanguageCode(sourceLanguage)
29
+
30
+ if (!(sourceLanguage in languageIdLookup)) {
31
+ throw new Error(`The language ${languageCodeToName(sourceLanguage)} is not supported by the Whisper engine.`)
32
+ }
33
+
34
+ if (!isMultiligualModel(modelName) && sourceLanguage != 'en') {
35
+ throw new Error(`The model '${modelName}' can only be used with English inputs. However, the given source language was ${languageCodeToName(sourceLanguage)}.`)
36
+ }
37
+
38
+ const whisper = new Whisper(modelName, modelDir, tokenizerDir)
39
+ await whisper.initialize()
40
+
41
+ const result = await whisper.recognize(sourceRawAudio, task, sourceLanguage, options)
42
+
43
+ return result
44
+ }
45
+
46
+ export async function align(sourceRawAudio: RawAudio, referenceText: string, modelName: WhisperModelName, modelDir: string, tokenizerDir: string, sourceLanguage: string) {
47
+ if (sourceRawAudio.sampleRate != 16000) {
48
+ throw new Error("Source audio must have a sampling rate of 16000")
49
+ }
50
+
51
+ sourceLanguage = getShortLanguageCode(sourceLanguage)
52
+
53
+ if (!(sourceLanguage in languageIdLookup)) {
54
+ throw new Error(`The language ${languageCodeToName(sourceLanguage)} is not supported by the Whisper engine.`)
55
+ }
56
+
57
+ if (!isMultiligualModel(modelName) && sourceLanguage != 'en') {
58
+ throw new Error(`The model '${modelName}' can only be used with English inputs. However, the given source language was ${languageCodeToName(sourceLanguage)}.`)
59
+ }
60
+
61
+ const whisper = new Whisper(modelName, modelDir, tokenizerDir)
62
+ await whisper.initialize()
63
+
64
+ const timeline = await whisper.align(sourceRawAudio, referenceText, sourceLanguage)
65
+
66
+ return timeline
67
+ }
68
+
69
+ export async function detectLanguage(sourceRawAudio: RawAudio, modelName: WhisperModelName, modelDir: string, tokenizerDir: string) {
70
+ if (sourceRawAudio.sampleRate != 16000) {
71
+ throw new Error("Source audio must have a sampling rate of 16000")
72
+ }
73
+
74
+ const whisper = new Whisper(modelName, modelDir, tokenizerDir)
75
+ await whisper.initialize()
76
+
77
+ async function detectLanguageForPart(partAudio: RawAudio) {
78
+ const audioFeatures = await whisper.encodeAudio(partAudio)
79
+ const partResults = await whisper.detectLanguage(audioFeatures)
80
+
81
+ return partResults
82
+ }
83
+
84
+ const results = await detectSpeechLanguageByParts(sourceRawAudio, detectLanguageForPart)
85
+
86
+ results.sort((entry1, entry2) => entry2.probability - entry1.probability)
87
+
88
+ return results
89
+ }
90
+
91
+ export class Whisper {
92
+ modelName: WhisperModelName
93
+ modelDir: string
94
+ tokenizerDir: string
95
+
96
+ isMultiligualModel: boolean
97
+
98
+ audioEncoder?: Onnx.InferenceSession
99
+ textDecoder?: Onnx.InferenceSession
100
+
101
+ textToTokenLookup = new Map<string, number>()
102
+ tokenToTextLookup = new Map<number, string>()
103
+
104
+ merges: [string, string][] = []
105
+
106
+ onnxOptions: Onnx.InferenceSession.SessionOptions = {
107
+ logSeverityLevel: 2,
108
+ executionProviders: ['cpu']
109
+ }
110
+
111
+ tokenConfig: {
112
+ suppressedTokens: number[]
113
+ sotToken: number
114
+ sotPrevToken: number
115
+ eotToken: number
116
+ noTimestampsToken: number
117
+ noSpeechToken: number
118
+ timestampTokensStart: number
119
+ }
120
+
121
+ randomGen = new XorShift32RNG(23948203)
122
+
123
+ constructor(modelName: WhisperModelName, modelDir: string, tokenizerDir: string) {
124
+ this.modelDir = modelDir
125
+ this.modelName = modelName
126
+ this.tokenizerDir = tokenizerDir
127
+
128
+ this.isMultiligualModel = isMultiligualModel(this.modelName)
129
+
130
+ if (this.isMultiligualModel) {
131
+ this.tokenConfig = {
132
+ sotToken: 50258,
133
+ sotPrevToken: 50361,
134
+ eotToken: 50257,
135
+ noSpeechToken: 50362,
136
+ noTimestampsToken: 50363,
137
+ timestampTokensStart: 50364,
138
+ suppressedTokens: [1, 2, 6, 7, 8, 9, 10, 12, 14, 25, 26, 27, 28, 29, 31, 58, 59, 60, 61, 62, 63, 90, 91, 92, 93, 359, 503, 522, 542, 873, 893, 902, 918, 922, 931, 1350, 1853, 1982, 2460, 2627, 3246, 3253, 3268, 3536, 3846, 3961, 4183, 4667, 6585, 6647, 7273, 9061, 9383, 10428, 10929, 11938, 12033, 12331, 12562, 13793, 14157, 14635, 15265, 15618, 16553, 16604, 18362, 18956, 20075, 21675, 22520, 26130, 26161, 26435, 28279, 29464, 31650, 32302, 32470, 36865, 42863, 47425, 49870, 50254, 50258, 50360, 50361, 50362]
139
+ }
140
+ } else {
141
+ this.tokenConfig = {
142
+ sotToken: 50257,
143
+ sotPrevToken: 50360,
144
+ eotToken: 50256,
145
+ noSpeechToken: 50361,
146
+ noTimestampsToken: 50362,
147
+ timestampTokensStart: 50363,
148
+ suppressedTokens: [1, 2, 7, 8, 9, 10, 14, 25, 26, 27, 28, 29, 31, 58, 59, 60, 61, 62, 63, 90, 91, 92, 93, 357, 366, 438, 532, 685, 705, 796, 930, 1058, 1220, 1267, 1279, 1303, 1343, 1377, 1391, 1635, 1782, 1875, 2162, 2361, 2488, 3467, 4008, 4211, 4600, 4808, 5299, 5855, 6329, 7203, 9609, 9959, 10563, 10786, 11420, 11709, 11907, 13163, 13697, 13700, 14808, 15306, 16410, 16791, 17992, 19203, 19510, 20724, 22305, 22935, 27007, 30109, 30420, 33409, 34949, 40283, 40493, 40549, 47282, 49146, 50257, 50359, 50360, 50361]
149
+ }
150
+ }
151
+ }
152
+
153
+ async initialize() {
154
+ const logger = new Logger()
155
+ await logger.startAsync("Load tokenizer data")
156
+
157
+ const encoderFilePath = path.join(this.modelDir, "encoder.onnx")
158
+ const decoderFilePath = path.join(this.modelDir, "decoder.onnx")
159
+
160
+ const vocabFilePath = path.join(this.tokenizerDir, "vocab.json")
161
+ const mergesFilePath = path.join(this.tokenizerDir, "merges.txt")
162
+
163
+ const vocabObject = await readAndParseJsonFile(vocabFilePath)
164
+
165
+ function bpeEncodedStrToString(str: string) {
166
+ const decodedChars = []
167
+
168
+ for (const char of str) {
169
+ const decodedChar = vocabCharacterSetLookup[char]
170
+
171
+ if (decodedChar == undefined) {
172
+ throw new Error(`Invalid char: '${char}'`)
173
+ }
174
+
175
+ decodedChars.push(decodedChar)
176
+ }
177
+
178
+ return Buffer.from(decodedChars).toString("utf-8")
179
+ }
180
+
181
+ for (const key in vocabObject) {
182
+ const value = vocabObject[key]
183
+
184
+ const decodedKey = bpeEncodedStrToString(key)
185
+
186
+ this.textToTokenLookup.set(decodedKey, value)
187
+ this.tokenToTextLookup.set(value, decodedKey)
188
+ }
189
+
190
+ const mergesFileRawLines = (await readFile(mergesFilePath, "utf8")).trim().split(/\r?\n/g)
191
+ const mergesFileRawEntries = mergesFileRawLines.map(line => line.trim().split(" "))
192
+ this.merges = mergesFileRawEntries.map(entry => [bpeEncodedStrToString(entry[0]), bpeEncodedStrToString(entry[1])])
193
+
194
+ await logger.startAsync(`Create ONNX inference session for model '${this.modelName}'`)
195
+
196
+ this.audioEncoder = await Onnx.InferenceSession.create(encoderFilePath, this.onnxOptions)
197
+ this.textDecoder = await Onnx.InferenceSession.create(decoderFilePath, this.onnxOptions)
198
+
199
+ logger.end()
200
+ }
201
+
202
+ async recognize(rawAudio: RawAudio, task: WhisperTask, language: string, options: WhisperOptions) {
203
+ const logger = new Logger()
204
+
205
+ const timestampTokensStart = this.tokenConfig.timestampTokensStart
206
+
207
+ const audioSamples = rawAudio.audioChannels[0]
208
+ const sampleRate = rawAudio.sampleRate
209
+ const audioDuration = getRawAudioDuration(rawAudio)
210
+ const prompt = options.prompt
211
+
212
+ const maxAudioSamples = sampleRate * 30
213
+
214
+ let previousPartTokens: number[] = []
215
+
216
+ let timeline: Timeline = []
217
+ let allDecodedTokens: number[] = []
218
+
219
+ for (let audioOffset = 0; audioOffset < audioSamples.length;) {
220
+ const segmentStartTime = audioOffset / sampleRate
221
+
222
+ await logger.startAsync(`\nPrepare audio part at time position ${segmentStartTime.toFixed(2)}`, undefined, chalk.magentaBright)
223
+
224
+ const audioPartSamples = audioSamples.slice(audioOffset, audioOffset + maxAudioSamples)
225
+ const audioPartRawAudio: RawAudio = { audioChannels: [audioPartSamples], sampleRate }
226
+ const audioPartDuration = getRawAudioDuration(audioPartRawAudio)
227
+
228
+ logger.end()
229
+
230
+ const audioPartFeatures = await this.encodeAudio(audioPartRawAudio)
231
+
232
+ const isFirstPart = audioOffset == 0
233
+ const isFinalPart = audioOffset + maxAudioSamples > audioSamples.length
234
+
235
+ let initialTokens: number[] = []
236
+
237
+ if (isFirstPart && prompt) {
238
+ const promptTokens = await this.textToTokens(prompt, language)
239
+
240
+ initialTokens = [this.tokenConfig.sotPrevToken, ...promptTokens]
241
+ } else if (options.autoPromptParts && previousPartTokens.length > 0) {
242
+ initialTokens = [this.tokenConfig.sotPrevToken, ...previousPartTokens]
243
+ }
244
+
245
+ initialTokens = [...initialTokens, ...this.getInitialTokens(language, task)]
246
+
247
+ logger.end()
248
+
249
+ let { decodedTokens: partTokens, crossAttentionQKs: partCrossAttentionQKs, decodedTokensConfidence } = await this.decodeTokens(audioPartFeatures, initialTokens, audioPartDuration, isFirstPart, isFinalPart, options)
250
+
251
+ const lastToken = partTokens[partTokens.length - 1]
252
+ const lastTokenIsTimestamp = lastToken >= timestampTokensStart
253
+
254
+ let audioEndOffset: number
255
+
256
+ if (!isFinalPart && lastTokenIsTimestamp) {
257
+ const timePosition = (lastToken - timestampTokensStart) * 0.02
258
+
259
+ audioEndOffset = audioOffset + Math.floor(timePosition * sampleRate)
260
+ } else {
261
+ audioEndOffset = Math.min(audioOffset + maxAudioSamples, audioSamples.length)
262
+ }
263
+
264
+ const segmentEndTime = audioEndOffset / sampleRate
265
+ const segmentFrameCount = Math.floor((segmentEndTime - segmentStartTime) / 0.02)
266
+
267
+ await logger.startAsync(`Extract timeline for part`)
268
+
269
+ if (partTokens.length != partCrossAttentionQKs.length) {
270
+ throw new Error("Unexpected: partTokens.length != partCrossAttentionQKs.length")
271
+ }
272
+
273
+ //partTokens = partTokens.filter(token => token < timestampTokensStart)
274
+ //partCrossAttentionQKs = await this.inferCrossAttentionQKs(partTokens, audioPartFeatures)
275
+
276
+ partTokens = partTokens.slice(initialTokens.length)
277
+ partCrossAttentionQKs = partCrossAttentionQKs.slice(initialTokens.length)
278
+
279
+ //await this.addWordsToTimeline(timeline, partTokens, audioPartRawAudio, partCrossAttentionQKs, initialAudioTimeOffset, audioPartSamples.length / sampleRate)
280
+
281
+ const alignmentPath = await this.findAlignmentPathFromQKs(partCrossAttentionQKs, partTokens, 0, segmentFrameCount) //, alignmentHeadsIndexes[this.modelName])
282
+ const partTimeline = await this.getWordTimelineFromAlignmentPath(alignmentPath, partTokens, segmentStartTime, segmentEndTime, decodedTokensConfidence)
283
+
284
+ timeline.push(...partTimeline)
285
+
286
+ audioOffset = audioEndOffset
287
+
288
+ previousPartTokens = partTokens.filter(token => token < this.tokenConfig.eotToken)
289
+
290
+ allDecodedTokens.push(...previousPartTokens)
291
+
292
+ logger.end()
293
+ }
294
+
295
+ if (timeline.length > 0) {
296
+ timeline[timeline.length - 1].endTime = audioDuration
297
+ }
298
+
299
+ timeline = this.mergeSuccessiveWordFragmentsInTimeline(timeline)
300
+ timeline.forEach(entry => { entry.text = entry.text.trim() })
301
+
302
+ const transcript = this.tokensToText(allDecodedTokens)
303
+
304
+ logger.end()
305
+
306
+ return { transcript, timeline }
307
+ }
308
+
309
+ async align(rawAudio: RawAudio, referenceText: string, language: string) {
310
+ const logger = new Logger()
311
+
312
+ await logger.startAsync("Prepare for alignment")
313
+ const audioDuration = Math.min(getRawAudioDuration(rawAudio), 30)
314
+ const audioFrameCount = Math.floor(audioDuration / 0.02)
315
+
316
+ const initialTokens = this.getInitialTokens(language, "transcribe", true)
317
+ const timestampTokensStart = this.tokenConfig.timestampTokensStart
318
+ const eotToken = this.tokenConfig.eotToken
319
+
320
+ let tokens = [...initialTokens, ...await this.textToTokens(referenceText, language), eotToken]
321
+
322
+ logger.end()
323
+ const audioFeatures = await this.encodeAudio(rawAudio)
324
+
325
+ await logger.startAsync("Infer cross-attention QKs")
326
+ let crossAttentionQKs = await this.inferCrossAttentionQKs(tokens, audioFeatures)
327
+
328
+ tokens = tokens.slice(initialTokens.length, tokens.length - 1)
329
+ crossAttentionQKs = crossAttentionQKs.slice(initialTokens.length, crossAttentionQKs.length - 1)
330
+
331
+ await logger.startAsync("Extract word timeline")
332
+ const alignmentPath = await this.findAlignmentPathFromQKs(crossAttentionQKs, tokens, 0, audioFrameCount)//, this.getAlignmentHeadIndexes())
333
+ let timeline = await this.getWordTimelineFromAlignmentPath(alignmentPath, tokens, 0, audioDuration)
334
+
335
+ timeline = this.mergeSuccessiveWordFragmentsInTimeline(timeline)
336
+ timeline.forEach(entry => { entry.text = entry.text.trim() })
337
+ //timeline = timeline.filter(entry => isWordOrSymbolWord(entry.text))
338
+
339
+ logger.end()
340
+
341
+ return timeline
342
+ }
343
+
344
+ async detectLanguage(audioFeatures: Onnx.Tensor): Promise<LanguageDetectionResults> {
345
+ const logger = new Logger()
346
+
347
+ if (!this.isMultiligualModel) {
348
+ throw new Error("Language detection only works for a multilingual model")
349
+ }
350
+
351
+ // Prepare and run decoder
352
+ await logger.startAsync("Detect language with Whisper model")
353
+
354
+ const sotToken = this.tokenConfig.sotToken
355
+
356
+ const initialTokens = [sotToken]
357
+ const offset = 0
358
+
359
+ const initialKvDimensions = this.getKvDimensions(1, initialTokens.length)
360
+ const kvCacheTensor = new Onnx.Tensor('float32', new Float32Array(initialKvDimensions[0] * initialKvDimensions[1] * initialKvDimensions[2] * initialKvDimensions[3]), initialKvDimensions)
361
+
362
+ const tokensTensor = new Onnx.Tensor('int64', new BigInt64Array(initialTokens.map(token => BigInt(token))), [1, initialTokens.length])
363
+ const offsetTensor = new Onnx.Tensor('int64', new BigInt64Array([BigInt(offset)]), [])
364
+
365
+ const decoderInputs = { tokens: tokensTensor, audio_features: audioFeatures, kv_cache: kvCacheTensor, offset: offsetTensor }
366
+
367
+ const decoderOutputs = await this.textDecoder!.run(decoderInputs)
368
+ const logitsBuffer = decoderOutputs["logits"].data as Float32Array
369
+
370
+ const languageTokensLogits = Array.from(logitsBuffer.slice(sotToken + 1, sotToken + 1 + 99))
371
+ const languageTokensProbabilities = softmax(languageTokensLogits, 1.0)
372
+
373
+ const results: LanguageDetectionResults = []
374
+
375
+ for (const language in languageIdLookup) {
376
+ const langId = languageIdLookup[language]
377
+ const probability = languageTokensProbabilities[langId]
378
+
379
+ results.push({
380
+ language,
381
+ languageName: languageCodeToName(language),
382
+ probability
383
+ })
384
+ }
385
+
386
+ logger.end()
387
+
388
+ return results
389
+ }
390
+
391
+ async decodeTokens(audioFeatures: Onnx.Tensor, initialTokens: number[], audioDuration: number, isFirstPart: boolean, isFinalPart: boolean, options: WhisperOptions) {
392
+ const logger = new Logger()
393
+ await logger.startAsync("Decode text tokens with Whisper decoder model")
394
+
395
+ options = extendDeep(whisperOptionsDefaults, options)
396
+
397
+ const noSpeechThreshold = 0.6
398
+
399
+ const blankToken = this.textToTokenLookup.get(" ")
400
+
401
+ const suppressedTokens = this.tokenConfig.suppressedTokens
402
+ const sotToken = this.tokenConfig.sotToken
403
+ const eotToken = this.tokenConfig.eotToken
404
+ const noTimestampsToken = this.tokenConfig.noTimestampsToken
405
+ const noSpeechToken = this.tokenConfig.noSpeechToken
406
+ const timestampTokensStart = this.tokenConfig.timestampTokensStart
407
+
408
+ const maxDecodedTokenCount = 250
409
+
410
+ let decodedTokens = initialTokens.slice()
411
+ const initialKvDimensions = this.getKvDimensions(1, decodedTokens.length)
412
+ let kvCacheTensor = new Onnx.Tensor('float32', new Float32Array(initialKvDimensions[0] * initialKvDimensions[1] * initialKvDimensions[2] * initialKvDimensions[3]), initialKvDimensions)
413
+
414
+ let decodedTokensTimestampLogits: number[][] = [new Array(1501)]
415
+
416
+ let lastTimestampTokenIndex = -1
417
+
418
+ let timestampsSeenCount = 0
419
+
420
+ const decodedTokensConfidence: number[] = []
421
+ let decodedTokensCrossAttentionQKs: Onnx.Tensor[] = []
422
+
423
+ for (let i = 0; i < decodedTokens.length; i++) {
424
+ decodedTokensCrossAttentionQKs.push(undefined as any)
425
+ }
426
+
427
+ // Start decoding loop
428
+ for (let decodedTokenCount = 0; decodedTokenCount < maxDecodedTokenCount; decodedTokenCount++) {
429
+ const isInitialState = decodedTokens.length == initialTokens.length
430
+
431
+ const tokensToDecode = isInitialState ? decodedTokens : [decodedTokens[decodedTokens.length - 1]]
432
+ const offset = isInitialState ? 0 : decodedTokens.length
433
+
434
+ if (!isInitialState) {
435
+ // Reshape KV Cache tensor
436
+ const dims = kvCacheTensor.dims
437
+
438
+ const currentKvCacheGroups = splitFloat32Array(kvCacheTensor.data as Float32Array, dims[2] * dims[3])
439
+
440
+ const reshapedKvCacheTensor = new Onnx.Tensor('float32', new Float32Array(dims[0] * dims[1] * (decodedTokens.length) * dims[3]), [dims[0], dims[1], decodedTokens.length, dims[3]])
441
+ const reshapedKvCacheGroups = splitFloat32Array(reshapedKvCacheTensor.data, decodedTokens.length * dims[3])
442
+
443
+ for (let i = 0; i < dims[0]; i++) {
444
+ reshapedKvCacheGroups[i].set(currentKvCacheGroups[i])
445
+ }
446
+
447
+ kvCacheTensor = reshapedKvCacheTensor
448
+ }
449
+
450
+ // Prepare and run decoder
451
+ const tokensTensor = new Onnx.Tensor('int64', new BigInt64Array(tokensToDecode.map(token => BigInt(token))), [1, tokensToDecode.length])
452
+ const offsetTensor = new Onnx.Tensor('int64', new BigInt64Array([BigInt(offset)]), [])
453
+
454
+ const decoderInputs = { tokens: tokensTensor, audio_features: audioFeatures, kv_cache: kvCacheTensor, offset: offsetTensor }
455
+
456
+ const decoderOutputs = await this.textDecoder!.run(decoderInputs)
457
+
458
+ const logitsBuffer = decoderOutputs["logits"].data as Float32Array
459
+ kvCacheTensor = decoderOutputs["output_kv_cache"] as any
460
+
461
+ // Compute logits
462
+ const resultLogits = splitFloat32Array(logitsBuffer, logitsBuffer.length / decoderOutputs["logits"].dims[1])
463
+ const tokenLogits = resultLogits[resultLogits.length - 1]
464
+ const tokenTimestampLogits = Array.from(tokenLogits.slice(timestampTokensStart))
465
+
466
+ // Suppress tokens
467
+ for (let logitIndex = 0; logitIndex < tokenLogits.length; logitIndex++) {
468
+ const isWrongTokenForInitialState = isInitialState && (logitIndex == blankToken || logitIndex == eotToken)
469
+ const isInSupressedList = suppressedTokens.includes(logitIndex)
470
+ const isNoTimestampsToken = logitIndex == noTimestampsToken
471
+
472
+ const shouldSupressToken = isWrongTokenForInitialState || isInSupressedList || isNoTimestampsToken
473
+
474
+ if (shouldSupressToken) {
475
+ tokenLogits[logitIndex] = -Infinity
476
+ }
477
+ }
478
+
479
+ // Find token distributions and best token
480
+ const probs = softmax(tokenLogits as any)
481
+ const logProbs = logOfVector(probs)
482
+
483
+ const textTokenLogProbs = logProbs.slice(0, timestampTokensStart)
484
+ const timestampTokenLogProbs = logProbs.slice(timestampTokensStart)
485
+
486
+ const indexOfMaxTextLogProb = indexOfMax(textTokenLogProbs)
487
+ const valueOfMaxTextLogProb = textTokenLogProbs[indexOfMaxTextLogProb]
488
+
489
+ const indexOfMaxTimestampLogProb = indexOfMax(timestampTokenLogProbs)
490
+
491
+ const logSumExpOfTimestampTokenLogProbs = logSumExp(timestampTokenLogProbs)
492
+
493
+ const isTimestampToken = logSumExpOfTimestampTokenLogProbs > valueOfMaxTextLogProb
494
+ const previousTokenWasTimestamp = decodedTokens[decodedTokens.length - 1] >= timestampTokensStart
495
+ const secondPreviousTokenWasTimestamp = decodedTokens.length < 2 || decodedTokens[decodedTokens.length - 2] >= timestampTokensStart
496
+
497
+ if (isTimestampToken && !previousTokenWasTimestamp) {
498
+ timestampsSeenCount += 1
499
+ }
500
+
501
+ //
502
+ //const topLogits = [...tokenLogits].map((logit, index) => ({ index, logit, token: this.tokenToTextLookup.get(index) || "", prob: probs[index] }))
503
+ //topLogits.sort((a, b) => b.logit - a.logit)
504
+ ///
505
+
506
+ // Add best token
507
+ function addToken(tokenToAdd: number, timestampLogits: number[], confidence: number) {
508
+ decodedTokens.push(tokenToAdd)
509
+ decodedTokensTimestampLogits.push(timestampLogits)
510
+ decodedTokensCrossAttentionQKs.push(decoderOutputs["cross_attention_qks"])
511
+ decodedTokensConfidence.push(confidence)
512
+ }
513
+
514
+ if (isTimestampToken || (previousTokenWasTimestamp && !secondPreviousTokenWasTimestamp)) {
515
+ if (previousTokenWasTimestamp) {
516
+ const previousToken = decodedTokens[decodedTokens.length - 1]
517
+ const previousTokenTimestampLogits = decodedTokensTimestampLogits[decodedTokensTimestampLogits.length - 1]
518
+ const previousTokenConfidence = decodedTokensConfidence[decodedTokensConfidence.length - 1]
519
+
520
+ addToken(previousToken, previousTokenTimestampLogits, previousTokenConfidence)
521
+
522
+ lastTimestampTokenIndex = decodedTokens.length
523
+
524
+ const previousTokenTimestamp = (previousToken - timestampTokensStart) * 0.02
525
+
526
+ if (previousTokenTimestamp >= audioDuration) {
527
+ break
528
+ }
529
+ } else {
530
+ const timestampToken = timestampTokensStart + indexOfMaxTimestampLogProb
531
+ const confidence = probs[timestampToken]
532
+
533
+ addToken(timestampToken, tokenTimestampLogits, confidence)
534
+ }
535
+ } else if (indexOfMaxTextLogProb == eotToken) {
536
+ break
537
+ } else {
538
+ let chosenTokenIndex: number
539
+
540
+ if (options.temperature == 0.0) {
541
+ chosenTokenIndex = indexOfMaxTextLogProb
542
+ } else {
543
+ const topLogitCount = options.topCandidateCount!
544
+
545
+ const textTokenLogits = tokenLogits.slice(0, timestampTokensStart)
546
+ const sortedTextTokenLogitsWithIndexes = Array.from(textTokenLogits).map((logit, index) => ({ logit, index }))
547
+ sortedTextTokenLogitsWithIndexes.sort((a, b) => b.logit - a.logit)
548
+ let topLogitsWithIndexes = sortedTextTokenLogitsWithIndexes.slice(0, topLogitCount)
549
+
550
+ ////
551
+ /*
552
+ topLogitsWithIndexes = topLogitsWithIndexes.filter(entry => {
553
+ const lastDecodedTextTokens = decodedTokens.filter(token => token < eotToken).reverse().slice(0, 20)
554
+ const { maxScore } = getRepetitionScoreRelativeToFirstSubstring([entry.index, ...lastDecodedTextTokens])
555
+
556
+ if (maxScore < 4) {
557
+ return true
558
+ } else {
559
+ return false
560
+ }
561
+ })
562
+ */
563
+ ////
564
+
565
+ const topLogits = topLogitsWithIndexes.map(a => a.logit)
566
+ const textTokenProbs = softmax(topLogits, options.temperature)
567
+
568
+ const topIndexOfPromisingPunctuationLogit = topLogitsWithIndexes.findIndex(entry => {
569
+ const tokenText = (this.tokenToTextLookup.get(entry.index) || "").trim()
570
+ const tokenProb = probs[entry.index]
571
+
572
+ return tokenProb >= options.punctuationThreshold! && [',', ',', '.', '。', '!', '?'].includes(tokenText)
573
+ })
574
+
575
+ let chosenTokenIndexInTopLogits: number
576
+
577
+ if (topIndexOfPromisingPunctuationLogit >= 0) {
578
+ chosenTokenIndexInTopLogits = topIndexOfPromisingPunctuationLogit
579
+ } else {
580
+ chosenTokenIndexInTopLogits = this.randomGen.selectRandomIndexFromDistribution(textTokenProbs)
581
+ }
582
+
583
+ chosenTokenIndex = sortedTextTokenLogitsWithIndexes[chosenTokenIndexInTopLogits].index
584
+ }
585
+
586
+ if (chosenTokenIndex < eotToken) {
587
+ let chosenTokenText = this.tokenToTextLookup.get(chosenTokenIndex) || ""
588
+
589
+ if (isFirstPart && decodedTokens.every(token => token >= eotToken)) {
590
+ chosenTokenText = chosenTokenText.trimStart()
591
+ }
592
+
593
+ writeToStderr(chosenTokenText)
594
+ }
595
+
596
+ const confidence = probs[chosenTokenIndex]
597
+
598
+ addToken(chosenTokenIndex, tokenTimestampLogits, confidence)
599
+ }
600
+
601
+ await yieldToEventLoop()
602
+ }
603
+
604
+ if (timestampsSeenCount >= 2 && !isFinalPart) {
605
+ decodedTokens = decodedTokens.slice(0, lastTimestampTokenIndex)
606
+ decodedTokensTimestampLogits = decodedTokensTimestampLogits.slice(0, lastTimestampTokenIndex)
607
+ decodedTokensCrossAttentionQKs = decodedTokensCrossAttentionQKs.slice(0, lastTimestampTokenIndex)
608
+ }
609
+
610
+ writeToStderr("\n")
611
+ logger.end()
612
+
613
+ // Return the tokens
614
+ return { decodedTokens, decodedTokensTimestampLogits, crossAttentionQKs: decodedTokensCrossAttentionQKs, decodedTokensConfidence }
615
+ }
616
+
617
+ async inferCrossAttentionQKs(tokens: number[], audioFeatures: Onnx.Tensor) {
618
+ const offset = 0
619
+
620
+ const tokensTensor = new Onnx.Tensor('int64', new BigInt64Array(tokens.map(token => BigInt(token))), [1, tokens.length])
621
+ const offsetTensor = new Onnx.Tensor('int64', new BigInt64Array([BigInt(offset)]), [])
622
+
623
+ const initialKvDimensions = this.getKvDimensions(1, tokens.length)
624
+ const kvCacheTensor = new Onnx.Tensor('float32', new Float32Array(initialKvDimensions[0] * initialKvDimensions[1] * initialKvDimensions[2] * initialKvDimensions[3]), initialKvDimensions)
625
+
626
+ const decoderInputs = { tokens: tokensTensor, audio_features: audioFeatures, kv_cache: kvCacheTensor, offset: offsetTensor }
627
+
628
+ const decoderOutputs = await this.textDecoder!.run(decoderInputs)
629
+
630
+ const crossAttentionQKsTensor = decoderOutputs["cross_attention_qks"]
631
+
632
+ const tensorShape = crossAttentionQKsTensor.dims.slice()
633
+
634
+ const ndarray = (await import('ndarray')).default
635
+
636
+ let qkArray = ndarray(crossAttentionQKsTensor.data, crossAttentionQKsTensor.dims.slice())
637
+ qkArray = qkArray.transpose(3, 0, 1, 2, 4)
638
+
639
+ const tokenCrossAttentionQKsTensors: Onnx.Tensor[] = []
640
+
641
+ for (let i0 = 0; i0 < qkArray.shape[0]; i0++) {
642
+ const dataForToken: number[] = []
643
+
644
+ for (let i1 = 0; i1 < qkArray.shape[1]; i1++) {
645
+ for (let i2 = 0; i2 < qkArray.shape[2]; i2++) {
646
+ for (let i3 = 0; i3 < qkArray.shape[3]; i3++) {
647
+ for (let i4 = 0; i4 < qkArray.shape[4]; i4++) {
648
+ dataForToken.push(qkArray.get(i0, i1, i2, i3, i4) as number)
649
+ }
650
+ }
651
+ }
652
+ }
653
+
654
+ const newTensorShape = tensorShape.slice()
655
+ newTensorShape[3] = 1
656
+
657
+ const newTensor = new Onnx.Tensor('float32', dataForToken, newTensorShape)
658
+
659
+ tokenCrossAttentionQKsTensors.push(newTensor)
660
+ }
661
+
662
+ return tokenCrossAttentionQKsTensors
663
+ }
664
+
665
+ async encodeAudio(rawAudio: RawAudio) {
666
+ const logger = new Logger()
667
+
668
+ const audioSamples = rawAudio.audioChannels[0]
669
+ const sampleRate = rawAudio.sampleRate
670
+
671
+ const fftOrder = 400
672
+ const hopLength = 160
673
+ const filterbankCount = 80
674
+
675
+ const maxAudioSamples = sampleRate * 30
676
+ const maxAudioFrames = 3000
677
+
678
+ await logger.startAsync("Extract mel spectogram from audio part")
679
+
680
+ const paddedAudioSamples = new Float32Array(maxAudioSamples)
681
+ paddedAudioSamples.set(audioSamples.subarray(0, maxAudioSamples), 0)
682
+
683
+ const rawAudioPart: RawAudio = { audioChannels: [paddedAudioSamples], sampleRate }
684
+
685
+ const { melSpectogram } = await computeMelSpectogramUsingFilterbanks(rawAudioPart, fftOrder, fftOrder, hopLength, filterbanks)
686
+
687
+ await logger.startAsync("Normalize mel spectogram")
688
+
689
+ const logMelSpectogram = melSpectogram.map(spectrum => spectrum.map(mel => Math.log10(Math.max(mel, 1e-10))))
690
+ let maxLogMel = -Infinity
691
+
692
+ for (const spectrum of logMelSpectogram) {
693
+ for (const mel of spectrum) {
694
+ if (mel > maxLogMel) {
695
+ maxLogMel = mel
696
+ }
697
+ }
698
+ }
699
+
700
+ const normalizedLogMelSpectogram = logMelSpectogram.map(spectrum => spectrum.map(
701
+ logMel => (Math.max(logMel, maxLogMel - 8) + 4) / 4))
702
+
703
+ const flattenedNormalizedLogMelSpectogram = new Float32Array(maxAudioFrames * filterbankCount)
704
+
705
+ for (let i = 0; i < filterbankCount; i++) {
706
+ for (let j = 0; j < maxAudioFrames; j++) {
707
+ flattenedNormalizedLogMelSpectogram[(i * maxAudioFrames) + j] = normalizedLogMelSpectogram[j][i]
708
+ }
709
+ }
710
+
711
+ await logger.startAsync("Encode mel spectogram with Whisper encoder model")
712
+
713
+ const inputTensor = new Onnx.Tensor('float32', flattenedNormalizedLogMelSpectogram, [1, filterbankCount, maxAudioFrames])
714
+
715
+ const encoderInputs = { mel: inputTensor }
716
+
717
+ const encoderOutputs = await this.audioEncoder!.run(encoderInputs)
718
+ const encodedAudioFeatures = encoderOutputs["output"]
719
+
720
+ logger.end()
721
+
722
+ return encodedAudioFeatures
723
+ }
724
+
725
+ addSegmentsToTimeline(timeline: Timeline, tokens: number[], initialTimeOffset: number, audioDuration: number) {
726
+ const timestampTokensStart = this.tokenConfig.timestampTokensStart
727
+
728
+ for (let i = 0; i < tokens.length; i++) {
729
+ const token = tokens[i]
730
+
731
+ if (token == this.tokenConfig.sotToken || token == this.tokenConfig.eotToken) {
732
+ continue
733
+ }
734
+
735
+ const tokenIsTimestamp = token >= timestampTokensStart
736
+ const previousTokenWasTimestamp = tokens.length > 1 && tokens[i - 1] >= timestampTokensStart
737
+
738
+ if (tokenIsTimestamp) {
739
+ if (previousTokenWasTimestamp) {
740
+ continue
741
+ }
742
+
743
+ let startTime = initialTimeOffset + (token - timestampTokensStart) * 0.02
744
+
745
+ startTime = Math.min(startTime, audioDuration)
746
+
747
+ if (timeline.length > 0) {
748
+ timeline[timeline.length - 1].endTime = startTime
749
+ }
750
+
751
+ timeline.push({
752
+ type: "segment",
753
+ text: "",
754
+ startTime,
755
+ endTime: -1,
756
+ })
757
+ } else {
758
+ if (timeline.length == 0) {
759
+ timeline.push({
760
+ type: "segment",
761
+ text: "",
762
+ startTime: initialTimeOffset,
763
+ endTime: -1,
764
+ })
765
+ }
766
+
767
+ const tokenText = this.tokenToTextLookup.get(token) || ""
768
+
769
+ timeline[timeline.length - 1].text += tokenText
770
+ }
771
+ }
772
+ }
773
+
774
+ async addWordsToTimeline(timeline: Timeline, tokens: number[], rawAudio: RawAudio, crossAttentionQKs: Onnx.Tensor[], initialAudioTimeOffset: number, duration: number) {
775
+ const timestampTokensStart = this.tokenConfig.timestampTokensStart
776
+
777
+ let segmentStartTime = 0
778
+ let segmentTokens: number[] = []
779
+ let segmentCrossAttentionQKs: Onnx.Tensor[] = []
780
+
781
+ for (let tokenIndex = 0; tokenIndex < tokens.length; tokenIndex++) {
782
+ const token = tokens[tokenIndex]
783
+ const tokenCrossAttentionQKs = crossAttentionQKs[tokenIndex]
784
+
785
+ const segmentTokensWithoutTimestamps = segmentTokens.filter(token => token < this.tokenConfig.timestampTokensStart)
786
+
787
+ const isTimestamp = token >= timestampTokensStart
788
+
789
+ if (isTimestamp || tokenIndex == tokens.length - 1) {
790
+ let tokenTime: number
791
+
792
+ if (isTimestamp) {
793
+ tokenTime = (token - timestampTokensStart) * 0.02
794
+ } else {
795
+ tokenTime = duration
796
+ }
797
+
798
+ if (segmentTokensWithoutTimestamps.length > 0) {
799
+ const segmentEndTime = tokenTime
800
+
801
+ const segmentStartFrame = Math.floor(segmentStartTime / 0.02)
802
+ let segmentEndFrame = Math.floor(segmentEndTime / 0.02)
803
+
804
+ if (segmentStartFrame == segmentEndFrame) {
805
+ segmentEndFrame += 1
806
+ }
807
+
808
+ const segmentFrameCount = segmentEndFrame - segmentStartFrame
809
+
810
+ const reinferCrossAttentionQKs = true
811
+
812
+ if (reinferCrossAttentionQKs) {
813
+ const initialTokens = this.getInitialTokens('en', 'transcribe')
814
+ const tokensToDecode = [...initialTokens, ...segmentTokensWithoutTimestamps]
815
+
816
+ //const segmentAudioFeaturesBuffer = audioFeatures.data.slice(segmentStartFrame * audioFeatures.dims[2], segmentEndFrame * audioFeatures.dims[2])
817
+ //const segmentAudioFeatures = new Onnx.Tensor('float32', segmentAudioFeaturesBuffer, [1, segmentFrameCount, audioFeatures.dims[2]])
818
+
819
+ const segmentAudioSamples = rawAudio.audioChannels[0].slice(Math.floor(segmentStartTime * rawAudio.sampleRate), Math.floor(segmentEndTime * rawAudio.sampleRate))
820
+ const segmentRawAudio: RawAudio = { audioChannels: [segmentAudioSamples], sampleRate: rawAudio.sampleRate }
821
+
822
+ const segmentAudioFeatures = await this.encodeAudio(segmentRawAudio)
823
+
824
+ const reinferredCrossAttentionQKs = await this.inferCrossAttentionQKs(tokensToDecode, segmentAudioFeatures)
825
+ reinferredCrossAttentionQKs.slice(initialTokens.length)
826
+
827
+ const alignmentPath = await this.findAlignmentPathFromQKs(reinferredCrossAttentionQKs, tokensToDecode, 0, segmentFrameCount)//, alignmentHeadsIndexes[modelName])
828
+ const wordTimeline = await this.getWordTimelineFromAlignmentPath(alignmentPath, segmentTokensWithoutTimestamps, initialAudioTimeOffset + segmentStartTime, initialAudioTimeOffset + segmentEndTime)
829
+
830
+ timeline.push(...wordTimeline)
831
+ } else {
832
+ const alignmentPath = await this.findAlignmentPathFromQKs(segmentCrossAttentionQKs, segmentTokens, segmentStartFrame, segmentEndFrame)//, alignmentHeadsIndexes[modelName])
833
+ const wordTimeline = await this.getWordTimelineFromAlignmentPath(alignmentPath, segmentTokens, initialAudioTimeOffset, initialAudioTimeOffset + segmentEndTime)
834
+
835
+ timeline.push(...wordTimeline)
836
+ }
837
+ }
838
+
839
+ segmentStartTime = tokenTime
840
+ segmentTokens = []
841
+ segmentCrossAttentionQKs = []
842
+ }
843
+
844
+ segmentTokens.push(token)
845
+ segmentCrossAttentionQKs.push(tokenCrossAttentionQKs)
846
+ }
847
+ }
848
+
849
+ mergeSuccessiveWordFragmentsInTimeline(timeline: Timeline) {
850
+ const resultTimeline: Timeline = []
851
+
852
+ const groups: TimelineEntry[][] = []
853
+
854
+ for (const entry of timeline) {
855
+ if (entry.type != "word") {
856
+ continue
857
+ }
858
+
859
+ if (groups.length == 0 || entry.text.startsWith(" ")) {
860
+ groups.push([entry])
861
+ } else {
862
+ groups[groups.length - 1].push(entry)
863
+ }
864
+ }
865
+
866
+ for (const group of groups) {
867
+ if (group.length == 1) {
868
+ resultTimeline.push(deepClone(group[0]))
869
+ } else {
870
+ const text = group.map(entry => entry.text).join("")
871
+ const startTime = group[0].startTime
872
+ const endTime = group[group.length - 1].endTime
873
+ let confidence: number | undefined = undefined
874
+
875
+ if (group[0].confidence != null) {
876
+ confidence = meanOfVector(group.map(entry => entry.confidence!))
877
+ }
878
+
879
+ const newEntry: TimelineEntry = {
880
+ type: "word",
881
+ text,
882
+ startTime,
883
+ endTime,
884
+ confidence
885
+ }
886
+
887
+ resultTimeline.push(newEntry)
888
+ }
889
+ }
890
+
891
+ return resultTimeline
892
+ }
893
+
894
+ async getWordTimelineFromAlignmentPath(alignmentPath: AlignmentPath, tokens: number[], startTimeOffset: number, endTimeOffset: number, tokensConfidence?: number[], correctionAmount = 0.0) {
895
+ if (alignmentPath.length == 0) {
896
+ return []
897
+ }
898
+
899
+ const wordTimeline: Timeline = []
900
+
901
+ for (let pathIndex = 0; pathIndex < alignmentPath.length; pathIndex++) {
902
+ if (pathIndex != 0 && alignmentPath[pathIndex].source == alignmentPath[pathIndex - 1].source) {
903
+ continue
904
+ }
905
+
906
+ const tokenMappingEntry = alignmentPath[pathIndex]
907
+
908
+ const tokenIndex = tokenMappingEntry.source
909
+ const token = tokens[tokenIndex]
910
+ const tokenConfidence = tokensConfidence ? tokensConfidence[tokenIndex] : undefined
911
+ const tokenText = this.tokenToTextLookup.get(token)
912
+
913
+ if (token >= this.tokenConfig.eotToken || !tokenText) {
914
+ continue
915
+ }
916
+
917
+ let startTime = startTimeOffset + (tokenMappingEntry.dest * 0.02)
918
+
919
+ startTime = Math.max(startTime + correctionAmount, startTimeOffset)
920
+
921
+ if (wordTimeline.length > 0) {
922
+ wordTimeline[wordTimeline.length - 1].endTime = startTime
923
+ }
924
+
925
+ wordTimeline.push({
926
+ type: "word",
927
+ text: tokenText,
928
+ startTime,
929
+ endTime: -1,
930
+ confidence: tokenConfidence
931
+ })
932
+ }
933
+
934
+ if (wordTimeline.length > 0) {
935
+ wordTimeline[wordTimeline.length - 1].endTime = endTimeOffset
936
+ }
937
+
938
+ return wordTimeline
939
+ }
940
+
941
+ async findAlignmentPathFromQKs(qksTensors: Onnx.Tensor[], tokens: number[], segmentStartFrame: number, segmentEndFrame: number, headIndexes?: number[]) {
942
+ const segmentFrameCount = segmentEndFrame - segmentStartFrame
943
+
944
+ if (segmentFrameCount == 0) {
945
+ //throw new Error("Segment has 0 frames")
946
+ return []
947
+ }
948
+
949
+ const tokenCount = qksTensors.length
950
+ const layerCount = qksTensors[0].dims[0]
951
+ const headCount = qksTensors[0].dims[2]
952
+ const frameCount = qksTensors[0].dims[4]
953
+
954
+ if (!headIndexes) {
955
+ headIndexes = []
956
+
957
+ for (let i = 0; i < layerCount * headCount; i++) {
958
+ //for (let i = Math.floor(layerCount * headCount / 2); i < layerCount * headCount; i++) {
959
+ headIndexes.push(i)
960
+ }
961
+ }
962
+
963
+ // Load attention head weights from tensors
964
+ const attentionHeads: number[][][] = [] // [heads, tokens, frames]
965
+
966
+ for (const headIndex of headIndexes) {
967
+ const attentionHead: number[][] = [] // [tokens, frames]
968
+
969
+ for (let tokenIndex = 0; tokenIndex < tokenCount; tokenIndex++) {
970
+ const bufferOffset = headIndex * frameCount
971
+ const startIndexInBuffer = bufferOffset + segmentStartFrame
972
+ const endIndexInBuffer = bufferOffset + segmentEndFrame
973
+
974
+ const framesForHead = qksTensors[tokenIndex].data.slice(startIndexInBuffer, endIndexInBuffer)
975
+
976
+ attentionHead.push(Array.from(framesForHead as any))
977
+ }
978
+
979
+ attentionHeads.push(attentionHead)
980
+ }
981
+
982
+ const applySoftmax = true
983
+ const normalize = true
984
+ const applyMedianFilter = true
985
+ const fixateTimestampTokens = false
986
+
987
+ const softmaxTemperature = 1.0
988
+ const medianFilterWidth = 7
989
+
990
+ if (applySoftmax) {
991
+ // Apply softmax to each token's frames
992
+ for (const head of attentionHeads) {
993
+ for (let tokenIndex = 0; tokenIndex < tokenCount; tokenIndex++) {
994
+ head[tokenIndex] = softmax(head[tokenIndex], softmaxTemperature)
995
+ }
996
+ }
997
+ }
998
+
999
+ if (normalize) {
1000
+ // Normalize all weights in each individual head
1001
+ for (const head of attentionHeads) {
1002
+ const allWeightsForHead = head.flatMap(tokenFrames => tokenFrames)
1003
+
1004
+ const meanOfAllWeights = meanOfVector(allWeightsForHead)
1005
+ const stdDeviationOfAllWeights = stdDeviationOfVector(allWeightsForHead)
1006
+
1007
+ for (const tokenFrames of head) {
1008
+ for (let frameIndex = 0; frameIndex < tokenFrames.length; frameIndex++) {
1009
+ tokenFrames[frameIndex] = (tokenFrames[frameIndex] - meanOfAllWeights) / stdDeviationOfAllWeights
1010
+ }
1011
+ }
1012
+ }
1013
+ }
1014
+
1015
+ if (applyMedianFilter) {
1016
+ // Apply median filter to each token's frames
1017
+ for (const head of attentionHeads) {
1018
+ for (let tokenIndex = 0; tokenIndex < tokenCount; tokenIndex++) {
1019
+ head[tokenIndex] = medianFilter(head[tokenIndex], medianFilterWidth)
1020
+ }
1021
+ }
1022
+ }
1023
+
1024
+ // Compute the mean for all layers and heads
1025
+ const frameMeansForToken: number[][] = []
1026
+
1027
+ for (let i = 0; i < tokenCount; i++) {
1028
+ const frameMeans = new Array(segmentFrameCount)
1029
+
1030
+ frameMeansForToken.push(frameMeans)
1031
+ }
1032
+
1033
+ for (let tokenIndex = 0; tokenIndex < tokenCount; tokenIndex++) {
1034
+ for (let frameIndex = 0; frameIndex < segmentFrameCount; frameIndex++) {
1035
+ let sum = 0
1036
+
1037
+ for (const head of attentionHeads) {
1038
+ sum += head[tokenIndex][frameIndex]
1039
+ }
1040
+
1041
+ const frameMean = sum / attentionHeads.length
1042
+
1043
+ frameMeansForToken[tokenIndex][frameIndex] = frameMean
1044
+ }
1045
+ }
1046
+
1047
+ if (fixateTimestampTokens) {
1048
+ const timestampTokensStart = this.tokenConfig.timestampTokensStart
1049
+
1050
+ for (let tokenIndex = 0; tokenIndex < tokens.length; tokenIndex++) {
1051
+ if (tokens[tokenIndex] >= timestampTokensStart) {
1052
+ let timestampFrame = tokens[tokenIndex] - timestampTokensStart
1053
+ timestampFrame = clip(timestampFrame, segmentStartFrame, segmentEndFrame - 1)
1054
+
1055
+ frameMeansForToken[tokenIndex][timestampFrame] = 100
1056
+ }
1057
+ }
1058
+ }
1059
+
1060
+ // Perform DTW
1061
+ const tokenIndexes = [...Array(tokenCount).keys()]
1062
+ const frameIndexes = [...Array(segmentFrameCount).keys()]
1063
+
1064
+ let { path } = await alignDTWWindowed(tokenIndexes, frameIndexes, (tokenIndex, frameIndex) => {
1065
+ return -frameMeansForToken[tokenIndex][frameIndex]
1066
+ }, 1000)
1067
+
1068
+ path = path.map(entry => ({ source: entry.source, dest: segmentStartFrame + entry.dest }))
1069
+
1070
+ return path
1071
+ }
1072
+
1073
+ getKvDimensions(groupCount: number, length: number) {
1074
+ const modelName = this.modelName
1075
+
1076
+ if (modelName == "tiny" || modelName == "tiny.en") {
1077
+ return [8, groupCount, length, 384]
1078
+ } else if (modelName == "base" || modelName == "base.en") {
1079
+ return [12, groupCount, length, 512]
1080
+ } else if (modelName == "small" || modelName == "small.en") {
1081
+ return [24, groupCount, length, 768]
1082
+ } else if (modelName == "medium" || modelName == "medium.en") {
1083
+ return [48, groupCount, length, 1024]
1084
+ } else if (modelName == "large" || modelName == "large-v1" || modelName == "large-v2") {
1085
+ return [64, groupCount, length, 1280]
1086
+ } else {
1087
+ throw new Error(`Unsupported model: ${modelName}`)
1088
+ }
1089
+ }
1090
+
1091
+ getInitialTokens(language: string, task: WhisperTask, disableTimestamps = false) {
1092
+ const sotToken = this.tokenConfig.sotToken
1093
+
1094
+ let initialTokens: number[]
1095
+
1096
+ if (this.isMultiligualModel) {
1097
+ const languageToken = sotToken + 1 + languageIdLookup[language]
1098
+ const translateTaskToken = 50358
1099
+ const transcribeTaskToken = 50359
1100
+ const taskToken = task == "transcribe" ? transcribeTaskToken : translateTaskToken
1101
+
1102
+ initialTokens = [sotToken, languageToken, taskToken]
1103
+ } else {
1104
+ initialTokens = [sotToken]
1105
+ }
1106
+
1107
+ if (disableTimestamps) {
1108
+ initialTokens.push(this.tokenConfig.noTimestampsToken)
1109
+ }
1110
+
1111
+ return initialTokens
1112
+ }
1113
+
1114
+ getAlignmentHeadIndexes() {
1115
+ return alignmentHeadsIndexes[this.modelName]
1116
+ }
1117
+
1118
+ tokensToText(tokens: number[]) {
1119
+ return tokens.map(token => this.tokenToTextLookup.get(token) || "").join("").trim()
1120
+ }
1121
+
1122
+ async textToTokens(text: string, language: string) {
1123
+ const resultTokens: number[] = []
1124
+
1125
+ const words = (await splitToWords(text, language)).filter(w => w.trim().length > 0)
1126
+
1127
+ //words = words.filter(word => wordCharacterPattern.test(word))
1128
+
1129
+ for (let i = 1; i < words.length; i++) {
1130
+ words[i] = ` ${words[i]}`
1131
+ }
1132
+
1133
+ const allResultingSubwords: string[][] = []
1134
+
1135
+ for (const word of words) {
1136
+ const tokenForEntireWord = this.textToTokenLookup.get(word)
1137
+
1138
+ if (tokenForEntireWord) {
1139
+ resultTokens.push(tokenForEntireWord)
1140
+ allResultingSubwords.push([word])
1141
+ continue
1142
+ }
1143
+
1144
+ const subwords = word.split("")
1145
+
1146
+ for (const mergeRule of this.merges) {
1147
+ for (let i = 0; i < subwords.length - 1; i++) {
1148
+ const currentSubword = subwords[i]
1149
+ const nextSubword = subwords[i + 1]
1150
+
1151
+ if (currentSubword == mergeRule[0] && nextSubword == mergeRule[1]) {
1152
+ subwords.splice(i, 2, mergeRule[0] + mergeRule[1])
1153
+ }
1154
+ }
1155
+ }
1156
+
1157
+ for (const subword of subwords) {
1158
+ const tokenForSubword = this.textToTokenLookup.get(subword)
1159
+
1160
+ if (!tokenForSubword) {
1161
+ throw new Error(`Failed tokenizing the given text. The word '${word}' contains a subword '${subword}' which is not in the vocabulary.`)
1162
+ }
1163
+
1164
+ resultTokens.push(tokenForSubword)
1165
+ }
1166
+
1167
+ allResultingSubwords.push(subwords)
1168
+ }
1169
+
1170
+ return resultTokens
1171
+ }
1172
+ }
1173
+
1174
+ const filterbanks: Filterbank[] = [
1175
+ /* 0 */ { startIndex: 1, weights: [0.02486259490251541,] },
1176
+
1177
+ /* 1 */ { startIndex: 1, weights: [0.001990821911022067, 0.022871771827340126,] },
1178
+
1179
+ /* 2 */ { startIndex: 2, weights: [0.003981643822044134, 0.02088095061480999,] },
1180
+
1181
+ /* 3 */ { startIndex: 3, weights: [0.0059724655002355576, 0.018890129402279854,] },
1182
+
1183
+ /* 4 */ { startIndex: 4, weights: [0.007963287644088268, 0.01689930632710457,] },
1184
+
1185
+ /* 5 */ { startIndex: 5, weights: [0.009954108856618404, 0.014908484183251858,] },
1186
+
1187
+ /* 6 */ { startIndex: 6, weights: [0.011944931000471115, 0.012917662039399147,] },
1188
+
1189
+ /* 7 */ { startIndex: 7, weights: [0.013935752213001251, 0.010926840826869011,] },
1190
+
1191
+ /* 8 */ { startIndex: 8, weights: [0.015926575288176537, 0.0089360186830163,] },
1192
+
1193
+ /* 9 */ { startIndex: 9, weights: [0.017917396500706673, 0.006945197004824877,] },
1194
+
1195
+ /* 10 */ { startIndex: 10, weights: [0.01990821771323681, 0.004954374860972166,] },
1196
+
1197
+ /* 11 */ { startIndex: 11, weights: [0.021899040788412094, 0.0029635531827807426,] },
1198
+
1199
+ /* 12 */ { startIndex: 12, weights: [0.02388986200094223, 0.0009727313299663365,] },
1200
+
1201
+ /* 13 */ { startIndex: 13, weights: [0.025880683213472366,] },
1202
+
1203
+ /* 14 */ { startIndex: 14, weights: [0.025835324078798294,] },
1204
+
1205
+ /* 15 */ { startIndex: 14, weights: [0.0010180906392633915, 0.023844502866268158,] },
1206
+
1207
+ /* 16 */ { startIndex: 15, weights: [0.003008912317454815, 0.021853681653738022,] },
1208
+
1209
+ /* 17 */ { startIndex: 16, weights: [0.004999734461307526, 0.019862858578562737,] },
1210
+
1211
+ /* 18 */ { startIndex: 17, weights: [0.006990555673837662, 0.0178720373660326,] },
1212
+
1213
+ /* 19 */ { startIndex: 18, weights: [0.008981377817690372, 0.015881216153502464,] },
1214
+
1215
+ /* 20 */ { startIndex: 19, weights: [0.010972199961543083, 0.013890394009649754,] },
1216
+
1217
+ /* 21 */ { startIndex: 20, weights: [0.01296302117407322, 0.011899571865797043,] },
1218
+
1219
+ /* 22 */ { startIndex: 21, weights: [0.01495384331792593, 0.009908749721944332,] },
1220
+
1221
+ /* 23 */ { startIndex: 22, weights: [0.01694466546177864, 0.007917927578091621,] },
1222
+
1223
+ /* 24 */ { startIndex: 23, weights: [0.018935488536953926, 0.005927106365561485,] },
1224
+
1225
+ /* 25 */ { startIndex: 24, weights: [0.020874010398983955, 0.004040425643324852,] },
1226
+
1227
+ /* 26 */ { startIndex: 25, weights: [0.022114217281341553, 0.0033186059445142746,] },
1228
+
1229
+ /* 27 */ { startIndex: 26, weights: [0.02173672430217266, 0.0036109676584601402,] },
1230
+
1231
+ /* 28 */ { startIndex: 27, weights: [0.020497702062129974, 0.004762193653732538,] },
1232
+
1233
+ /* 29 */ { startIndex: 28, weights: [0.018486659973859787, 0.006592618301510811,] },
1234
+
1235
+ /* 30 */ { startIndex: 29, weights: [0.01585603691637516, 0.00896277092397213,] },
1236
+
1237
+ /* 31 */ { startIndex: 30, weights: [0.012738768011331558, 0.011751330457627773,] },
1238
+
1239
+ /* 32 */ { startIndex: 31, weights: [0.009250369854271412, 0.014853144995868206,] },
1240
+
1241
+ /* 33 */ { startIndex: 32, weights: [0.005490840878337622, 0.018177473917603493, 0.0028155462350696325,] },
1242
+
1243
+ /* 34 */ { startIndex: 33, weights: [0.0015463664894923568, 0.01632951945066452, 0.007420188747346401,] },
1244
+
1245
+ /* 35 */ { startIndex: 35, weights: [0.011181050911545753, 0.012018864043056965,] },
1246
+
1247
+ /* 36 */ { startIndex: 36, weights: [0.006065350491553545, 0.016561277210712433, 0.004360878840088844,] },
1248
+
1249
+ /* 37 */ { startIndex: 37, weights: [0.0010297985281795263, 0.012770536355674267, 0.009707189165055752,] },
1250
+
1251
+ /* 38 */ { startIndex: 39, weights: [0.006986402906477451, 0.01485429983586073, 0.004391219466924667,] },
1252
+
1253
+ /* 39 */ { startIndex: 40, weights: [0.001418047584593296, 0.011486922390758991, 0.010089744813740253, 0.00040022286702878773,] },
1254
+
1255
+ /* 40 */ { startIndex: 42, weights: [0.005411104764789343, 0.014735566452145576, 0.006518189795315266,] },
1256
+
1257
+ /* 41 */ { startIndex: 44, weights: [0.00827841367572546, 0.012277561239898205, 0.00396781275048852,] },
1258
+
1259
+ /* 42 */ { startIndex: 45, weights: [0.002187808509916067, 0.010184479877352715, 0.00998187530785799, 0.0022864851634949446,] },
1260
+
1261
+ /* 43 */ { startIndex: 47, weights: [0.00386943481862545, 0.011274894699454308, 0.008466221392154694, 0.0013397691072896123,] },
1262
+
1263
+ /* 44 */ { startIndex: 49, weights: [0.004820294212549925, 0.011678251437842846, 0.007608682848513126, 0.0010091039584949613,] },
1264
+
1265
+ /* 45 */ { startIndex: 51, weights: [0.005156961735337973, 0.011507894843816757, 0.007301822770386934, 0.0011901655234396458,] },
1266
+
1267
+ /* 46 */ { startIndex: 53, weights: [0.004982104524970055, 0.010863498784601688, 0.007451189681887627, 0.001791381393559277,] },
1268
+
1269
+ /* 47 */ { startIndex: 55, weights: [0.004385921638458967, 0.009832492098212242, 0.007973956875503063, 0.002732589840888977,] },
1270
+
1271
+ /* 48 */ { startIndex: 57, weights: [0.0034474546555429697, 0.008491347543895245, 0.008797688409686089, 0.00394382793456316,] },
1272
+
1273
+ /* 49 */ { startIndex: 59, weights: [0.0022357646375894547, 0.0069067515432834625, 0.009859241545200348, 0.005364237818866968, 0.0008692338014952838,] },
1274
+
1275
+ /* 50 */ { startIndex: 61, weights: [0.0008110002381727099, 0.005136650986969471, 0.00946230161935091, 0.0069410777650773525, 0.0027783995028585196,] },
1276
+
1277
+ /* 51 */ { startIndex: 64, weights: [0.003231203882023692, 0.007237049750983715, 0.00862883497029543, 0.004773912951350212, 0.0009189908159896731,] },
1278
+
1279
+ /* 52 */ { startIndex: 66, weights: [0.001233637798577547, 0.0049433219246566296, 0.008653006516397, 0.006818502210080624, 0.003248583758249879,] },
1280
+
1281
+ /* 53 */ { startIndex: 69, weights: [0.0026164355222135782, 0.006051854696124792, 0.008880467154085636, 0.005574479699134827, 0.002268492942675948,] },
1282
+
1283
+ /* 54 */ { startIndex: 71, weights: [0.0002863667905330658, 0.003467798000201583, 0.0066492292098701, 0.00787146482616663, 0.004809896927326918, 0.0017483289120718837,] },
1284
+
1285
+ /* 55 */ { startIndex: 74, weights: [0.0009245910914614797, 0.0038708120118826628, 0.00681703258305788, 0.007283343467861414, 0.004448124207556248, 0.0016129047144204378,] },
1286
+
1287
+ /* 56 */ { startIndex: 77, weights: [0.0011703289346769452, 0.003898728871718049, 0.006627128925174475, 0.0070473202504217625, 0.004421714693307877, 0.0017961094854399562,] },
1288
+
1289
+ /* 57 */ { startIndex: 80, weights: [0.0010892992140725255, 0.003615982597693801, 0.006142666097730398, 0.007102936040610075, 0.004671447444707155, 0.002239959081634879,] },
1290
+
1291
+ /* 58 */ { startIndex: 83, weights: [0.0007392280967906117, 0.0030791081953793764, 0.005418988410383463, 0.007397185545414686, 0.005145462695509195, 0.002893739379942417, 0.0006420162972062826,] },
1292
+
1293
+ /* 59 */ { startIndex: 86, weights: [0.00017068670422304422, 0.0023375742603093386, 0.004504461772739887, 0.0066713495180010796, 0.005798479542136192, 0.003713231300935149, 0.0016279831761494279,] },
1294
+
1295
+ /* 60 */ { startIndex: 90, weights: [0.0014345343224704266, 0.0034412189852446318, 0.005447904113680124, 0.006591092795133591, 0.004660011734813452, 0.002728930441662669, 0.0007978491485118866,] },
1296
+
1297
+ /* 61 */ { startIndex: 93, weights: [0.0004075043834745884, 0.002265830524265766, 0.004124156199395657, 0.005982482805848122, 0.005700822453945875, 0.003912510350346565, 0.0021241982467472553, 0.0003358862304594368,] },
1298
+
1299
+ /* 62 */ { startIndex: 97, weights: [0.0010099108330905437, 0.002730846870690584, 0.004451782442629337, 0.006172718480229378, 0.005150905344635248, 0.0034948070533573627, 0.0018387088784947991, 0.0001826105872169137,] },
1300
+
1301
+ /* 63 */ { startIndex: 101, weights: [0.0012943691108375788, 0.002888072282075882, 0.004481775686144829, 0.006075479090213776, 0.0048866597935557365, 0.003353001084178686, 0.0018193417927250266, 0.00028568264679051936,] },
1302
+
1303
+ /* 64 */ { startIndex: 105, weights: [0.0013131388695910573, 0.0027890161145478487, 0.004264893010258675, 0.0057407706044614315, 0.004859979264438152, 0.0034397069830447435, 0.0020194342359900475, 0.0005991620710119605,] },
1304
+
1305
+ /* 65 */ { startIndex: 109, weights: [0.0011121684219688177, 0.002478930866345763, 0.0038456933107227087, 0.0052124555222690105, 0.005028639920055866, 0.0037133716978132725, 0.002398103242740035, 0.0010828346712514758,] },
1306
+
1307
+ /* 66 */ { startIndex: 113, weights: [0.0007317548734135926, 0.0019974694587290287, 0.003263183869421482, 0.004528898745775223, 0.005355686880648136, 0.004137659445405006, 0.0029196315445005894, 0.0017016039928421378, 0.0004835762665607035,] },
1308
+
1309
+ /* 67 */ { startIndex: 117, weights: [0.00020713974663522094, 0.0013792773243039846, 0.0025514145381748676, 0.003723552217707038, 0.004895689897239208, 0.004680895246565342, 0.0035529187880456448, 0.0024249425623565912, 0.0012969663366675377, 0.00016899015463422984,] },
1310
+
1311
+ /* 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,] },
1312
+
1313
+ /* 69 */ { startIndex: 127, weights: [0.000854626705404371, 0.001859853626228869, 0.002865080488845706, 0.003870307235047221, 0.00487553421407938, 0.00408313749358058, 0.003115783678367734, 0.0021484296303242445, 0.001181075582280755, 0.0002137213887181133,] },
1314
+
1315
+ /* 70 */ { startIndex: 132, weights: [0.0008483415003865957, 0.0017792496364563704, 0.0027101580053567886, 0.0036410661414265633, 0.004571974277496338, 0.004079728852957487, 0.003183893393725157, 0.002288057701662183, 0.0013922222424298525, 0.0004963868414051831,] },
1316
+
1317
+ /* 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,] },
1318
+
1319
+ /* 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,] },
1320
+
1321
+ /* 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,] },
1322
+
1323
+ /* 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,] },
1324
+
1325
+ /* 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,] },
1326
+
1327
+ /* 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,] },
1328
+
1329
+ /* 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,] },
1330
+
1331
+ /* 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,] },
1332
+
1333
+ /* 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,] },
1334
+ ]
1335
+
1336
+ export async function loadPackagesAndGetPaths(modelName: WhisperModelName | undefined, languageCode: string | undefined) {
1337
+ if (!modelName) {
1338
+ if (languageCode) {
1339
+ const shortLanguageCode = getShortLanguageCode(languageCode)
1340
+
1341
+ modelName = shortLanguageCode == "en" ? "tiny.en" : "tiny"
1342
+ } else {
1343
+ modelName = "tiny"
1344
+ }
1345
+ }
1346
+
1347
+ const packageName = modelNameToPackageName[modelName]
1348
+
1349
+ const modelDir = await loadPackage(packageName)
1350
+
1351
+ const tokenizerPackagePath = await loadPackage(tokenizerPackageName)
1352
+ const tokenizerDir = isMultiligualModel(modelName) ? path.join(tokenizerPackagePath, "multilingual") : path.join(tokenizerPackagePath, "gpt2")
1353
+
1354
+ return { modelName, modelDir, tokenizerDir }
1355
+ }
1356
+
1357
+ export function isMultiligualModel(modelName: WhisperModelName) {
1358
+ return !modelName.endsWith(".en")
1359
+ }
1360
+
1361
+ export type WhisperModelName = "tiny" | "tiny.en" | "base" | "base.en" | "small" | "small.en" | "medium" | "medium.en" | "large" | "large-v1" | "large-v2"
1362
+ export type WhisperTask = "transcribe" | "translate"
1363
+
1364
+ export const modelNameToPackageName: { [modelName in WhisperModelName]: string } = {
1365
+ "tiny": "whisper-tiny",
1366
+ "tiny.en": "whisper-tiny.en",
1367
+ "base": "whisper-base",
1368
+ "base.en": "whisper-base.en",
1369
+ "small": "whisper-small",
1370
+ "small.en": "whisper-small.en",
1371
+ "medium": "whisper-medium",
1372
+ "medium.en": "whisper-medium.en",
1373
+ "large": "whisper-large-v2",
1374
+ "large-v1": "whisper-large-v1",
1375
+ "large-v2": "whisper-large-v2"
1376
+ }
1377
+
1378
+ export const tokenizerPackageName = "whisper-tokenizer"
1379
+
1380
+ const vocabCharacterSetLookup: { [s: string]: number } = {
1381
+ "!": 33, "\"": 34, "#": 35, "$": 36, "%": 37, "&": 38, "'": 39, "(": 40, ")": 41, "*": 42, "+": 43, ",": 44, "-": 45, ".": 46, "/": 47, "0": 48, "1": 49, "2": 50, "3": 51, "4": 52, "5": 53, "6": 54,
1382
+ "7": 55, "8": 56, "9": 57, ":": 58, ";": 59, "<": 60, "=": 61, ">": 62, "?": 63, "@": 64, "A": 65, "B": 66, "C": 67, "D": 68, "E": 69, "F": 70, "G": 71, "H": 72, "I": 73, "J": 74, "K": 75, "L": 76, "M": 77, "N": 78, "O": 79, "P": 80, "Q": 81, "R": 82, "S": 83, "T": 84, "U": 85, "V": 86, "W": 87, "X": 88, "Y": 89, "Z": 90, "[": 91, "\\": 92, "]": 93, "^": 94, "_": 95, "`": 96, "a": 97, "b": 98, "c": 99, "d": 100, "e": 101, "f": 102, "g": 103, "h": 104, "i": 105, "j": 106, "k": 107, "l": 108, "m": 109, "n": 110, "o": 111, "p": 112, "q": 113, "r": 114, "s": 115, "t": 116, "u": 117, "v": 118, "w": 119, "x": 120, "y": 121, "z": 122, "{": 123, "|": 124, "}": 125, "~": 126, "¡": 161, "¢": 162, "£": 163, "¤": 164, "¥": 165, "¦": 166, "§": 167, "¨": 168, "©": 169, "ª": 170, "«": 171, "¬": 172, "®": 174, "¯": 175, "°": 176, "±": 177, "²": 178, "³": 179, "´": 180, "µ": 181, "¶": 182, "·": 183, "¸": 184, "¹": 185, "º": 186, "»": 187, "¼": 188, "½": 189, "¾": 190, "¿": 191, "À": 192, "Á": 193, "Â": 194, "Ã": 195, "Ä": 196, "Å": 197, "Æ": 198, "Ç": 199, "È": 200, "É": 201, "Ê": 202, "Ë": 203, "Ì": 204, "Í": 205, "Î": 206, "Ï": 207, "Ð": 208, "Ñ": 209, "Ò": 210, "Ó": 211, "Ô": 212, "Õ": 213, "Ö": 214, "×": 215, "Ø": 216, "Ù": 217, "Ú": 218, "Û": 219, "Ü": 220, "Ý": 221, "Þ": 222, "ß": 223, "à": 224, "á": 225, "â": 226, "ã": 227, "ä": 228, "å": 229, "æ": 230, "ç": 231, "è": 232, "é": 233, "ê": 234, "ë": 235, "ì": 236, "í": 237, "î": 238, "ï": 239, "ð": 240, "ñ": 241, "ò": 242, "ó": 243, "ô": 244, "õ": 245, "ö": 246, "÷": 247, "ø": 248, "ù": 249, "ú": 250, "û": 251, "ü": 252, "ý": 253, "þ": 254, "ÿ": 255, "Ā": 0, "ā": 1, "Ă": 2, "ă": 3, "Ą": 4, "ą": 5, "Ć": 6, "ć": 7, "Ĉ": 8, "ĉ": 9, "Ċ": 10, "ċ": 11, "Č": 12, "č": 13, "Ď": 14, "ď": 15, "Đ": 16, "đ": 17, "Ē": 18, "ē": 19, "Ĕ": 20, "ĕ":
1383
+ 21, "Ė": 22, "ė": 23, "Ę": 24, "ę": 25, "Ě": 26, "ě": 27, "Ĝ": 28, "ĝ": 29, "Ğ": 30, "ğ": 31, "Ġ": 32, "ġ": 127, "Ģ": 128, "ģ": 129, "Ĥ": 130, "ĥ": 131, "Ħ": 132, "ħ": 133, "Ĩ": 134, "ĩ": 135, "Ī": 136, "ī": 137, "Ĭ": 138, "ĭ": 139, "Į": 140, "į": 141, "İ": 142, "ı": 143, "IJ": 144, "ij": 145, "Ĵ": 146, "ĵ": 147, "Ķ": 148, "ķ": 149, "ĸ": 150, "Ĺ": 151, "ĺ": 152, "Ļ": 153, "ļ": 154, "Ľ": 155, "ľ": 156, "Ŀ": 157, "ŀ": 158, "Ł": 159, "ł": 160, "Ń": 173
1384
+ }
1385
+
1386
+ const languageIdLookup: { [s: string]: number } = {
1387
+ "en": 0,
1388
+ "zh": 1,
1389
+ "de": 2,
1390
+ "es": 3,
1391
+ "ru": 4,
1392
+ "ko": 5,
1393
+ "fr": 6,
1394
+ "ja": 7,
1395
+ "pt": 8,
1396
+ "tr": 9,
1397
+ "pl": 10,
1398
+ "ca": 11,
1399
+ "nl": 12,
1400
+ "ar": 13,
1401
+ "sv": 14,
1402
+ "it": 15,
1403
+ "id": 16,
1404
+ "hi": 17,
1405
+ "fi": 18,
1406
+ "vi": 19,
1407
+ "iw": 20,
1408
+ "uk": 21,
1409
+ "el": 22,
1410
+ "ms": 23,
1411
+ "cs": 24,
1412
+ "ro": 25,
1413
+ "da": 26,
1414
+ "hu": 27,
1415
+ "ta": 28,
1416
+ "no": 29,
1417
+ "th": 30,
1418
+ "ur": 31,
1419
+ "hr": 32,
1420
+ "bg": 33,
1421
+ "lt": 34,
1422
+ "la": 35,
1423
+ "mi": 36,
1424
+ "ml": 37,
1425
+ "cy": 38,
1426
+ "sk": 39,
1427
+ "te": 40,
1428
+ "fa": 41,
1429
+ "lv": 42,
1430
+ "bn": 43,
1431
+ "sr": 44,
1432
+ "az": 45,
1433
+ "sl": 46,
1434
+ "kn": 47,
1435
+ "et": 48,
1436
+ "mk": 49,
1437
+ "br": 50,
1438
+ "eu": 51,
1439
+ "is": 52,
1440
+ "hy": 53,
1441
+ "ne": 54,
1442
+ "mn": 55,
1443
+ "bs": 56,
1444
+ "kk": 57,
1445
+ "sq": 58,
1446
+ "sw": 59,
1447
+ "gl": 60,
1448
+ "mr": 61,
1449
+ "pa": 62,
1450
+ "si": 63,
1451
+ "km": 64,
1452
+ "sn": 65,
1453
+ "yo": 66,
1454
+ "so": 67,
1455
+ "af": 68,
1456
+ "oc": 69,
1457
+ "ka": 70,
1458
+ "be": 71,
1459
+ "tg": 72,
1460
+ "sd": 73,
1461
+ "gu": 74,
1462
+ "am": 75,
1463
+ "yi": 76,
1464
+ "lo": 77,
1465
+ "uz": 78,
1466
+ "fo": 79,
1467
+ "ht": 80,
1468
+ "ps": 81,
1469
+ "tk": 82,
1470
+ "nn": 83,
1471
+ "mt": 84,
1472
+ "sa": 85,
1473
+ "lb": 86,
1474
+ "my": 87,
1475
+ "bo": 88,
1476
+ "tl": 89,
1477
+ "mg": 90,
1478
+ "as": 91,
1479
+ "tt": 92,
1480
+ "haw": 93,
1481
+ "ln": 94,
1482
+ "ha": 95,
1483
+ "ba": 96,
1484
+ "jw": 97,
1485
+ "su": 98,
1486
+ }
1487
+
1488
+ const alignmentHeadsIndexes: { [name in WhisperModelName]: number[] } = {
1489
+ "tiny.en": [6, 12, 17, 18, 19, 20, 21, 22],
1490
+ "tiny": [14, 18, 20, 21, 22, 23],
1491
+ "base.en": [27, 39, 41, 45, 47],
1492
+ "base": [25, 34, 35, 39, 41, 42, 44, 46],
1493
+ "small.en": [78, 84, 87, 92, 98, 101, 103, 108, 112, 116, 118, 120, 121, 122, 123, 126, 131, 134, 136],
1494
+ "small": [63, 69, 96, 100, 103, 104, 108, 115, 117, 125],
1495
+ "medium.en": [180, 225, 236, 238, 244, 256, 260, 265, 284, 286, 295, 298, 303, 320, 323, 329, 334, 348],
1496
+ "medium": [223, 244, 255, 257, 320, 372],
1497
+ "large-v1": [199, 222, 224, 237, 447, 451, 457, 462, 475],
1498
+ "large-v2": [212, 277, 331, 332, 333, 355, 356, 364, 371, 379, 391, 422, 423, 443, 449, 452, 465, 467, 473, 505, 521, 532, 555],
1499
+ "large": [212, 277, 331, 332, 333, 355, 356, 364, 371, 379, 391, 422, 423, 443, 449, 452, 465, 467, 473, 505, 521, 532, 555],
1500
+ }
1501
+
1502
+ export interface WhisperOptions {
1503
+ model?: WhisperModelName
1504
+ temperature?: number
1505
+ prompt?: string
1506
+ topCandidateCount?: number
1507
+ punctuationThreshold?: number
1508
+ autoPromptParts?: boolean
1509
+ }
1510
+
1511
+ export const whisperOptionsDefaults: WhisperOptions = {
1512
+ model: undefined,
1513
+ temperature: 0.1,
1514
+ prompt: undefined,
1515
+ topCandidateCount: 5,
1516
+ punctuationThreshold: 0.2,
1517
+ autoPromptParts: true
1518
+ }