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.
- package/data/schemas/options.json +16 -0
- package/dist/api/Alignment.js +2 -2
- package/dist/api/Alignment.js.map +1 -1
- package/dist/api/Recognition.js +2 -2
- package/dist/api/Recognition.js.map +1 -1
- package/dist/api/Synthesis.js +5 -4
- package/dist/api/Synthesis.js.map +1 -1
- package/dist/api/Translation.js +2 -2
- package/dist/api/Translation.js.map +1 -1
- package/dist/audio/AudioUtilities.d.ts +1 -0
- package/dist/audio/AudioUtilities.js +25 -7
- package/dist/audio/AudioUtilities.js.map +1 -1
- package/dist/cli/CLI.js +2 -2
- package/dist/cli/CLI.js.map +1 -1
- package/dist/recognition/WhisperSTT.js +2 -2
- package/dist/recognition/WhisperSTT.js.map +1 -1
- package/dist/subtitles/Subtitles.d.ts +10 -7
- package/dist/subtitles/Subtitles.js +268 -207
- package/dist/subtitles/Subtitles.js.map +1 -1
- package/docs/Options.md +4 -2
- package/package.json +7 -6
- package/src/alignment/DTWMfccSequenceAlignment.ts +43 -0
- package/src/alignment/DTWSequenceAlignment.ts +121 -0
- package/src/alignment/DTWSequenceAlignmentWindowed.ts +210 -0
- package/src/alignment/LevenshteinSequenceAlignment.ts +126 -0
- package/src/alignment/SpeechAlignment.ts +488 -0
- package/src/api/API.ts +12 -0
- package/src/api/APIOptions.ts +15 -0
- package/src/api/Alignment.ts +329 -0
- package/src/api/Common.ts +16 -0
- package/src/api/Denoising.ts +120 -0
- package/src/api/LanguageDetection.ts +286 -0
- package/src/api/Recognition.ts +344 -0
- package/src/api/Synthesis.ts +1735 -0
- package/src/api/Translation.ts +143 -0
- package/src/api/Vad.ts +172 -0
- package/src/audio/AudioBufferConversion.ts +248 -0
- package/src/audio/AudioPlayer.ts +358 -0
- package/src/audio/AudioRecorder.ts +91 -0
- package/src/audio/AudioUtilities.ts +392 -0
- package/src/audio/SoxPath.ts +24 -0
- package/src/cli/CLI.ts +1360 -0
- package/src/cli/CLIConfigFile.ts +91 -0
- package/src/cli/CLILauncher.ts +26 -0
- package/src/cli/CLIOptionsSchema.ts +54 -0
- package/src/cli/CLIParser.ts +41 -0
- package/src/cli/CLIStarter.ts +40 -0
- package/src/codecs/FFMpegTranscoder.ts +214 -0
- package/src/codecs/TIMITCodec.ts +17 -0
- package/src/codecs/WaveCodec.ts +260 -0
- package/src/denoising/RNNoise.ts +95 -0
- package/src/dsp/BiquadFilter.ts +488 -0
- package/src/dsp/FFT.ts +187 -0
- package/src/dsp/MFCC.ts +227 -0
- package/src/dsp/MelSpectogram.ts +145 -0
- package/src/dsp/Rubberband.ts +249 -0
- package/src/dsp/Sonic.ts +59 -0
- package/src/dsp/SpeexResampler.ts +79 -0
- package/src/math/VectorMath.ts +812 -0
- package/src/nlp/ChineseSegmentation.ts +68 -0
- package/src/nlp/CompromiseNLP.ts +113 -0
- package/src/nlp/EspeakPhonemizer.ts +168 -0
- package/src/nlp/IPA.ts +139 -0
- package/src/nlp/JapaneseSegmentation.ts +53 -0
- package/src/nlp/Lexicon.ts +119 -0
- package/src/nlp/PhoneConversion.ts +508 -0
- package/src/nlp/Segmentation.ts +237 -0
- package/src/nlp/TextNormalizer.ts +160 -0
- package/src/recognition/AmazonTranscribeSTT.ts +112 -0
- package/src/recognition/AzureCognitiveServicesSTT.ts +76 -0
- package/src/recognition/GoogleCloudSTT.ts +92 -0
- package/src/recognition/SileroSTT.ts +173 -0
- package/src/recognition/VoskSTT.ts +112 -0
- package/src/recognition/WhisperSTT.ts +1518 -0
- package/src/server/Client.ts +297 -0
- package/src/server/Server.ts +178 -0
- package/src/server/ServerStarter.ts +12 -0
- package/src/server/Worker.ts +400 -0
- package/src/server/WorkerStarter.ts +38 -0
- package/src/speech-language-detection/SileroLanguageDetection.ts +105 -0
- package/src/subtitles/Subtitles.ts +478 -0
- package/src/synthesis/AwsPollyTTS.ts +78 -0
- package/src/synthesis/AzureCognitiveServicesTTS.ts +146 -0
- package/src/synthesis/CoquiServerTTS.ts +29 -0
- package/src/synthesis/ElevenLabsTTS.ts +104 -0
- package/src/synthesis/EspeakTTS.ts +552 -0
- package/src/synthesis/FliteTTS.ts +387 -0
- package/src/synthesis/GoogleCloudTTS.ts +112 -0
- package/src/synthesis/GoogleTranslateTTS.ts +210 -0
- package/src/synthesis/MicrosoftEdgeTTS.ts +298 -0
- package/src/synthesis/SamTTS.ts +30 -0
- package/src/synthesis/SapiTTS.ts +222 -0
- package/src/synthesis/StreamlabsPollyTTS.ts +114 -0
- package/src/synthesis/SvoxPicoTTS.ts +318 -0
- package/src/synthesis/VitsTTS.ts +734 -0
- package/src/tests/Test.ts +24 -0
- package/src/text-language-detection/FastTextLanguageDetection.ts +53 -0
- package/src/text-language-detection/TinyLDLanguageDetection.ts +16 -0
- package/src/typings/Fillers.d.ts +41 -0
- package/src/utilities/BinaryArrayConversion.ts +159 -0
- package/src/utilities/Compression.ts +91 -0
- package/src/utilities/FileDownloader.ts +201 -0
- package/src/utilities/FileSystem.ts +265 -0
- package/src/utilities/Hashing.ts +230 -0
- package/src/utilities/Locale.ts +119 -0
- package/src/utilities/Logger.ts +72 -0
- package/src/utilities/NdArrayUtilities.ts +31 -0
- package/src/utilities/ObjectUtilities.ts +169 -0
- package/src/utilities/OpenPromise.ts +13 -0
- package/src/utilities/PackageManager.ts +97 -0
- package/src/utilities/Queue.ts +17 -0
- package/src/utilities/RandomGenerator.ts +237 -0
- package/src/utilities/SignalChannel.ts +22 -0
- package/src/utilities/TarballMaker.ts +68 -0
- package/src/utilities/Timeline.ts +231 -0
- package/src/utilities/Timer.ts +93 -0
- package/src/utilities/Utilities.ts +574 -0
- package/src/utilities/WasmMemoryManager.ts +516 -0
- package/src/utilities/WebReader.ts +55 -0
- package/src/utilities/WikipediaReader.ts +41 -0
- package/src/voice-activity-detection/SileroVAD.ts +86 -0
- 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
|
+
}
|