echogarden 1.0.4 → 1.1.0

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (122) hide show
  1. package/README.md +26 -23
  2. package/data/schemas/options.json +177 -36
  3. package/dist/alignment/SpeechAlignment.d.ts +1 -1
  4. package/dist/alignment/SpeechAlignment.js +1 -1
  5. package/dist/alignment/SpeechAlignment.js.map +1 -1
  6. package/dist/api/API.d.ts +1 -0
  7. package/dist/api/API.js +1 -0
  8. package/dist/api/API.js.map +1 -1
  9. package/dist/api/APIOptions.d.ts +1 -0
  10. package/dist/api/Alignment.d.ts +3 -3
  11. package/dist/api/Alignment.js +5 -10
  12. package/dist/api/Alignment.js.map +1 -1
  13. package/dist/api/LanguageDetection.d.ts +5 -7
  14. package/dist/api/LanguageDetection.js +3 -2
  15. package/dist/api/LanguageDetection.js.map +1 -1
  16. package/dist/api/Recognition.d.ts +4 -5
  17. package/dist/api/Recognition.js +5 -8
  18. package/dist/api/Recognition.js.map +1 -1
  19. package/dist/api/SourceSeparation.d.ts +2 -0
  20. package/dist/api/SourceSeparation.js +4 -2
  21. package/dist/api/SourceSeparation.js.map +1 -1
  22. package/dist/api/Synthesis.d.ts +3 -1
  23. package/dist/api/Synthesis.js +9 -10
  24. package/dist/api/Synthesis.js.map +1 -1
  25. package/dist/api/Translation.d.ts +1 -1
  26. package/dist/api/Translation.js +4 -8
  27. package/dist/api/Translation.js.map +1 -1
  28. package/dist/api/TranslationAlignment.d.ts +31 -0
  29. package/dist/api/TranslationAlignment.js +121 -0
  30. package/dist/api/TranslationAlignment.js.map +1 -0
  31. package/dist/api/VoiceActivityDetection.d.ts +5 -1
  32. package/dist/api/VoiceActivityDetection.js +38 -2
  33. package/dist/api/VoiceActivityDetection.js.map +1 -1
  34. package/dist/audio/AudioPlayer.js +6 -1
  35. package/dist/audio/AudioPlayer.js.map +1 -1
  36. package/dist/cli/CLI.js +85 -0
  37. package/dist/cli/CLI.js.map +1 -1
  38. package/dist/dsp/FFT.js.map +1 -1
  39. package/dist/math/MedianFilter.d.ts +5 -0
  40. package/dist/math/MedianFilter.js +102 -0
  41. package/dist/math/MedianFilter.js.map +1 -0
  42. package/dist/math/VectorMath.d.ts +0 -2
  43. package/dist/math/VectorMath.js +1 -25
  44. package/dist/math/VectorMath.js.map +1 -1
  45. package/dist/recognition/OpenAICloudSTT.d.ts +1 -1
  46. package/dist/recognition/OpenAICloudSTT.js.map +1 -1
  47. package/dist/recognition/SileroSTT.d.ts +22 -1
  48. package/dist/recognition/SileroSTT.js +122 -95
  49. package/dist/recognition/SileroSTT.js.map +1 -1
  50. package/dist/recognition/WhisperCppSTT.js +1 -1
  51. package/dist/recognition/WhisperCppSTT.js.map +1 -1
  52. package/dist/recognition/WhisperSTT.d.ts +52 -19
  53. package/dist/recognition/WhisperSTT.js +645 -494
  54. package/dist/recognition/WhisperSTT.js.map +1 -1
  55. package/dist/server/Server.js.map +1 -1
  56. package/dist/source-separation/MDXNetSourceSeparation.d.ts +5 -3
  57. package/dist/source-separation/MDXNetSourceSeparation.js +26 -19
  58. package/dist/source-separation/MDXNetSourceSeparation.js.map +1 -1
  59. package/dist/speech-language-detection/SileroLanguageDetection.d.ts +15 -9
  60. package/dist/speech-language-detection/SileroLanguageDetection.js +23 -16
  61. package/dist/speech-language-detection/SileroLanguageDetection.js.map +1 -1
  62. package/dist/synthesis/EspeakTTS.js +4 -0
  63. package/dist/synthesis/EspeakTTS.js.map +1 -1
  64. package/dist/synthesis/GoogleCloudTTS.js.map +1 -1
  65. package/dist/synthesis/VitsTTS.d.ts +8 -6
  66. package/dist/synthesis/VitsTTS.js +36 -31
  67. package/dist/synthesis/VitsTTS.js.map +1 -1
  68. package/dist/tests/Test.js.map +1 -1
  69. package/dist/utilities/OnnxUtilities.d.ts +14 -0
  70. package/dist/utilities/OnnxUtilities.js +43 -0
  71. package/dist/utilities/OnnxUtilities.js.map +1 -0
  72. package/dist/utilities/Utilities.d.ts +4 -8
  73. package/dist/utilities/Utilities.js +35 -58
  74. package/dist/utilities/Utilities.js.map +1 -1
  75. package/dist/voice-activity-detection/SileroVAD.d.ts +5 -3
  76. package/dist/voice-activity-detection/SileroVAD.js +9 -11
  77. package/dist/voice-activity-detection/SileroVAD.js.map +1 -1
  78. package/docs/API.md +54 -34
  79. package/docs/CLI.md +25 -13
  80. package/docs/Contributing.md +4 -2
  81. package/docs/Engines.md +43 -32
  82. package/docs/Licenses.md +3 -4
  83. package/docs/Options.md +47 -11
  84. package/docs/Releases.md +4 -0
  85. package/docs/Server.md +8 -6
  86. package/docs/Tasklist.md +39 -52
  87. package/docs/Technical.md +1 -1
  88. package/package.json +8 -12
  89. package/src/alignment/SpeechAlignment.ts +1 -1
  90. package/src/api/API.ts +1 -0
  91. package/src/api/APIOptions.ts +1 -0
  92. package/src/api/Alignment.ts +10 -14
  93. package/src/api/LanguageDetection.ts +14 -10
  94. package/src/api/Recognition.ts +17 -10
  95. package/src/api/SourceSeparation.ts +7 -2
  96. package/src/api/Synthesis.ts +26 -11
  97. package/src/api/Translation.ts +14 -8
  98. package/src/api/TranslationAlignment.ts +213 -0
  99. package/src/api/VoiceActivityDetection.ts +66 -3
  100. package/src/audio/AudioPlayer.ts +6 -2
  101. package/src/cli/CLI.ts +121 -2
  102. package/src/dsp/FFT.ts +3 -0
  103. package/src/math/MedianFilter.ts +124 -0
  104. package/src/math/VectorMath.ts +1 -36
  105. package/src/recognition/OpenAICloudSTT.ts +27 -27
  106. package/src/recognition/SileroSTT.ts +149 -102
  107. package/src/recognition/WhisperCppSTT.ts +1 -1
  108. package/src/recognition/WhisperSTT.ts +961 -684
  109. package/src/server/Server.ts +1 -1
  110. package/src/source-separation/MDXNetSourceSeparation.ts +35 -19
  111. package/src/speech-language-detection/SileroLanguageDetection.ts +53 -33
  112. package/src/synthesis/EspeakTTS.ts +8 -0
  113. package/src/synthesis/GoogleCloudTTS.ts +12 -1
  114. package/src/synthesis/VitsTTS.ts +57 -46
  115. package/src/tests/Test.ts +1 -1
  116. package/src/utilities/OnnxUtilities.ts +68 -0
  117. package/src/utilities/Utilities.ts +38 -66
  118. package/src/voice-activity-detection/SileroVAD.ts +15 -15
  119. package/dist/utilities/NdArrayUtilities.d.ts +0 -3
  120. package/dist/utilities/NdArrayUtilities.js +0 -23
  121. package/dist/utilities/NdArrayUtilities.js.map +0 -1
  122. package/src/utilities/NdArrayUtilities.ts +0 -31
@@ -1,10 +1,10 @@
1
1
  import { Logger } from '../utilities/Logger.js';
2
2
  import { computeMelSpectogramUsingFilterbanks } from '../dsp/MelSpectogram.js';
3
- import { clip, getIntegerRange, getRepetitionScoreRelativeToFirstSubstring, getUTF32Chars, splitFloat32Array, yieldToEventLoop } from '../utilities/Utilities.js';
4
- import { indexOfMax, logOfVector, logSumExp, meanOfVector, medianFilter, softmax, stdDeviationOfVector } from '../math/VectorMath.js';
3
+ import { clip, containsInvalidCodepoint, getIntegerRange, getTokenRepetitionScore, splitFloat32Array, yieldToEventLoop } from '../utilities/Utilities.js';
4
+ import { indexOfMax, logOfVector, logSumExp, meanOfVector, softmax, stdDeviationOfVector } from '../math/VectorMath.js';
5
5
  import { alignDTWWindowed } from '../alignment/DTWSequenceAlignmentWindowed.js';
6
6
  import { extendDeep } from '../utilities/ObjectUtilities.js';
7
- import { getRawAudioDuration } from '../audio/AudioUtilities.js';
7
+ import { getRawAudioDuration, sliceRawAudio } from '../audio/AudioUtilities.js';
8
8
  import { readFile } from '../utilities/FileSystem.js';
9
9
  import path from 'path';
10
10
  import { getShortLanguageCode, languageCodeToName } from '../utilities/Locale.js';
@@ -12,8 +12,12 @@ import { loadPackage } from '../utilities/PackageManager.js';
12
12
  import chalk from 'chalk';
13
13
  import { XorShift32RNG } from '../utilities/RandomGenerator.js';
14
14
  import { detectSpeechLanguageByParts } from '../api/LanguageDetection.js';
15
- import { isPunctuation, isWhitespace } from '../nlp/Segmentation.js';
15
+ import { isPunctuation, isWhitespace, isWord, splitToSentences, splitToWords } from '../nlp/Segmentation.js';
16
+ import { medianOf5Filter } from '../math/MedianFilter.js';
17
+ import { getDeflateCompressionMetricsForString } from '../utilities/Compression.js';
18
+ import { getOnnxSessionOptions, makeOnnxLikeFloat32Tensor } from '../utilities/OnnxUtilities.js';
16
19
  export async function recognize(sourceRawAudio, modelName, modelDir, task, sourceLanguage, options) {
20
+ options = extendDeep(defaultWhisperOptions, options);
17
21
  if (sourceRawAudio.sampleRate != 16000) {
18
22
  throw new Error('Source audio must have a sampling rate of 16000');
19
23
  }
@@ -31,11 +35,14 @@ export async function recognize(sourceRawAudio, modelName, modelDir, task, sourc
31
35
  if (seed) {
32
36
  seed = Math.max(Math.floor(seed), 1) | 0;
33
37
  }
34
- const whisper = new Whisper(modelName, modelDir, seed);
38
+ const encoderProviders = options.encoderProvider ? [options.encoderProvider] : ['dml', 'cpu'];
39
+ const decoderProviders = options.decoderProvider ? [options.decoderProvider] : [];
40
+ const whisper = new Whisper(modelName, modelDir, encoderProviders, decoderProviders, seed);
35
41
  const result = await whisper.recognize(sourceRawAudio, task, sourceLanguage, options);
36
42
  return result;
37
43
  }
38
- export async function align(sourceRawAudio, referenceText, modelName, modelDir, sourceLanguage) {
44
+ export async function align(sourceRawAudio, transcript, modelName, modelDir, sourceLanguage, options) {
45
+ options = extendDeep(defaultWhisperAlignmentOptions, options);
39
46
  if (sourceRawAudio.sampleRate != 16000) {
40
47
  throw new Error('Source audio must have a sampling rate of 16000');
41
48
  }
@@ -46,46 +53,101 @@ export async function align(sourceRawAudio, referenceText, modelName, modelDir,
46
53
  if (isEnglishOnlyModel(modelName) && sourceLanguage != 'en') {
47
54
  throw new Error(`The model '${modelName}' can only be used with English inputs. However, the given source language was ${languageCodeToName(sourceLanguage)}.`);
48
55
  }
49
- const whisper = new Whisper(modelName, modelDir);
50
- const timeline = await whisper.align(sourceRawAudio, referenceText, sourceLanguage);
56
+ const encoderProviders = options.encoderProvider ? [options.encoderProvider] : ['dml', 'cpu'];
57
+ const decoderProviders = options.decoderProvider ? [options.decoderProvider] : [];
58
+ const whisper = new Whisper(modelName, modelDir, encoderProviders, decoderProviders);
59
+ const timeline = await whisper.align(sourceRawAudio, transcript, sourceLanguage, 'transcribe', options);
51
60
  return timeline;
52
61
  }
53
- export async function detectLanguage(sourceRawAudio, modelName, modelDir, temperature) {
62
+ export async function alignEnglishTranslation(sourceRawAudio, translatedTranscript, modelName, modelDir, sourceLanguage, options) {
63
+ options = extendDeep(defaultWhisperAlignmentOptions, options);
64
+ if (sourceRawAudio.sampleRate != 16000) {
65
+ throw new Error('Source audio must have a sampling rate of 16000');
66
+ }
67
+ sourceLanguage = getShortLanguageCode(sourceLanguage);
68
+ if (!(sourceLanguage in languageIdLookup)) {
69
+ throw new Error(`The source language ${languageCodeToName(sourceLanguage)} is not supported by the Whisper engine.`);
70
+ }
71
+ if (isEnglishOnlyModel(modelName)) {
72
+ throw new Error(`Translation alignment can only be done with multilingual models.`);
73
+ }
74
+ const encoderProviders = options.encoderProvider ? [options.encoderProvider] : ['dml', 'cpu'];
75
+ const decoderProviders = options.decoderProvider ? [options.decoderProvider] : [];
76
+ const whisper = new Whisper(modelName, modelDir, encoderProviders, decoderProviders);
77
+ const timeline = await whisper.align(sourceRawAudio, translatedTranscript, sourceLanguage, 'translate', options);
78
+ return timeline;
79
+ }
80
+ export async function detectLanguage(sourceRawAudio, modelName, modelDir, options) {
81
+ options = extendDeep(defaultWhisperLanguageDetectionOptions, options);
54
82
  if (sourceRawAudio.sampleRate != 16000) {
55
83
  throw new Error('Source audio must have a sampling rate of 16000');
56
84
  }
57
85
  if (!isMultilingualModel(modelName)) {
58
86
  throw new Error(`Language detection is only supported with multilingual models.`);
59
87
  }
60
- if (temperature < 0) {
88
+ if (options.temperature < 0) {
61
89
  throw new Error(`Temperature cannot be negative`);
62
90
  }
63
- const whisper = new Whisper(modelName, modelDir);
91
+ const encoderProviders = options.encoderProvider ? [options.encoderProvider] : ['dml', 'cpu'];
92
+ const decoderProviders = options.decoderProvider ? [options.decoderProvider] : [];
93
+ const whisper = new Whisper(modelName, modelDir, encoderProviders, decoderProviders);
64
94
  async function detectLanguageForPart(partAudio) {
65
95
  const audioFeatures = await whisper.encodeAudio(partAudio);
66
- const partResults = await whisper.detectLanguage(audioFeatures, temperature);
96
+ const partResults = await whisper.detectLanguage(audioFeatures, options.temperature);
67
97
  return partResults;
68
98
  }
69
99
  const results = await detectSpeechLanguageByParts(sourceRawAudio, detectLanguageForPart);
70
100
  results.sort((entry1, entry2) => entry2.probability - entry1.probability);
71
101
  return results;
72
102
  }
103
+ export async function detectVoiceActivity(sourceRawAudio, modelName, modelDir, options) {
104
+ options = extendDeep(defaultWhisperVADOptions, options);
105
+ if (sourceRawAudio.sampleRate != 16000) {
106
+ throw new Error('Source audio must have a sampling rate of 16000');
107
+ }
108
+ if (options.temperature < 0) {
109
+ throw new Error(`Temperature cannot be negative`);
110
+ }
111
+ const audioSamples = sourceRawAudio.audioChannels[0];
112
+ const partDuration = 5;
113
+ const maxSamplesCountForPart = sourceRawAudio.sampleRate * partDuration;
114
+ const encoderProviders = options.encoderProvider ? [options.encoderProvider] : ['dml', 'cpu'];
115
+ const decoderProviders = options.decoderProvider ? [options.decoderProvider] : [];
116
+ const whisper = new Whisper(modelName, modelDir, encoderProviders, decoderProviders);
117
+ const partProbabilities = [];
118
+ for (let sampleOffset = 0; sampleOffset < audioSamples.length; sampleOffset += maxSamplesCountForPart) {
119
+ const partSamples = sliceRawAudio(sourceRawAudio, sampleOffset, sampleOffset + maxSamplesCountForPart);
120
+ const samplesCountForPart = partSamples.audioChannels[0].length;
121
+ const startTime = sampleOffset / sourceRawAudio.sampleRate;
122
+ const endTime = (sampleOffset + samplesCountForPart) / sourceRawAudio.sampleRate;
123
+ const encodedPartSamples = await whisper.encodeAudio(partSamples);
124
+ const probabilityForPart = await whisper.detectVoiceActivity(encodedPartSamples, options.temperature);
125
+ partProbabilities.push({
126
+ type: 'segment',
127
+ text: '',
128
+ startTime,
129
+ endTime,
130
+ confidence: probabilityForPart,
131
+ });
132
+ }
133
+ return { partProbabilities };
134
+ }
73
135
  export class Whisper {
74
136
  modelName;
75
137
  modelDir;
138
+ encoderExecutionProviders;
139
+ decoderExecutionProviders;
76
140
  isMultiligualModel;
77
141
  audioEncoder;
78
142
  textDecoder;
79
143
  tiktoken;
80
- onnxOptions = {
81
- logSeverityLevel: 2,
82
- executionProviders: ['cpu']
83
- };
84
144
  tokenConfig;
85
145
  randomGen;
86
- constructor(modelName, modelDir, rngSeed = 461845907) {
146
+ constructor(modelName, modelDir, encoderExecutionProviders, decoderExecutionProviders, rngSeed = 461845907) {
87
147
  this.modelName = modelName;
88
148
  this.modelDir = modelDir;
149
+ this.encoderExecutionProviders = encoderExecutionProviders;
150
+ this.decoderExecutionProviders = decoderExecutionProviders;
89
151
  this.isMultiligualModel = isMultilingualModel(this.modelName);
90
152
  if (this.isMultiligualModel) {
91
153
  this.tokenConfig = {
@@ -119,87 +181,25 @@ export class Whisper {
119
181
  }
120
182
  this.randomGen = new XorShift32RNG(rngSeed);
121
183
  }
122
- async initializeIfNeeded() {
123
- await this.initializeTokenizerIfNeeded();
124
- await this.initializeEncoderSessionIfNeeded();
125
- await this.initializeDecoderSessionIfNeeded();
126
- }
127
- async initializeTokenizerIfNeeded() {
128
- if (this.tiktoken) {
129
- return;
130
- }
131
- const logger = new Logger();
132
- await logger.startAsync('Load tokenizer data');
133
- const tiktokenModulePackagePath = await loadPackage('whisper-tiktoken-data');
134
- const tiktokenDataFilePath = path.join(tiktokenModulePackagePath, this.isMultiligualModel ? 'multilingual.tiktoken' : 'gpt2.tiktoken');
135
- let tiktokenData = await readFile(tiktokenDataFilePath, { encoding: 'utf8' });
136
- const tokenConfig = this.tokenConfig;
137
- const metadataTokens = {
138
- [tokenConfig.endOfTextToken]: '[EndOfText]',
139
- [tokenConfig.startOfTextToken]: '[StartOfText]',
140
- [tokenConfig.translateTaskToken]: '[TranslateTask]',
141
- [tokenConfig.transcribeTaskToken]: '[TranscribeTask]',
142
- [tokenConfig.startOfPromptToken]: '[StartOfPrompt]',
143
- [tokenConfig.nonSpeechToken]: '[NonSpeech]',
144
- [tokenConfig.noTimestampsToken]: '[NoTimestamps]',
145
- };
146
- if (this.isMultiligualModel) {
147
- metadataTokens[50256] = '[Unused_50256]';
148
- metadataTokens[50360] = '[Unused_50360]';
149
- }
150
- const languageTokenCount = tokenConfig.languageTokensEnd - tokenConfig.languageTokensStart;
151
- for (let i = 0; i < languageTokenCount; i++) {
152
- const tokenIndex = this.tokenConfig.languageTokensStart + i;
153
- metadataTokens[tokenIndex] = `[Language_${i}]`;
154
- }
155
- const timestampTokensCount = 1501;
156
- for (let i = 0; i < timestampTokensCount; i++) {
157
- const tokenIndex = this.tokenConfig.timestampTokensStart + i;
158
- const tokenTime = this.timestampTokenToSeconds(tokenIndex);
159
- metadataTokens[tokenIndex] = `[Timestamp_${tokenTime.toFixed(2)}]`;
160
- }
161
- const inverseMetadataTokensLookup = {};
162
- for (const [key, value] of Object.entries(metadataTokens)) {
163
- inverseMetadataTokensLookup[value] = parseInt(key);
164
- }
165
- const patternString = `'s|'t|'re|'ve|'m|'ll|'d| ?\\p{L}+| ?\\p{N}+| ?[^\\s\\p{L}\\p{N}]+|\\s+(?!\\S)|\\s+`;
166
- const { Tiktoken } = await import('tiktoken/lite');
167
- this.tiktoken = new Tiktoken(tiktokenData, inverseMetadataTokensLookup, patternString);
168
- logger.end();
169
- }
170
- async initializeEncoderSessionIfNeeded() {
171
- if (this.audioEncoder) {
172
- return;
173
- }
174
- const logger = new Logger();
175
- await logger.startAsync(`Create encoder model inference session for model '${this.modelName}'`);
176
- const encoderFilePath = path.join(this.modelDir, 'encoder.onnx');
177
- const Onnx = await import('onnxruntime-node');
178
- this.audioEncoder = await Onnx.InferenceSession.create(encoderFilePath, this.onnxOptions);
179
- logger.end();
180
- }
181
- async initializeDecoderSessionIfNeeded() {
182
- if (this.textDecoder) {
183
- return;
184
- }
185
- const logger = new Logger();
186
- await logger.startAsync(`Create decoder model inference session for model '${this.modelName}'`);
187
- const decoderFilePath = path.join(this.modelDir, 'decoder.onnx');
188
- const Onnx = await import('onnxruntime-node');
189
- this.textDecoder = await Onnx.InferenceSession.create(decoderFilePath, this.onnxOptions);
190
- logger.end();
191
- }
192
- async recognize(rawAudio, task, language, options) {
184
+ async recognize(rawAudio, task, language, options, logitFilter) {
193
185
  await this.initializeIfNeeded();
194
186
  const logger = new Logger();
187
+ options = extendDeep(defaultWhisperOptions, options);
188
+ options.model = this.modelName;
195
189
  const audioSamples = rawAudio.audioChannels[0];
196
190
  const sampleRate = rawAudio.sampleRate;
197
191
  const prompt = options.prompt;
198
- const maxAudioSamplesPerPart = sampleRate * 30;
199
192
  const decodeTimestampTokens = options.decodeTimestampTokens;
193
+ const maxAudioSamplesPerPart = sampleRate * 30;
200
194
  let previousPartTextTokens = [];
201
195
  let timeline = [];
202
196
  let allDecodedTokens = [];
197
+ let wrappedLogitFilter;
198
+ if (logitFilter) {
199
+ wrappedLogitFilter = (logits, partDecodedTokens, isFirstPart, isFinalPart) => {
200
+ return logitFilter(logits, [...allDecodedTokens, ...partDecodedTokens], isFirstPart, isFinalPart);
201
+ };
202
+ }
203
203
  for (let audioOffset = 0; audioOffset < audioSamples.length;) {
204
204
  const segmentStartTime = audioOffset / sampleRate;
205
205
  await logger.startAsync(`\nPrepare audio part at time position ${segmentStartTime.toFixed(2)}`, undefined, chalk.magentaBright);
@@ -220,7 +220,7 @@ export class Whisper {
220
220
  }
221
221
  initialTokens = [...initialTokens, ...this.getTextStartTokens(language, task, !decodeTimestampTokens)];
222
222
  logger.end();
223
- let { decodedTokens: partTokens, crossAttentionQKs: partCrossAttentionQKs, decodedTokensConfidence } = await this.decodeTokens(audioPartFeatures, initialTokens, audioPartDuration, isFirstPart, isFinalPart, options);
223
+ let { decodedTokens: partTokens, decodedTokensConfidence: partTokensConfidence, decodedTokensCrossAttentionQKs: partCrossAttentionQKs, } = await this.decodeTokens(audioPartFeatures, initialTokens, audioPartDuration, isFirstPart, isFinalPart, options, wrappedLogitFilter);
224
224
  const lastToken = partTokens[partTokens.length - 1];
225
225
  const lastTokenIsTimestamp = this.isTimestampToken(lastToken);
226
226
  let audioEndOffset;
@@ -237,50 +237,96 @@ export class Whisper {
237
237
  if (partTokens.length != partCrossAttentionQKs.length) {
238
238
  throw new Error('Unexpected: partTokens.length != partCrossAttentionQKs.length');
239
239
  }
240
+ // Prepare tokens
240
241
  partTokens = partTokens.slice(initialTokens.length);
241
- //const compressionRatioForPart = (await getDeflateCompressionMetricsForString(this.tokensToText(partTokens))).ratio
242
+ partTokensConfidence = partTokensConfidence.slice(initialTokens.length);
242
243
  partCrossAttentionQKs = partCrossAttentionQKs.slice(initialTokens.length);
244
+ // Compute compression ratio for part
245
+ if (false) {
246
+ const compressionRatioForPart = (await getDeflateCompressionMetricsForString(this.tokensToText(partTokens))).ratio;
247
+ }
248
+ // Find alignment path
243
249
  const alignmentPath = await this.findAlignmentPathFromQKs(partCrossAttentionQKs, partTokens, 0, segmentFrameCount); //, alignmentHeadsIndexes[this.modelName])
244
- const partTimeline = await this.getTokenTimelineFromAlignmentPath(alignmentPath, partTokens, segmentStartTime, segmentEndTime, decodedTokensConfidence);
245
- audioOffset = audioEndOffset;
250
+ // Generate timeline from alignment path
251
+ const partTimeline = await this.getTokenTimelineFromAlignmentPath(alignmentPath, partTokens, segmentStartTime, segmentEndTime, partTokensConfidence);
246
252
  allDecodedTokens.push(...partTokens);
247
253
  timeline.push(...partTimeline);
248
254
  previousPartTextTokens = partTokens.filter(token => this.isTextToken(token));
255
+ audioOffset = audioEndOffset;
249
256
  logger.end();
250
257
  }
258
+ // Convert token timeline to word timeline
251
259
  timeline = this.tokenTimelineToWordTimeline(timeline, language);
260
+ // Convert tokens to transcript
252
261
  const transcript = this.tokensToText(allDecodedTokens).trim();
253
262
  logger.end();
254
263
  return { transcript, timeline };
255
264
  }
256
- async align(rawAudio, referenceText, language) {
257
- await this.initializeIfNeeded();
258
- const logger = new Logger();
259
- await logger.startAsync('Prepare for alignment');
260
- referenceText = referenceText.replaceAll(/\s+/g, ' ');
261
- const audioDuration = Math.min(getRawAudioDuration(rawAudio), 30);
262
- const audioFrameCount = this.secondsToFrame(audioDuration);
263
- const initialTokens = this.getTextStartTokens(language, 'transcribe', true);
265
+ async align(rawAudio, transcript, sourceLanguage, task, whisperAlignmentOptions) {
266
+ await this.initializeTokenizerIfNeeded();
267
+ whisperAlignmentOptions = extendDeep(defaultWhisperAlignmentOptions, whisperAlignmentOptions);
268
+ const shouldSplitToSentences = false;
269
+ let simplifiedTranscript = '';
270
+ if (shouldSplitToSentences) {
271
+ const sentences = splitToSentences(transcript, 'en');
272
+ for (const sentence of sentences) {
273
+ let sentenceWords = await splitToWords(sentence, 'en');
274
+ sentenceWords = sentenceWords.filter(word => isWord(word));
275
+ simplifiedTranscript += sentenceWords.join(' ');
276
+ simplifiedTranscript += ' ';
277
+ }
278
+ }
279
+ else {
280
+ let words = await splitToWords(transcript, 'en');
281
+ words = words.map(word => word.trim());
282
+ words = words.filter(word => isWord(word));
283
+ simplifiedTranscript = words.join(' ');
284
+ }
285
+ // Tokenize the transcript
286
+ const simplifiedTranscriptTokens = this.textToTokens(simplifiedTranscript);
287
+ // Initialize custom logit filter that allows only the transcript tokens to be decoded
288
+ // in order.
264
289
  const endOfTextToken = this.tokenConfig.endOfTextToken;
265
- let tokens = [...initialTokens, ...this.textToTokens(referenceText), endOfTextToken];
266
- logger.end();
267
- const audioFeatures = await this.encodeAudio(rawAudio);
268
- await logger.startAsync('Infer cross-attention QKs');
269
- let crossAttentionQKs = await this.inferCrossAttentionQKs(tokens, audioFeatures);
270
- tokens = tokens.slice(initialTokens.length, tokens.length - 1);
271
- crossAttentionQKs = crossAttentionQKs.slice(initialTokens.length, crossAttentionQKs.length - 1);
272
- await logger.startAsync('Extract word timeline');
273
- const alignmentPath = await this.findAlignmentPathFromQKs(crossAttentionQKs, tokens, 0, audioFrameCount); //, this.getAlignmentHeadIndexes())
274
- const tokenTimeline = await this.getTokenTimelineFromAlignmentPath(alignmentPath, tokens, 0, audioDuration);
275
- const wordTimeline = this.tokenTimelineToWordTimeline(tokenTimeline, language);
276
- logger.end();
277
- return wordTimeline;
290
+ const logitFilter = (logits, decodedTokens, isFirstPart, isFinalPart) => {
291
+ const decodedTextTokens = decodedTokens.filter(token => this.isTextToken(token));
292
+ const nextTokenToDecode = simplifiedTranscriptTokens[decodedTextTokens.length] ?? endOfTextToken;
293
+ const newLogits = logits.map((logit, index) => {
294
+ if (index === nextTokenToDecode) {
295
+ return logit;
296
+ }
297
+ // If it's the final part, the ent-of-text token logit is set to -Infinity.
298
+ // This will force to force all transcript tokens to be decoded even if the model doesn't
299
+ // recognize them.
300
+ if (!isFinalPart && index === endOfTextToken) {
301
+ return logit;
302
+ }
303
+ return -Infinity;
304
+ });
305
+ return newLogits;
306
+ };
307
+ // Set options for alignment
308
+ const options = {
309
+ model: this.modelName,
310
+ temperature: 0.0,
311
+ prompt: undefined,
312
+ topCandidateCount: 1,
313
+ punctuationThreshold: Infinity,
314
+ autoPromptParts: false,
315
+ maxTokensPerPart: Infinity,
316
+ suppressRepetition: false,
317
+ decodeTimestampTokens: true,
318
+ endTokenThreshold: whisperAlignmentOptions.endTokenThreshold,
319
+ includeEndTokenInCandidates: false,
320
+ seed: undefined,
321
+ };
322
+ // Recognize
323
+ const { timeline } = await this.recognize(rawAudio, task, sourceLanguage, options, logitFilter);
324
+ return timeline;
278
325
  }
279
326
  async detectLanguage(audioFeatures, temperature) {
280
327
  if (!this.isMultiligualModel) {
281
328
  throw new Error('Language detection is only supported with multilingual models');
282
329
  }
283
- await this.initializeTokenizerIfNeeded();
284
330
  await this.initializeDecoderSessionIfNeeded();
285
331
  // Prepare and run decoder
286
332
  const logger = new Logger();
@@ -317,37 +363,76 @@ export class Whisper {
317
363
  logger.end();
318
364
  return results;
319
365
  }
320
- async decodeTokens(audioFeatures, initialTokens, audioDuration, isFirstPart, isFinalPart, options) {
366
+ async detectVoiceActivity(audioFeatures, temperature) {
367
+ await this.initializeDecoderSessionIfNeeded();
368
+ // Prepare and run decoder
369
+ const logger = new Logger();
370
+ await logger.startAsync('Detect voice activity with Whisper model');
371
+ const sotToken = this.tokenConfig.startOfTextToken;
372
+ const initialTokens = [sotToken];
373
+ const offset = 0;
374
+ const Onnx = await import('onnxruntime-node');
375
+ const initialKvDimensions = this.getKvDimensions(1, initialTokens.length);
376
+ const kvCacheTensor = new Onnx.Tensor('float32', new Float32Array(initialKvDimensions[0] * initialKvDimensions[1] * initialKvDimensions[2] * initialKvDimensions[3]), initialKvDimensions);
377
+ const tokensTensor = new Onnx.Tensor('int64', new BigInt64Array(initialTokens.map(token => BigInt(token))), [1, initialTokens.length]);
378
+ const offsetTensor = new Onnx.Tensor('int64', new BigInt64Array([BigInt(offset)]), []);
379
+ const decoderInputs = {
380
+ tokens: tokensTensor,
381
+ audio_features: audioFeatures,
382
+ kv_cache: kvCacheTensor,
383
+ offset: offsetTensor
384
+ };
385
+ const decoderOutputs = await this.textDecoder.run(decoderInputs);
386
+ const logitsBuffer = decoderOutputs['logits'].data;
387
+ const tokenConfig = this.tokenConfig;
388
+ const logits = Array.from(logitsBuffer);
389
+ const probabilities = softmax(logits, temperature);
390
+ const noSpeechProbability = probabilities[tokenConfig.nonSpeechToken];
391
+ return 1.0 - noSpeechProbability;
392
+ }
393
+ // Decode tokens using the decoder model
394
+ async decodeTokens(audioFeatures, initialTokens, audioDuration, isFirstPart, isFinalPart, options, logitFilter) {
395
+ // Initialize
321
396
  await this.initializeTokenizerIfNeeded();
322
397
  await this.initializeDecoderSessionIfNeeded();
323
398
  const logger = new Logger();
324
- const allowedPunctuationMarks = this.getAllowedPunctuationMarks();
325
399
  await logger.startAsync('Decode text tokens with Whisper decoder model');
326
400
  options = extendDeep(defaultWhisperOptions, options);
327
401
  const Onnx = await import('onnxruntime-node');
402
+ // Get token information
328
403
  const endOfTextToken = this.tokenConfig.endOfTextToken;
329
404
  const timestampTokensStart = this.tokenConfig.timestampTokensStart;
330
- const suppressedTokens = new Set(this.getSuppressedTokens());
405
+ const suppressedTextTokens = this.getSuppressedTextTokens();
406
+ const suppressedMetadataTokens = this.getSuppressedMetadataTokens();
407
+ const allowedPunctuationMarks = this.getAllowedPunctuationMarks();
331
408
  const spaceToken = this.textToTokens(' ')[0];
332
- const maxDecodedTokenCount = options.maxTokensPerPart;
409
+ // Initialize variables for decoding loop
333
410
  let decodedTokens = initialTokens.slice();
334
411
  const initialKvDimensions = this.getKvDimensions(1, decodedTokens.length);
335
412
  let kvCacheTensor = new Onnx.Tensor('float32', new Float32Array(initialKvDimensions[0] * initialKvDimensions[1] * initialKvDimensions[2] * initialKvDimensions[3]), initialKvDimensions);
336
- let decodedTokensTimestampLogits = [new Array(1501)];
337
- let lastTimestampTokenIndex = -1;
338
- let timestampsSeenCount = 0;
339
- const decodedTokensConfidence = [];
413
+ let decodedTokensTimestampLogits = [];
414
+ let decodedTokensConfidence = [];
340
415
  let decodedTokensCrossAttentionQKs = [];
341
416
  for (let i = 0; i < decodedTokens.length; i++) {
417
+ decodedTokensTimestampLogits.push(new Array(1501));
418
+ decodedTokensConfidence.push(1.0);
342
419
  decodedTokensCrossAttentionQKs.push(undefined);
343
420
  }
421
+ let lastTimestampTokenIndex = -1;
422
+ let timestampTokenSeenCount = 0;
423
+ let bufferedTokensToPrint = [];
424
+ // Define method to add a token to output
425
+ function addToken(tokenToAdd, timestampLogits, confidence, crossAttentionQKs) {
426
+ decodedTokens.push(tokenToAdd);
427
+ decodedTokensTimestampLogits.push(timestampLogits);
428
+ decodedTokensConfidence.push(confidence);
429
+ decodedTokensCrossAttentionQKs.push(crossAttentionQKs);
430
+ }
344
431
  // Start decoding loop
345
- for (let decodedTokenCount = 0; decodedTokenCount < maxDecodedTokenCount; decodedTokenCount++) {
432
+ for (let decodedTokenCount = 0; decodedTokenCount < options.maxTokensPerPart; decodedTokenCount++) {
346
433
  const isInitialState = decodedTokens.length == initialTokens.length;
347
- const tokensToDecode = isInitialState ? decodedTokens : [decodedTokens[decodedTokens.length - 1]];
348
- const offset = isInitialState ? 0 : decodedTokens.length;
434
+ // If not in initial state, reshape KV Cache tensor to accomodate a new output token
349
435
  if (!isInitialState) {
350
- // Reshape KV Cache tensor
351
436
  const dims = kvCacheTensor.dims;
352
437
  const currentKvCacheGroups = splitFloat32Array(kvCacheTensor.data, dims[2] * dims[3]);
353
438
  const reshapedKvCacheTensor = new Onnx.Tensor('float32', new Float32Array(dims[0] * dims[1] * (decodedTokens.length) * dims[3]), [dims[0], dims[1], decodedTokens.length, dims[3]]);
@@ -357,200 +442,252 @@ export class Whisper {
357
442
  }
358
443
  kvCacheTensor = reshapedKvCacheTensor;
359
444
  }
360
- // Prepare and run decoder
445
+ // Prepare values for decoder
446
+ const tokensToDecode = isInitialState ? decodedTokens : [decodedTokens[decodedTokens.length - 1]];
447
+ const offset = isInitialState ? 0 : decodedTokens.length;
361
448
  const tokensTensor = new Onnx.Tensor('int64', new BigInt64Array(tokensToDecode.map(token => BigInt(token))), [1, tokensToDecode.length]);
362
449
  const offsetTensor = new Onnx.Tensor('int64', new BigInt64Array([BigInt(offset)]), []);
363
- const decoderInputs = { tokens: tokensTensor, audio_features: audioFeatures, kv_cache: kvCacheTensor, offset: offsetTensor };
450
+ const decoderInputs = {
451
+ tokens: tokensTensor,
452
+ audio_features: audioFeatures,
453
+ kv_cache: kvCacheTensor,
454
+ offset: offsetTensor
455
+ };
456
+ // Run decoder
364
457
  const decoderOutputs = await this.textDecoder.run(decoderInputs);
458
+ // Store results
365
459
  const logitsBuffer = decoderOutputs['logits'].data;
366
460
  kvCacheTensor = decoderOutputs['output_kv_cache'];
367
- // Compute logits
368
- const resultLogits = splitFloat32Array(logitsBuffer, logitsBuffer.length / decoderOutputs['logits'].dims[1]);
369
- const allTokenLogits = Array.from(resultLogits[resultLogits.length - 1]);
370
- const timestampTokenLogits = allTokenLogits.slice(timestampTokensStart);
371
- // Suppress tokens
372
- for (let logitIndex = 0; logitIndex < allTokenLogits.length; logitIndex++) {
373
- const isWrongTokenForInitialState = isInitialState &&
374
- (logitIndex === spaceToken || logitIndex === endOfTextToken);
375
- const isInSuppressedList = suppressedTokens.has(logitIndex);
376
- const shouldSuppressToken = isWrongTokenForInitialState || isInSuppressedList;
377
- if (shouldSuppressToken) {
378
- allTokenLogits[logitIndex] = -Infinity;
379
- }
461
+ const crossAttentionQKsForTokenOnnx = decoderOutputs['cross_attention_qks'];
462
+ const crossAttentionQKsForToken = makeOnnxLikeFloat32Tensor(crossAttentionQKsForTokenOnnx);
463
+ crossAttentionQKsForTokenOnnx.dispose();
464
+ // Get logits
465
+ const resultLogitsFloatArrays = splitFloat32Array(logitsBuffer, logitsBuffer.length / decoderOutputs['logits'].dims[1]);
466
+ const allTokenLogits = Array.from(resultLogitsFloatArrays[resultLogitsFloatArrays.length - 1]);
467
+ // Suppress metadata tokens in the suppression set
468
+ for (const suppressedTokenIndex of suppressedMetadataTokens) {
469
+ allTokenLogits[suppressedTokenIndex] = -Infinity;
380
470
  }
381
- // Derive token probabilities
382
- let bufferedTokensToPrint = [];
383
- // Add best token
384
- function addToken(tokenToAdd, timestampLogits, confidence) {
385
- decodedTokens.push(tokenToAdd);
386
- decodedTokensTimestampLogits.push(timestampLogits);
387
- decodedTokensCrossAttentionQKs.push(decoderOutputs['cross_attention_qks']);
388
- decodedTokensConfidence.push(confidence);
471
+ if (isInitialState) {
472
+ // If in initial state, suppress end-of-text token
473
+ allTokenLogits[endOfTextToken] = -Infinity;
389
474
  }
390
- let shouldDecodeNonTimestampToken = true;
391
- if (options.decodeTimestampTokens) {
475
+ const timestampTokenLogits = allTokenLogits.slice(timestampTokensStart);
476
+ const decodeTimestampTokenIfNeeded = () => {
477
+ // Try to decode a timestamp token, if needed
478
+ // If timestamp tokens is disabled in options, don't decode a timestamp
479
+ if (!options.decodeTimestampTokens) {
480
+ return false;
481
+ }
482
+ const previousTokenWasTimestamp = this.isTimestampToken(decodedTokens[decodedTokens.length - 1]);
483
+ const secondPreviousTokenWasTimestamp = this.isTimestampToken(decodedTokens[decodedTokens.length - 2]);
484
+ // If there are two successive timestamp tokens decoded, or the previous timestamp was the first token,
485
+ // don't decode a timestamp
486
+ if (previousTokenWasTimestamp &&
487
+ (decodedTokens.length === initialTokens.length + 1) || secondPreviousTokenWasTimestamp) {
488
+ return false;
489
+ }
490
+ // Derive token probabilities
392
491
  const probabilities = softmax(allTokenLogits, 1.0);
393
492
  const logProbabilities = logOfVector(probabilities);
394
493
  const nonTimestampTokenLogProbs = logProbabilities.slice(0, timestampTokensStart);
494
+ // Find highest non-timestamp token
395
495
  const indexOfMaxNonTimestampLogProb = indexOfMax(nonTimestampTokenLogProbs);
396
496
  const valueOfMaxNonTimestampLogProb = nonTimestampTokenLogProbs[indexOfMaxNonTimestampLogProb];
497
+ // Find highest timestamp token
397
498
  const timestampTokenLogProbs = logProbabilities.slice(timestampTokensStart);
398
499
  const indexOfMaxTimestampLogProb = indexOfMax(timestampTokenLogProbs);
500
+ // Compute the log of the sum of exponentials of the log probabilities
501
+ // of the timestamp tokens
399
502
  const logSumExpOfTimestampTokenLogProbs = logSumExp(timestampTokenLogProbs);
400
- const shouldDecodeTimestampToken = logSumExpOfTimestampTokenLogProbs > valueOfMaxNonTimestampLogProb;
401
- const previousTokenWasTimestamp = this.isTimestampToken(decodedTokens[decodedTokens.length - 1]);
402
- const secondPreviousTokenWasTimestamp = decodedTokens.length < 2 || this.isTimestampToken(decodedTokens[decodedTokens.length - 2]);
403
- if (shouldDecodeTimestampToken && !previousTokenWasTimestamp) {
404
- timestampsSeenCount += 1;
503
+ // If the sum isn't greater than the log probability of the highest non-timestamp token,
504
+ // don't decode a timestamp
505
+ if (logSumExpOfTimestampTokenLogProbs <= valueOfMaxNonTimestampLogProb) {
506
+ return false;
405
507
  }
406
- if (shouldDecodeTimestampToken || (previousTokenWasTimestamp && !secondPreviousTokenWasTimestamp)) {
407
- if (previousTokenWasTimestamp) {
408
- const previousToken = decodedTokens[decodedTokens.length - 1];
409
- const previousTokenTimestampLogits = decodedTokensTimestampLogits[decodedTokensTimestampLogits.length - 1];
410
- const previousTokenConfidence = decodedTokensConfidence[decodedTokensConfidence.length - 1];
411
- addToken(previousToken, previousTokenTimestampLogits, previousTokenConfidence);
412
- lastTimestampTokenIndex = decodedTokens.length;
413
- const previousTokenTimestamp = this.timestampTokenToSeconds(previousToken);
414
- if (previousTokenTimestamp >= audioDuration) {
415
- break;
416
- }
417
- }
418
- else {
419
- const timestampToken = timestampTokensStart + indexOfMaxTimestampLogProb;
420
- const confidence = probabilities[timestampToken];
421
- addToken(timestampToken, timestampTokenLogits, confidence);
422
- }
423
- shouldDecodeNonTimestampToken = false;
508
+ // Decode a timestamp token
509
+ timestampTokenSeenCount += 1;
510
+ if (previousTokenWasTimestamp) {
511
+ const previousToken = decodedTokens[decodedTokens.length - 1];
512
+ const previousTokenTimestampLogits = decodedTokensTimestampLogits[decodedTokensTimestampLogits.length - 1];
513
+ const previousTokenConfidence = decodedTokensConfidence[decodedTokensConfidence.length - 1];
514
+ addToken(previousToken, previousTokenTimestampLogits, previousTokenConfidence, crossAttentionQKsForToken);
515
+ lastTimestampTokenIndex = decodedTokens.length;
516
+ }
517
+ else {
518
+ const timestampToken = timestampTokensStart + indexOfMaxTimestampLogProb;
519
+ const confidence = probabilities[timestampToken];
520
+ addToken(timestampToken, timestampTokenLogits, confidence, crossAttentionQKsForToken);
424
521
  }
522
+ return true;
523
+ };
524
+ const timestampTokenDecoded = decodeTimestampTokenIfNeeded();
525
+ if (timestampTokenDecoded) {
526
+ await yieldToEventLoop();
527
+ continue;
425
528
  }
426
- if (shouldDecodeNonTimestampToken) {
427
- const topLogitCount = options.topCandidateCount;
428
- const nonTimestampTokenLogits = allTokenLogits.slice(0, timestampTokensStart);
429
- const sortedNonTimestampTokenLogitsWithIndexes = Array.from(nonTimestampTokenLogits).map((logit, index) => ({ token: index, logit }));
430
- sortedNonTimestampTokenLogitsWithIndexes.sort((a, b) => b.logit - a.logit);
431
- let topCandidates = sortedNonTimestampTokenLogitsWithIndexes.slice(0, topLogitCount)
432
- .map(entry => ({
433
- token: entry.token,
434
- logit: entry.logit,
435
- text: this.tokenToText(entry.token, true)
436
- }));
437
- //// Repetition suppression code
438
- if (options.suppressRepetition) {
439
- const topCandidatesRepetitionScores = topCandidates.map(entry => {
440
- const lastDecodedTextTokens = decodedTokens.filter(token => this.isTextToken(token)).reverse().slice(0, 20);
441
- const { maxScore } = getRepetitionScoreRelativeToFirstSubstring([entry.token, ...lastDecodedTextTokens]);
442
- return maxScore;
443
- });
444
- const thresholdRepetitionScore = 4;
445
- if (topCandidatesRepetitionScores.every(score => score >= thresholdRepetitionScore)) {
446
- const indexOfMaxScore = topCandidatesRepetitionScores.indexOf(Math.max(...topCandidatesRepetitionScores));
447
- topCandidates = [topCandidates[indexOfMaxScore]];
448
- }
449
- else {
450
- topCandidates = topCandidates.filter((candidate, index) => topCandidatesRepetitionScores[index] < thresholdRepetitionScore);
451
- }
529
+ // Decode a non-timestamp token
530
+ let nonTimestampTokenLogits = allTokenLogits.slice(0, timestampTokensStart);
531
+ let shouldDecodeEndfOfTextToken = false;
532
+ // If not in initial state, and the end-of-text token's probability is sufficiently higher than
533
+ // the second highest ranked token, then set to accept it
534
+ if (!isInitialState) {
535
+ const endOfTextTokenLogit = nonTimestampTokenLogits[endOfTextToken];
536
+ const otherTokensLogits = nonTimestampTokenLogits.slice();
537
+ otherTokensLogits[endOfTextToken] = -Infinity;
538
+ const indexOfMaximumOtherTokenLogit = indexOfMax(otherTokensLogits);
539
+ const maximumOtherTokenLogit = nonTimestampTokenLogits[indexOfMaximumOtherTokenLogit];
540
+ const endProbabilities = softmax([endOfTextTokenLogit, maximumOtherTokenLogit], 1.0);
541
+ if (endProbabilities[0] > options.endTokenThreshold) {
542
+ shouldDecodeEndfOfTextToken = true;
452
543
  }
453
- ////
454
- const topCandidateProbabilities = softmax(topCandidates.map(a => a.logit), options.temperature);
455
- //// Remove end-of-text token from candidates if its probability isn't high enough
456
- if (options.decodeTimestampTokens === false) {
457
- topCandidates = topCandidates.filter((candidate, index) => {
458
- if (candidate.token === endOfTextToken) {
459
- return topCandidateProbabilities[index] >= 0.9;
460
- }
461
- return true;
462
- });
544
+ }
545
+ if (logitFilter) {
546
+ // Apply custom logit filter function if given
547
+ nonTimestampTokenLogits = logitFilter(nonTimestampTokenLogits, decodedTokens, isFirstPart, isFinalPart);
548
+ // If the custom filter set the end-of-text token to be Infinity, or -Infinity,
549
+ // then override any previous decision and accept or reject it, respectively
550
+ if (nonTimestampTokenLogits[endOfTextToken] === Infinity) {
551
+ shouldDecodeEndfOfTextToken = true;
463
552
  }
464
- ////
465
- const rankOfPromisingPunctuationToken = topCandidates.findIndex((entry, index) => {
466
- const tokenText = this.tokenToText(entry.token).trim();
467
- const isPunctuationToken = allowedPunctuationMarks.includes(tokenText);
468
- if (!isPunctuationToken) {
469
- return false;
470
- }
471
- const tokenProb = topCandidateProbabilities[index];
472
- return tokenProb >= options.punctuationThreshold;
473
- });
474
- let rankOfSpaceToken = topCandidates.findIndex(candidate => candidate.token === spaceToken);
475
- if (rankOfSpaceToken < 0) {
476
- rankOfSpaceToken = Infinity;
553
+ else if (nonTimestampTokenLogits[endOfTextToken] === -Infinity) {
554
+ shouldDecodeEndfOfTextToken = false;
477
555
  }
478
- let chosenCandidateRank;
479
- if (rankOfPromisingPunctuationToken >= 0 &&
480
- rankOfPromisingPunctuationToken < rankOfSpaceToken) {
481
- chosenCandidateRank = rankOfPromisingPunctuationToken;
556
+ // If filter caused all word token logits to be -Infinity, then there is no
557
+ // other token to decode. Fall back to accept end-of-text
558
+ if (nonTimestampTokenLogits.slice(0, endOfTextToken).every(logit => logit === -Infinity)) {
559
+ shouldDecodeEndfOfTextToken = true;
482
560
  }
483
- else {
484
- chosenCandidateRank = this.randomGen.selectRandomIndexFromDistribution(topCandidateProbabilities);
561
+ }
562
+ else {
563
+ // Otherwise, suppress text tokens in the suppression set
564
+ for (const suppressedTokenIndex of suppressedTextTokens) {
565
+ nonTimestampTokenLogits[suppressedTokenIndex] = -Infinity;
485
566
  }
486
- const chosenToken = topCandidates[chosenCandidateRank].token;
487
- if (this.isTextToken(chosenToken)) {
488
- bufferedTokensToPrint.push(chosenToken);
489
- let textToPrint = this.tokensToText(bufferedTokensToPrint);
490
- if (textToPrint.codePointAt(0) !== 65533) {
491
- if (isFirstPart && decodedTokens.every(token => this.isMetadataToken(token))) {
492
- textToPrint = textToPrint.trimStart();
493
- }
494
- logger.write(textToPrint);
495
- bufferedTokensToPrint = [];
567
+ // Suppress space token if at initial state
568
+ if (isInitialState) {
569
+ nonTimestampTokenLogits[spaceToken] = -Infinity;
570
+ }
571
+ }
572
+ // If end-of-text token should be decoded, then add it and break
573
+ // out of the loop
574
+ if (shouldDecodeEndfOfTextToken) {
575
+ addToken(endOfTextToken, timestampTokenLogits, 1.0, crossAttentionQKsForToken);
576
+ break;
577
+ }
578
+ // Suppress end-of-text if it shouldn't be included in candidates
579
+ if (!options.includeEndTokenInCandidates) {
580
+ nonTimestampTokenLogits[endOfTextToken] = -Infinity;
581
+ }
582
+ // Find top candidates
583
+ const sortedNonTimestampLogitsWithIndexes = Array.from(nonTimestampTokenLogits).map((logit, index) => ({ token: index, logit }));
584
+ sortedNonTimestampLogitsWithIndexes.sort((a, b) => b.logit - a.logit);
585
+ let topCandidates = sortedNonTimestampLogitsWithIndexes.slice(0, options.topCandidateCount)
586
+ .map(entry => ({
587
+ token: entry.token,
588
+ logit: entry.logit,
589
+ text: this.tokenToText(entry.token, true)
590
+ }));
591
+ // Apply repetition suppression if enabled
592
+ if (options.suppressRepetition) {
593
+ // Using some hardcoded constants, for now
594
+ const tokenWindowSize = 30;
595
+ const thresholdMatchLength = 6;
596
+ const thresholdCycleRepetition = 2.0;
597
+ const filteredCandidates = [];
598
+ for (const candidate of topCandidates) {
599
+ const lastDecodedTextTokens = decodedTokens
600
+ .filter(token => this.isTextToken(token))
601
+ .reverse()
602
+ .slice(0, tokenWindowSize);
603
+ const { longestMatch, longestCycleRepetition } = getTokenRepetitionScore([candidate.token, ...lastDecodedTextTokens]);
604
+ if (longestMatch >= thresholdMatchLength || longestCycleRepetition >= thresholdCycleRepetition) {
605
+ continue;
496
606
  }
607
+ filteredCandidates.push(candidate);
608
+ }
609
+ // If all candidates have been filtered out, accept an end-of-text token
610
+ if (filteredCandidates.length === 0) {
611
+ filteredCandidates.push({
612
+ token: endOfTextToken,
613
+ logit: Infinity,
614
+ text: this.tokenToText(endOfTextToken, true)
615
+ });
616
+ }
617
+ topCandidates = filteredCandidates;
618
+ }
619
+ // Compute top candidate probabilities
620
+ const topCandidateProbabilities = softmax(topCandidates.map(a => a.logit), options.temperature);
621
+ // Find highest ranking punctuation token
622
+ const rankOfPromisingPunctuationToken = topCandidates.findIndex((entry, index) => {
623
+ const tokenText = this.tokenToText(entry.token).trim();
624
+ const isPunctuationToken = allowedPunctuationMarks.includes(tokenText);
625
+ if (!isPunctuationToken) {
626
+ return false;
497
627
  }
498
- const confidence = topCandidateProbabilities[chosenCandidateRank];
499
- addToken(chosenToken, timestampTokenLogits, confidence);
500
- if (chosenToken === endOfTextToken) {
501
- break;
628
+ const tokenProb = topCandidateProbabilities[index];
629
+ return tokenProb >= options.punctuationThreshold;
630
+ });
631
+ // Find rank of space token
632
+ let rankOfSpaceToken = topCandidates.findIndex(candidate => candidate.token === spaceToken);
633
+ if (rankOfSpaceToken < 0) {
634
+ rankOfSpaceToken = Infinity;
635
+ }
636
+ // Choose token
637
+ let chosenCandidateRank;
638
+ // Select a high-ranking punctuation token if found, and it has
639
+ // a rank higher than the space token,
640
+ if (rankOfPromisingPunctuationToken >= 0 &&
641
+ rankOfPromisingPunctuationToken < rankOfSpaceToken) {
642
+ chosenCandidateRank = rankOfPromisingPunctuationToken;
643
+ }
644
+ else {
645
+ // Otherwise, select randomly from top k distribution
646
+ chosenCandidateRank = this.randomGen.selectRandomIndexFromDistribution(topCandidateProbabilities);
647
+ }
648
+ // Add chosen token
649
+ const chosenToken = topCandidates[chosenCandidateRank].token;
650
+ const chosenTokenConfidence = topCandidateProbabilities[chosenCandidateRank];
651
+ addToken(chosenToken, timestampTokenLogits, chosenTokenConfidence, crossAttentionQKsForToken);
652
+ // If chosen token is the end-of-text token, break
653
+ if (chosenToken === endOfTextToken) {
654
+ break;
655
+ }
656
+ // Print token if needed
657
+ if (this.isTextToken(chosenToken)) {
658
+ bufferedTokensToPrint.push(chosenToken);
659
+ let textToPrint = this.tokensToText(bufferedTokensToPrint);
660
+ // If the decoded text is valid, print it
661
+ if (!containsInvalidCodepoint(textToPrint)) {
662
+ if (isFirstPart && decodedTokens.every(token => this.isMetadataToken(token))) {
663
+ textToPrint = textToPrint.trimStart();
664
+ }
665
+ logger.write(textToPrint);
666
+ bufferedTokensToPrint = [];
502
667
  }
503
668
  }
504
669
  await yieldToEventLoop();
505
670
  }
506
- if (timestampsSeenCount >= 2 && !isFinalPart) {
507
- decodedTokens = decodedTokens.slice(0, lastTimestampTokenIndex);
508
- decodedTokensTimestampLogits = decodedTokensTimestampLogits.slice(0, lastTimestampTokenIndex);
509
- decodedTokensCrossAttentionQKs = decodedTokensCrossAttentionQKs.slice(0, lastTimestampTokenIndex);
671
+ // If at least two timestamp tokens were decoded and it's not the final part,
672
+ // truncate up to the last timestamp token
673
+ if (timestampTokenSeenCount >= 2 && !isFinalPart) {
674
+ const sliceEndTokenIndex = lastTimestampTokenIndex;
675
+ decodedTokens = decodedTokens.slice(0, sliceEndTokenIndex);
676
+ decodedTokensTimestampLogits = decodedTokensTimestampLogits.slice(0, sliceEndTokenIndex);
677
+ decodedTokensCrossAttentionQKs = decodedTokensCrossAttentionQKs.slice(0, sliceEndTokenIndex);
678
+ decodedTokensConfidence = decodedTokensConfidence.slice(0, sliceEndTokenIndex);
510
679
  }
511
680
  logger.write('\n');
512
681
  logger.end();
513
- // Return the tokens
682
+ // Return the decoded tokens
514
683
  return {
515
684
  decodedTokens,
516
685
  decodedTokensTimestampLogits,
517
- crossAttentionQKs: decodedTokensCrossAttentionQKs,
518
- decodedTokensConfidence
686
+ decodedTokensConfidence,
687
+ decodedTokensCrossAttentionQKs,
519
688
  };
520
689
  }
521
- async inferCrossAttentionQKs(tokens, audioFeatures) {
522
- const offset = 0;
523
- const Onnx = await import('onnxruntime-node');
524
- const tokensTensor = new Onnx.Tensor('int64', new BigInt64Array(tokens.map(token => BigInt(token))), [1, tokens.length]);
525
- const offsetTensor = new Onnx.Tensor('int64', new BigInt64Array([BigInt(offset)]), []);
526
- const initialKvDimensions = this.getKvDimensions(1, tokens.length);
527
- const kvCacheTensor = new Onnx.Tensor('float32', new Float32Array(initialKvDimensions[0] * initialKvDimensions[1] * initialKvDimensions[2] * initialKvDimensions[3]), initialKvDimensions);
528
- const decoderInputs = { tokens: tokensTensor, audio_features: audioFeatures, kv_cache: kvCacheTensor, offset: offsetTensor };
529
- const decoderOutputs = await this.textDecoder.run(decoderInputs);
530
- const crossAttentionQKsTensor = decoderOutputs['cross_attention_qks'];
531
- const tensorShape = crossAttentionQKsTensor.dims.slice();
532
- const ndarray = (await import('ndarray')).default;
533
- let qkArray = ndarray(crossAttentionQKsTensor.data, crossAttentionQKsTensor.dims.slice());
534
- qkArray = qkArray.transpose(3, 0, 1, 2, 4);
535
- const tokenCrossAttentionQKsTensors = [];
536
- for (let i0 = 0; i0 < qkArray.shape[0]; i0++) {
537
- const dataForToken = [];
538
- for (let i1 = 0; i1 < qkArray.shape[1]; i1++) {
539
- for (let i2 = 0; i2 < qkArray.shape[2]; i2++) {
540
- for (let i3 = 0; i3 < qkArray.shape[3]; i3++) {
541
- for (let i4 = 0; i4 < qkArray.shape[4]; i4++) {
542
- dataForToken.push(qkArray.get(i0, i1, i2, i3, i4));
543
- }
544
- }
545
- }
546
- }
547
- const newTensorShape = tensorShape.slice();
548
- newTensorShape[3] = 1;
549
- const newTensor = new Onnx.Tensor('float32', dataForToken, newTensorShape);
550
- tokenCrossAttentionQKsTensors.push(newTensor);
551
- }
552
- return tokenCrossAttentionQKsTensors;
553
- }
690
+ // Encode audio using the encoder model
554
691
  async encodeAudio(rawAudio) {
555
692
  await this.initializeEncoderSessionIfNeeded();
556
693
  const Onnx = await import('onnxruntime-node');
@@ -562,9 +699,12 @@ export class Whisper {
562
699
  const filterbankCount = 80;
563
700
  const maxAudioSamples = sampleRate * 30;
564
701
  const maxAudioFrames = 3000;
702
+ if (audioSamples.length > maxAudioSamples) {
703
+ throw new Error(`Audio part is longer than 30 seconds`);
704
+ }
565
705
  await logger.startAsync('Extract mel spectogram from audio part');
566
706
  const paddedAudioSamples = new Float32Array(maxAudioSamples);
567
- paddedAudioSamples.set(audioSamples.subarray(0, maxAudioSamples), 0);
707
+ paddedAudioSamples.set(audioSamples, 0);
568
708
  const rawAudioPart = { audioChannels: [paddedAudioSamples], sampleRate };
569
709
  const { melSpectogram } = await computeMelSpectogramUsingFilterbanks(rawAudioPart, fftOrder, fftOrder, hopLength, filterbanks);
570
710
  await logger.startAsync('Normalize mel spectogram');
@@ -592,99 +732,6 @@ export class Whisper {
592
732
  logger.end();
593
733
  return encodedAudioFeatures;
594
734
  }
595
- addSegmentsToTimeline(timeline, tokens, initialTimeOffset, audioDuration) {
596
- const timestampTokensStart = this.tokenConfig.timestampTokensStart;
597
- for (let i = 0; i < tokens.length; i++) {
598
- const token = tokens[i];
599
- if (token == this.tokenConfig.startOfTextToken || token == this.tokenConfig.endOfTextToken) {
600
- continue;
601
- }
602
- const tokenIsTimestamp = token >= timestampTokensStart;
603
- const previousTokenWasTimestamp = tokens.length > 1 && tokens[i - 1] >= timestampTokensStart;
604
- if (tokenIsTimestamp) {
605
- if (previousTokenWasTimestamp) {
606
- continue;
607
- }
608
- let startTime = initialTimeOffset + this.timestampTokenToSeconds(token);
609
- startTime = Math.min(startTime, audioDuration);
610
- if (timeline.length > 0) {
611
- timeline[timeline.length - 1].endTime = startTime;
612
- }
613
- timeline.push({
614
- type: 'segment',
615
- text: '',
616
- startTime,
617
- endTime: -1,
618
- });
619
- }
620
- else {
621
- if (timeline.length == 0) {
622
- timeline.push({
623
- type: 'segment',
624
- text: '',
625
- startTime: initialTimeOffset,
626
- endTime: -1,
627
- });
628
- }
629
- const tokenText = this.tokenToText(token);
630
- timeline[timeline.length - 1].text += tokenText;
631
- }
632
- }
633
- }
634
- async addWordsToTimeline(timeline, tokens, rawAudio, crossAttentionQKs, initialAudioTimeOffset, duration) {
635
- let segmentStartTime = 0;
636
- let segmentTokens = [];
637
- let segmentCrossAttentionQKs = [];
638
- for (let tokenIndex = 0; tokenIndex < tokens.length; tokenIndex++) {
639
- const token = tokens[tokenIndex];
640
- const tokenCrossAttentionQKs = crossAttentionQKs[tokenIndex];
641
- const segmentTokensWithoutTimestamps = segmentTokens.filter(token => this.isNonTimestampToken(token));
642
- const isTimestamp = this.isTimestampToken(token);
643
- if (isTimestamp || tokenIndex == tokens.length - 1) {
644
- let tokenTime;
645
- if (isTimestamp) {
646
- tokenTime = this.timestampTokenToSeconds(token);
647
- }
648
- else {
649
- tokenTime = duration;
650
- }
651
- if (segmentTokensWithoutTimestamps.length > 0) {
652
- const segmentEndTime = tokenTime;
653
- const segmentStartFrame = this.secondsToFrame(segmentStartTime);
654
- let segmentEndFrame = this.secondsToFrame(segmentEndTime);
655
- if (segmentStartFrame == segmentEndFrame) {
656
- segmentEndFrame += 1;
657
- }
658
- const segmentFrameCount = segmentEndFrame - segmentStartFrame;
659
- const reinferCrossAttentionQKs = true;
660
- if (reinferCrossAttentionQKs) {
661
- const initialTokens = this.getTextStartTokens('en', 'transcribe');
662
- const tokensToDecode = [...initialTokens, ...segmentTokensWithoutTimestamps];
663
- //const segmentAudioFeaturesBuffer = audioFeatures.data.slice(segmentStartFrame * audioFeatures.dims[2], segmentEndFrame * audioFeatures.dims[2])
664
- //const segmentAudioFeatures = new Onnx.Tensor('float32', segmentAudioFeaturesBuffer, [1, segmentFrameCount, audioFeatures.dims[2]])
665
- const segmentAudioSamples = rawAudio.audioChannels[0].slice(Math.floor(segmentStartTime * rawAudio.sampleRate), Math.floor(segmentEndTime * rawAudio.sampleRate));
666
- const segmentRawAudio = { audioChannels: [segmentAudioSamples], sampleRate: rawAudio.sampleRate };
667
- const segmentAudioFeatures = await this.encodeAudio(segmentRawAudio);
668
- const reinferredCrossAttentionQKs = await this.inferCrossAttentionQKs(tokensToDecode, segmentAudioFeatures);
669
- reinferredCrossAttentionQKs.slice(initialTokens.length);
670
- const alignmentPath = await this.findAlignmentPathFromQKs(reinferredCrossAttentionQKs, tokensToDecode, 0, segmentFrameCount); //, alignmentHeadsIndexes[modelName])
671
- const tokenTimeline = await this.getTokenTimelineFromAlignmentPath(alignmentPath, segmentTokensWithoutTimestamps, initialAudioTimeOffset + segmentStartTime, initialAudioTimeOffset + segmentEndTime);
672
- timeline.push(...tokenTimeline);
673
- }
674
- else {
675
- const alignmentPath = await this.findAlignmentPathFromQKs(segmentCrossAttentionQKs, segmentTokens, segmentStartFrame, segmentEndFrame); //, alignmentHeadsIndexes[modelName])
676
- const tokenTimeline = await this.getTokenTimelineFromAlignmentPath(alignmentPath, segmentTokens, initialAudioTimeOffset, initialAudioTimeOffset + segmentEndTime);
677
- timeline.push(...tokenTimeline);
678
- }
679
- }
680
- segmentStartTime = tokenTime;
681
- segmentTokens = [];
682
- segmentCrossAttentionQKs = [];
683
- }
684
- segmentTokens.push(token);
685
- segmentCrossAttentionQKs.push(tokenCrossAttentionQKs);
686
- }
687
- }
688
735
  tokenTimelineToWordTimeline(tokenTimeline, language) {
689
736
  function isSeparatorCharacter(char) {
690
737
  const nonSeparatingPunctuation = [`'`, `-`, `.`, `·`, `•`];
@@ -699,6 +746,9 @@ export class Whisper {
699
746
  function endsWithSeparatorCharacter(text) {
700
747
  return isSeparatorCharacter(text[text.length - 1]);
701
748
  }
749
+ if (language != 'zh' && language != 'ja') {
750
+ tokenTimeline = tokenTimeline.filter(entry => this.isTextToken(entry.id));
751
+ }
702
752
  const resultTimeline = [];
703
753
  let groups = [];
704
754
  for (let tokenIndex = 0; tokenIndex < tokenTimeline.length; tokenIndex++) {
@@ -716,21 +766,23 @@ export class Whisper {
716
766
  groups[groups.length - 1].push(entry);
717
767
  }
718
768
  }
719
- const newGroups = [];
720
- for (let groupIndex = 0; groupIndex < groups.length; groupIndex++) {
721
- const group = groups[groupIndex];
722
- const nextGroup = groups[groupIndex + 1];
723
- if (group.length > 1 &&
724
- group[group.length - 1].text === '.' &&
725
- (!nextGroup || [' ', '['].includes(nextGroup[0].text[0]))) {
726
- newGroups.push(group.slice(0, group.length - 1));
727
- newGroups.push(group.slice(group.length - 1));
728
- }
729
- else {
730
- newGroups.push(group);
769
+ {
770
+ const splitGroups = [];
771
+ for (let groupIndex = 0; groupIndex < groups.length; groupIndex++) {
772
+ const group = groups[groupIndex];
773
+ const nextGroup = groups[groupIndex + 1];
774
+ if (group.length > 1 &&
775
+ group[group.length - 1].text === '.' &&
776
+ (!nextGroup || [' ', '['].includes(nextGroup[0].text[0]))) {
777
+ splitGroups.push(group.slice(0, group.length - 1));
778
+ splitGroups.push(group.slice(group.length - 1));
779
+ }
780
+ else {
781
+ splitGroups.push(group);
782
+ }
731
783
  }
784
+ groups = splitGroups;
732
785
  }
733
- groups = newGroups;
734
786
  for (const group of groups) {
735
787
  let groupText = this.tokensToText(group.map(entry => entry.id));
736
788
  if (groupText === '') {
@@ -819,35 +871,35 @@ export class Whisper {
819
871
  const applySoftmax = true;
820
872
  const normalize = true;
821
873
  const applyMedianFilter = true;
822
- const fixateTimestampTokens = false;
874
+ const anchorTimestampTokens = false;
823
875
  const softmaxTemperature = 1.0;
824
- const medianFilterWidth = 7;
876
+ // Apply softmax to each token's frames, if enabled
825
877
  if (applySoftmax) {
826
- // Apply softmax to each token's frames
827
878
  for (const head of attentionHeads) {
828
879
  for (let tokenIndex = 0; tokenIndex < tokenCount; tokenIndex++) {
829
880
  head[tokenIndex] = softmax(head[tokenIndex], softmaxTemperature);
830
881
  }
831
882
  }
832
883
  }
884
+ // Normalize all weights in each individual head, if enabled
833
885
  if (normalize) {
834
- // Normalize all weights in each individual head
835
886
  for (const head of attentionHeads) {
836
887
  const allWeightsForHead = head.flatMap(tokenFrames => tokenFrames);
837
- const meanOfAllWeights = meanOfVector(allWeightsForHead);
838
- const stdDeviationOfAllWeights = stdDeviationOfVector(allWeightsForHead) + 1e-10;
888
+ const meanOfAllWeightsForHead = meanOfVector(allWeightsForHead);
889
+ const stdDeviationOfAllWeightsForHead = stdDeviationOfVector(allWeightsForHead, 'population', meanOfAllWeightsForHead) + 1e-10;
890
+ const stdDeviationReciprocal = 1.0 / (stdDeviationOfAllWeightsForHead + 1e-10);
839
891
  for (const tokenFrames of head) {
840
892
  for (let frameIndex = 0; frameIndex < tokenFrames.length; frameIndex++) {
841
- tokenFrames[frameIndex] = (tokenFrames[frameIndex] - meanOfAllWeights) / stdDeviationOfAllWeights;
893
+ tokenFrames[frameIndex] = (tokenFrames[frameIndex] - meanOfAllWeightsForHead) * stdDeviationReciprocal;
842
894
  }
843
895
  }
844
896
  }
845
897
  }
898
+ // Apply median filter to each token's frames, if enabled
846
899
  if (applyMedianFilter) {
847
- // Apply median filter to each token's frames
848
900
  for (const head of attentionHeads) {
849
901
  for (let tokenIndex = 0; tokenIndex < tokenCount; tokenIndex++) {
850
- head[tokenIndex] = medianFilter(head[tokenIndex], medianFilterWidth);
902
+ head[tokenIndex] = medianOf5Filter(head[tokenIndex]);
851
903
  }
852
904
  }
853
905
  }
@@ -867,8 +919,8 @@ export class Whisper {
867
919
  frameMeansForToken[tokenIndex][frameIndex] = frameMean;
868
920
  }
869
921
  }
870
- if (fixateTimestampTokens) {
871
- // Fixate timestamp tokens to the original ones detected
922
+ // Anchor timestamp tokens timestamps to their original values, if enabled
923
+ if (anchorTimestampTokens) {
872
924
  const timestampTokensStart = this.tokenConfig.timestampTokensStart;
873
925
  for (let tokenIndex = 0; tokenIndex < tokens.length; tokenIndex++) {
874
926
  const token = tokens[tokenIndex];
@@ -880,14 +932,86 @@ export class Whisper {
880
932
  }
881
933
  }
882
934
  // Perform DTW
883
- const tokenIndexes = [...Array(tokenCount).keys()];
884
- const frameIndexes = [...Array(segmentFrameCount).keys()];
935
+ const tokenIndexes = getIntegerRange(0, tokenCount);
936
+ const frameIndexes = getIntegerRange(0, segmentFrameCount);
885
937
  let { path } = alignDTWWindowed(tokenIndexes, frameIndexes, (tokenIndex, frameIndex) => {
886
938
  return -frameMeansForToken[tokenIndex][frameIndex];
887
939
  }, segmentFrameCount);
888
940
  path = path.map(entry => ({ source: entry.source, dest: segmentStartFrame + entry.dest }));
889
941
  return path;
890
942
  }
943
+ async initializeIfNeeded() {
944
+ await this.initializeTokenizerIfNeeded();
945
+ await this.initializeEncoderSessionIfNeeded();
946
+ await this.initializeDecoderSessionIfNeeded();
947
+ }
948
+ async initializeTokenizerIfNeeded() {
949
+ if (this.tiktoken) {
950
+ return;
951
+ }
952
+ const logger = new Logger();
953
+ await logger.startAsync('Load tokenizer data');
954
+ const tiktokenModulePackagePath = await loadPackage('whisper-tiktoken-data');
955
+ const tiktokenDataFilePath = path.join(tiktokenModulePackagePath, this.isMultiligualModel ? 'multilingual.tiktoken' : 'gpt2.tiktoken');
956
+ let tiktokenData = await readFile(tiktokenDataFilePath, { encoding: 'utf8' });
957
+ const tokenConfig = this.tokenConfig;
958
+ const metadataTokens = {
959
+ [tokenConfig.endOfTextToken]: '[EndOfText]',
960
+ [tokenConfig.startOfTextToken]: '[StartOfText]',
961
+ [tokenConfig.translateTaskToken]: '[TranslateTask]',
962
+ [tokenConfig.transcribeTaskToken]: '[TranscribeTask]',
963
+ [tokenConfig.startOfPromptToken]: '[StartOfPrompt]',
964
+ [tokenConfig.nonSpeechToken]: '[NonSpeech]',
965
+ [tokenConfig.noTimestampsToken]: '[NoTimestamps]',
966
+ };
967
+ if (this.isMultiligualModel) {
968
+ metadataTokens[50256] = '[Unused_50256]';
969
+ metadataTokens[50360] = '[Unused_50360]';
970
+ }
971
+ const languageTokenCount = tokenConfig.languageTokensEnd - tokenConfig.languageTokensStart;
972
+ for (let i = 0; i < languageTokenCount; i++) {
973
+ const tokenIndex = this.tokenConfig.languageTokensStart + i;
974
+ metadataTokens[tokenIndex] = `[Language_${i}]`;
975
+ }
976
+ const timestampTokensCount = 1501;
977
+ for (let i = 0; i < timestampTokensCount; i++) {
978
+ const tokenIndex = this.tokenConfig.timestampTokensStart + i;
979
+ const tokenTime = this.timestampTokenToSeconds(tokenIndex);
980
+ metadataTokens[tokenIndex] = `[Timestamp_${tokenTime.toFixed(2)}]`;
981
+ }
982
+ const inverseMetadataTokensLookup = {};
983
+ for (const [key, value] of Object.entries(metadataTokens)) {
984
+ inverseMetadataTokensLookup[value] = parseInt(key);
985
+ }
986
+ const patternString = `'s|'t|'re|'ve|'m|'ll|'d| ?\\p{L}+| ?\\p{N}+| ?[^\\s\\p{L}\\p{N}]+|\\s+(?!\\S)|\\s+`;
987
+ const { Tiktoken } = await import('tiktoken/lite');
988
+ this.tiktoken = new Tiktoken(tiktokenData, inverseMetadataTokensLookup, patternString);
989
+ logger.end();
990
+ }
991
+ async initializeEncoderSessionIfNeeded() {
992
+ if (this.audioEncoder) {
993
+ return;
994
+ }
995
+ const logger = new Logger();
996
+ await logger.startAsync(`Create encoder inference session for model '${this.modelName}'`);
997
+ const encoderFilePath = path.join(this.modelDir, 'encoder.onnx');
998
+ const Onnx = await import('onnxruntime-node');
999
+ const onnxSessionOptions = getOnnxSessionOptions({ executionProviders: this.encoderExecutionProviders });
1000
+ this.audioEncoder = await Onnx.InferenceSession.create(encoderFilePath, onnxSessionOptions);
1001
+ logger.end();
1002
+ }
1003
+ async initializeDecoderSessionIfNeeded() {
1004
+ if (this.textDecoder) {
1005
+ return;
1006
+ }
1007
+ const logger = new Logger();
1008
+ await logger.startAsync(`Create decoder inference session for model '${this.modelName}'`);
1009
+ const decoderFilePath = path.join(this.modelDir, 'decoder.onnx');
1010
+ const Onnx = await import('onnxruntime-node');
1011
+ const onnxSessionOptions = getOnnxSessionOptions({ executionProviders: this.decoderExecutionProviders });
1012
+ this.textDecoder = await Onnx.InferenceSession.create(decoderFilePath, onnxSessionOptions);
1013
+ logger.end();
1014
+ }
891
1015
  getKvDimensions(groupCount, length) {
892
1016
  const modelName = this.modelName;
893
1017
  if (modelName == 'tiny' || modelName == 'tiny.en') {
@@ -1008,7 +1132,7 @@ export class Whisper {
1008
1132
  }
1009
1133
  getSuppressedTextTokens() {
1010
1134
  const allowedPunctuationMarks = this.getAllowedPunctuationMarks();
1011
- const nonWordTokensData = this.getNonWordTokenData();
1135
+ const nonWordTokensData = this.getWordTokenData().nonWordTokenData;
1012
1136
  const suppressedTextTokens = nonWordTokensData
1013
1137
  .filter(entry => !allowedPunctuationMarks.includes(entry.text))
1014
1138
  .map(entry => entry.id);
@@ -1037,22 +1161,27 @@ export class Whisper {
1037
1161
  }
1038
1162
  return allowedPunctuation;
1039
1163
  }
1040
- getNonWordTokenData() {
1164
+ getWordTokenData() {
1165
+ const wordTokenData = [];
1041
1166
  const nonWordTokenData = [];
1042
- const invalidUTF8Char = String.fromCharCode(65533);
1043
1167
  for (let i = 0; i < this.tokenConfig.endOfTextToken; i++) {
1044
1168
  const tokenText = this.tokenToText(i, false);
1045
- const tokenTextWithoutWhitespace = tokenText.replaceAll(/\s/g, '');
1046
- const isNonWordToken = /^[\p{Punctuation}\p{Symbol}]+$/u.test(tokenTextWithoutWhitespace);
1047
- const containsInvalidUTF8 = getUTF32Chars(tokenTextWithoutWhitespace).utf32chars.includes(invalidUTF8Char);
1048
- if (isNonWordToken && !containsInvalidUTF8) {
1169
+ const isNonWordToken = /^[\s\p{Punctuation}\p{Symbol}]+$/u.test(tokenText);
1170
+ const containsInvalidUTF8 = containsInvalidCodepoint(tokenText);
1171
+ if (isNonWordToken && (this.isEnglishOnlyModel || !containsInvalidUTF8)) {
1049
1172
  nonWordTokenData.push({
1050
1173
  id: i,
1051
1174
  text: tokenText,
1052
1175
  });
1053
1176
  }
1177
+ else {
1178
+ wordTokenData.push({
1179
+ id: i,
1180
+ text: tokenText,
1181
+ });
1182
+ }
1054
1183
  }
1055
- return nonWordTokenData;
1184
+ return { wordTokenData, nonWordTokenData };
1056
1185
  }
1057
1186
  getTokensData(tokens) {
1058
1187
  const tokensData = [];
@@ -1065,88 +1194,6 @@ export class Whisper {
1065
1194
  return tokensData;
1066
1195
  }
1067
1196
  }
1068
- const filterbanks = [
1069
- /* 0 */ { startIndex: 1, weights: [0.02486259490251541,] },
1070
- /* 1 */ { startIndex: 1, weights: [0.001990821911022067, 0.022871771827340126,] },
1071
- /* 2 */ { startIndex: 2, weights: [0.003981643822044134, 0.02088095061480999,] },
1072
- /* 3 */ { startIndex: 3, weights: [0.0059724655002355576, 0.018890129402279854,] },
1073
- /* 4 */ { startIndex: 4, weights: [0.007963287644088268, 0.01689930632710457,] },
1074
- /* 5 */ { startIndex: 5, weights: [0.009954108856618404, 0.014908484183251858,] },
1075
- /* 6 */ { startIndex: 6, weights: [0.011944931000471115, 0.012917662039399147,] },
1076
- /* 7 */ { startIndex: 7, weights: [0.013935752213001251, 0.010926840826869011,] },
1077
- /* 8 */ { startIndex: 8, weights: [0.015926575288176537, 0.0089360186830163,] },
1078
- /* 9 */ { startIndex: 9, weights: [0.017917396500706673, 0.006945197004824877,] },
1079
- /* 10 */ { startIndex: 10, weights: [0.01990821771323681, 0.004954374860972166,] },
1080
- /* 11 */ { startIndex: 11, weights: [0.021899040788412094, 0.0029635531827807426,] },
1081
- /* 12 */ { startIndex: 12, weights: [0.02388986200094223, 0.0009727313299663365,] },
1082
- /* 13 */ { startIndex: 13, weights: [0.025880683213472366,] },
1083
- /* 14 */ { startIndex: 14, weights: [0.025835324078798294,] },
1084
- /* 15 */ { startIndex: 14, weights: [0.0010180906392633915, 0.023844502866268158,] },
1085
- /* 16 */ { startIndex: 15, weights: [0.003008912317454815, 0.021853681653738022,] },
1086
- /* 17 */ { startIndex: 16, weights: [0.004999734461307526, 0.019862858578562737,] },
1087
- /* 18 */ { startIndex: 17, weights: [0.006990555673837662, 0.0178720373660326,] },
1088
- /* 19 */ { startIndex: 18, weights: [0.008981377817690372, 0.015881216153502464,] },
1089
- /* 20 */ { startIndex: 19, weights: [0.010972199961543083, 0.013890394009649754,] },
1090
- /* 21 */ { startIndex: 20, weights: [0.01296302117407322, 0.011899571865797043,] },
1091
- /* 22 */ { startIndex: 21, weights: [0.01495384331792593, 0.009908749721944332,] },
1092
- /* 23 */ { startIndex: 22, weights: [0.01694466546177864, 0.007917927578091621,] },
1093
- /* 24 */ { startIndex: 23, weights: [0.018935488536953926, 0.005927106365561485,] },
1094
- /* 25 */ { startIndex: 24, weights: [0.020874010398983955, 0.004040425643324852,] },
1095
- /* 26 */ { startIndex: 25, weights: [0.022114217281341553, 0.0033186059445142746,] },
1096
- /* 27 */ { startIndex: 26, weights: [0.02173672430217266, 0.0036109676584601402,] },
1097
- /* 28 */ { startIndex: 27, weights: [0.020497702062129974, 0.004762193653732538,] },
1098
- /* 29 */ { startIndex: 28, weights: [0.018486659973859787, 0.006592618301510811,] },
1099
- /* 30 */ { startIndex: 29, weights: [0.01585603691637516, 0.00896277092397213,] },
1100
- /* 31 */ { startIndex: 30, weights: [0.012738768011331558, 0.011751330457627773,] },
1101
- /* 32 */ { startIndex: 31, weights: [0.009250369854271412, 0.014853144995868206,] },
1102
- /* 33 */ { startIndex: 32, weights: [0.005490840878337622, 0.018177473917603493, 0.0028155462350696325,] },
1103
- /* 34 */ { startIndex: 33, weights: [0.0015463664894923568, 0.01632951945066452, 0.007420188747346401,] },
1104
- /* 35 */ { startIndex: 35, weights: [0.011181050911545753, 0.012018864043056965,] },
1105
- /* 36 */ { startIndex: 36, weights: [0.006065350491553545, 0.016561277210712433, 0.004360878840088844,] },
1106
- /* 37 */ { startIndex: 37, weights: [0.0010297985281795263, 0.012770536355674267, 0.009707189165055752,] },
1107
- /* 38 */ { startIndex: 39, weights: [0.006986402906477451, 0.01485429983586073, 0.004391219466924667,] },
1108
- /* 39 */ { startIndex: 40, weights: [0.001418047584593296, 0.011486922390758991, 0.010089744813740253, 0.00040022286702878773,] },
1109
- /* 40 */ { startIndex: 42, weights: [0.005411104764789343, 0.014735566452145576, 0.006518189795315266,] },
1110
- /* 41 */ { startIndex: 44, weights: [0.00827841367572546, 0.012277561239898205, 0.00396781275048852,] },
1111
- /* 42 */ { startIndex: 45, weights: [0.002187808509916067, 0.010184479877352715, 0.00998187530785799, 0.0022864851634949446,] },
1112
- /* 43 */ { startIndex: 47, weights: [0.00386943481862545, 0.011274894699454308, 0.008466221392154694, 0.0013397691072896123,] },
1113
- /* 44 */ { startIndex: 49, weights: [0.004820294212549925, 0.011678251437842846, 0.007608682848513126, 0.0010091039584949613,] },
1114
- /* 45 */ { startIndex: 51, weights: [0.005156961735337973, 0.011507894843816757, 0.007301822770386934, 0.0011901655234396458,] },
1115
- /* 46 */ { startIndex: 53, weights: [0.004982104524970055, 0.010863498784601688, 0.007451189681887627, 0.001791381393559277,] },
1116
- /* 47 */ { startIndex: 55, weights: [0.004385921638458967, 0.009832492098212242, 0.007973956875503063, 0.002732589840888977,] },
1117
- /* 48 */ { startIndex: 57, weights: [0.0034474546555429697, 0.008491347543895245, 0.008797688409686089, 0.00394382793456316,] },
1118
- /* 49 */ { startIndex: 59, weights: [0.0022357646375894547, 0.0069067515432834625, 0.009859241545200348, 0.005364237818866968, 0.0008692338014952838,] },
1119
- /* 50 */ { startIndex: 61, weights: [0.0008110002381727099, 0.005136650986969471, 0.00946230161935091, 0.0069410777650773525, 0.0027783995028585196,] },
1120
- /* 51 */ { startIndex: 64, weights: [0.003231203882023692, 0.007237049750983715, 0.00862883497029543, 0.004773912951350212, 0.0009189908159896731,] },
1121
- /* 52 */ { startIndex: 66, weights: [0.001233637798577547, 0.0049433219246566296, 0.008653006516397, 0.006818502210080624, 0.003248583758249879,] },
1122
- /* 53 */ { startIndex: 69, weights: [0.0026164355222135782, 0.006051854696124792, 0.008880467154085636, 0.005574479699134827, 0.002268492942675948,] },
1123
- /* 54 */ { startIndex: 71, weights: [0.0002863667905330658, 0.003467798000201583, 0.0066492292098701, 0.00787146482616663, 0.004809896927326918, 0.0017483289120718837,] },
1124
- /* 55 */ { startIndex: 74, weights: [0.0009245910914614797, 0.0038708120118826628, 0.00681703258305788, 0.007283343467861414, 0.004448124207556248, 0.0016129047144204378,] },
1125
- /* 56 */ { startIndex: 77, weights: [0.0011703289346769452, 0.003898728871718049, 0.006627128925174475, 0.0070473202504217625, 0.004421714693307877, 0.0017961094854399562,] },
1126
- /* 57 */ { startIndex: 80, weights: [0.0010892992140725255, 0.003615982597693801, 0.006142666097730398, 0.007102936040610075, 0.004671447444707155, 0.002239959081634879,] },
1127
- /* 58 */ { startIndex: 83, weights: [0.0007392280967906117, 0.0030791081953793764, 0.005418988410383463, 0.007397185545414686, 0.005145462695509195, 0.002893739379942417, 0.0006420162972062826,] },
1128
- /* 59 */ { startIndex: 86, weights: [0.00017068670422304422, 0.0023375742603093386, 0.004504461772739887, 0.0066713495180010796, 0.005798479542136192, 0.003713231300935149, 0.0016279831761494279,] },
1129
- /* 60 */ { startIndex: 90, weights: [0.0014345343224704266, 0.0034412189852446318, 0.005447904113680124, 0.006591092795133591, 0.004660011734813452, 0.002728930441662669, 0.0007978491485118866,] },
1130
- /* 61 */ { startIndex: 93, weights: [0.0004075043834745884, 0.002265830524265766, 0.004124156199395657, 0.005982482805848122, 0.005700822453945875, 0.003912510350346565, 0.0021241982467472553, 0.0003358862304594368,] },
1131
- /* 62 */ { startIndex: 97, weights: [0.0010099108330905437, 0.002730846870690584, 0.004451782442629337, 0.006172718480229378, 0.005150905344635248, 0.0034948070533573627, 0.0018387088784947991, 0.0001826105872169137,] },
1132
- /* 63 */ { startIndex: 101, weights: [0.0012943691108375788, 0.002888072282075882, 0.004481775686144829, 0.006075479090213776, 0.0048866597935557365, 0.003353001084178686, 0.0018193417927250266, 0.00028568264679051936,] },
1133
- /* 64 */ { startIndex: 105, weights: [0.0013131388695910573, 0.0027890161145478487, 0.004264893010258675, 0.0057407706044614315, 0.004859979264438152, 0.0034397069830447435, 0.0020194342359900475, 0.0005991620710119605,] },
1134
- /* 65 */ { startIndex: 109, weights: [0.0011121684219688177, 0.002478930866345763, 0.0038456933107227087, 0.0052124555222690105, 0.005028639920055866, 0.0037133716978132725, 0.002398103242740035, 0.0010828346712514758,] },
1135
- /* 66 */ { startIndex: 113, weights: [0.0007317548734135926, 0.0019974694587290287, 0.003263183869421482, 0.004528898745775223, 0.005355686880648136, 0.004137659445405006, 0.0029196315445005894, 0.0017016039928421378, 0.0004835762665607035,] },
1136
- /* 67 */ { startIndex: 117, weights: [0.00020713974663522094, 0.0013792773243039846, 0.0025514145381748676, 0.003723552217707038, 0.004895689897239208, 0.004680895246565342, 0.0035529187880456448, 0.0024249425623565912, 0.0012969663366675377, 0.00016899015463422984,] },
1137
- /* 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,] },
1138
- /* 69 */ { startIndex: 127, weights: [0.000854626705404371, 0.001859853626228869, 0.002865080488845706, 0.003870307235047221, 0.00487553421407938, 0.00408313749358058, 0.003115783678367734, 0.0021484296303242445, 0.001181075582280755, 0.0002137213887181133,] },
1139
- /* 70 */ { startIndex: 132, weights: [0.0008483415003865957, 0.0017792496364563704, 0.0027101580053567886, 0.0036410661414265633, 0.004571974277496338, 0.004079728852957487, 0.003183893393725157, 0.002288057701662183, 0.0013922222424298525, 0.0004963868414051831,] },
1140
- /* 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,] },
1141
- /* 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,] },
1142
- /* 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,] },
1143
- /* 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,] },
1144
- /* 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,] },
1145
- /* 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,] },
1146
- /* 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,] },
1147
- /* 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,] },
1148
- /* 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,] },
1149
- ];
1150
1197
  export async function loadPackagesAndGetPaths(modelName, languageCode) {
1151
1198
  if (modelName) {
1152
1199
  modelName = normalizeWhisperModelName(modelName, languageCode);
@@ -1312,6 +1359,88 @@ const alignmentHeadsIndexes = {
1312
1359
  'large-v3': [212, 277, 331, 332, 333, 355, 356, 364, 371, 379, 391, 422, 423, 443, 449, 452, 465, 467, 473, 505, 521, 532, 555], // Temporary (may not be correct)
1313
1360
  'large': [212, 277, 331, 332, 333, 355, 356, 364, 371, 379, 391, 422, 423, 443, 449, 452, 465, 467, 473, 505, 521, 532, 555],
1314
1361
  };
1362
+ const filterbanks = [
1363
+ /* 0 */ { startIndex: 1, weights: [0.02486259490251541,] },
1364
+ /* 1 */ { startIndex: 1, weights: [0.001990821911022067, 0.022871771827340126,] },
1365
+ /* 2 */ { startIndex: 2, weights: [0.003981643822044134, 0.02088095061480999,] },
1366
+ /* 3 */ { startIndex: 3, weights: [0.0059724655002355576, 0.018890129402279854,] },
1367
+ /* 4 */ { startIndex: 4, weights: [0.007963287644088268, 0.01689930632710457,] },
1368
+ /* 5 */ { startIndex: 5, weights: [0.009954108856618404, 0.014908484183251858,] },
1369
+ /* 6 */ { startIndex: 6, weights: [0.011944931000471115, 0.012917662039399147,] },
1370
+ /* 7 */ { startIndex: 7, weights: [0.013935752213001251, 0.010926840826869011,] },
1371
+ /* 8 */ { startIndex: 8, weights: [0.015926575288176537, 0.0089360186830163,] },
1372
+ /* 9 */ { startIndex: 9, weights: [0.017917396500706673, 0.006945197004824877,] },
1373
+ /* 10 */ { startIndex: 10, weights: [0.01990821771323681, 0.004954374860972166,] },
1374
+ /* 11 */ { startIndex: 11, weights: [0.021899040788412094, 0.0029635531827807426,] },
1375
+ /* 12 */ { startIndex: 12, weights: [0.02388986200094223, 0.0009727313299663365,] },
1376
+ /* 13 */ { startIndex: 13, weights: [0.025880683213472366,] },
1377
+ /* 14 */ { startIndex: 14, weights: [0.025835324078798294,] },
1378
+ /* 15 */ { startIndex: 14, weights: [0.0010180906392633915, 0.023844502866268158,] },
1379
+ /* 16 */ { startIndex: 15, weights: [0.003008912317454815, 0.021853681653738022,] },
1380
+ /* 17 */ { startIndex: 16, weights: [0.004999734461307526, 0.019862858578562737,] },
1381
+ /* 18 */ { startIndex: 17, weights: [0.006990555673837662, 0.0178720373660326,] },
1382
+ /* 19 */ { startIndex: 18, weights: [0.008981377817690372, 0.015881216153502464,] },
1383
+ /* 20 */ { startIndex: 19, weights: [0.010972199961543083, 0.013890394009649754,] },
1384
+ /* 21 */ { startIndex: 20, weights: [0.01296302117407322, 0.011899571865797043,] },
1385
+ /* 22 */ { startIndex: 21, weights: [0.01495384331792593, 0.009908749721944332,] },
1386
+ /* 23 */ { startIndex: 22, weights: [0.01694466546177864, 0.007917927578091621,] },
1387
+ /* 24 */ { startIndex: 23, weights: [0.018935488536953926, 0.005927106365561485,] },
1388
+ /* 25 */ { startIndex: 24, weights: [0.020874010398983955, 0.004040425643324852,] },
1389
+ /* 26 */ { startIndex: 25, weights: [0.022114217281341553, 0.0033186059445142746,] },
1390
+ /* 27 */ { startIndex: 26, weights: [0.02173672430217266, 0.0036109676584601402,] },
1391
+ /* 28 */ { startIndex: 27, weights: [0.020497702062129974, 0.004762193653732538,] },
1392
+ /* 29 */ { startIndex: 28, weights: [0.018486659973859787, 0.006592618301510811,] },
1393
+ /* 30 */ { startIndex: 29, weights: [0.01585603691637516, 0.00896277092397213,] },
1394
+ /* 31 */ { startIndex: 30, weights: [0.012738768011331558, 0.011751330457627773,] },
1395
+ /* 32 */ { startIndex: 31, weights: [0.009250369854271412, 0.014853144995868206,] },
1396
+ /* 33 */ { startIndex: 32, weights: [0.005490840878337622, 0.018177473917603493, 0.0028155462350696325,] },
1397
+ /* 34 */ { startIndex: 33, weights: [0.0015463664894923568, 0.01632951945066452, 0.007420188747346401,] },
1398
+ /* 35 */ { startIndex: 35, weights: [0.011181050911545753, 0.012018864043056965,] },
1399
+ /* 36 */ { startIndex: 36, weights: [0.006065350491553545, 0.016561277210712433, 0.004360878840088844,] },
1400
+ /* 37 */ { startIndex: 37, weights: [0.0010297985281795263, 0.012770536355674267, 0.009707189165055752,] },
1401
+ /* 38 */ { startIndex: 39, weights: [0.006986402906477451, 0.01485429983586073, 0.004391219466924667,] },
1402
+ /* 39 */ { startIndex: 40, weights: [0.001418047584593296, 0.011486922390758991, 0.010089744813740253, 0.00040022286702878773,] },
1403
+ /* 40 */ { startIndex: 42, weights: [0.005411104764789343, 0.014735566452145576, 0.006518189795315266,] },
1404
+ /* 41 */ { startIndex: 44, weights: [0.00827841367572546, 0.012277561239898205, 0.00396781275048852,] },
1405
+ /* 42 */ { startIndex: 45, weights: [0.002187808509916067, 0.010184479877352715, 0.00998187530785799, 0.0022864851634949446,] },
1406
+ /* 43 */ { startIndex: 47, weights: [0.00386943481862545, 0.011274894699454308, 0.008466221392154694, 0.0013397691072896123,] },
1407
+ /* 44 */ { startIndex: 49, weights: [0.004820294212549925, 0.011678251437842846, 0.007608682848513126, 0.0010091039584949613,] },
1408
+ /* 45 */ { startIndex: 51, weights: [0.005156961735337973, 0.011507894843816757, 0.007301822770386934, 0.0011901655234396458,] },
1409
+ /* 46 */ { startIndex: 53, weights: [0.004982104524970055, 0.010863498784601688, 0.007451189681887627, 0.001791381393559277,] },
1410
+ /* 47 */ { startIndex: 55, weights: [0.004385921638458967, 0.009832492098212242, 0.007973956875503063, 0.002732589840888977,] },
1411
+ /* 48 */ { startIndex: 57, weights: [0.0034474546555429697, 0.008491347543895245, 0.008797688409686089, 0.00394382793456316,] },
1412
+ /* 49 */ { startIndex: 59, weights: [0.0022357646375894547, 0.0069067515432834625, 0.009859241545200348, 0.005364237818866968, 0.0008692338014952838,] },
1413
+ /* 50 */ { startIndex: 61, weights: [0.0008110002381727099, 0.005136650986969471, 0.00946230161935091, 0.0069410777650773525, 0.0027783995028585196,] },
1414
+ /* 51 */ { startIndex: 64, weights: [0.003231203882023692, 0.007237049750983715, 0.00862883497029543, 0.004773912951350212, 0.0009189908159896731,] },
1415
+ /* 52 */ { startIndex: 66, weights: [0.001233637798577547, 0.0049433219246566296, 0.008653006516397, 0.006818502210080624, 0.003248583758249879,] },
1416
+ /* 53 */ { startIndex: 69, weights: [0.0026164355222135782, 0.006051854696124792, 0.008880467154085636, 0.005574479699134827, 0.002268492942675948,] },
1417
+ /* 54 */ { startIndex: 71, weights: [0.0002863667905330658, 0.003467798000201583, 0.0066492292098701, 0.00787146482616663, 0.004809896927326918, 0.0017483289120718837,] },
1418
+ /* 55 */ { startIndex: 74, weights: [0.0009245910914614797, 0.0038708120118826628, 0.00681703258305788, 0.007283343467861414, 0.004448124207556248, 0.0016129047144204378,] },
1419
+ /* 56 */ { startIndex: 77, weights: [0.0011703289346769452, 0.003898728871718049, 0.006627128925174475, 0.0070473202504217625, 0.004421714693307877, 0.0017961094854399562,] },
1420
+ /* 57 */ { startIndex: 80, weights: [0.0010892992140725255, 0.003615982597693801, 0.006142666097730398, 0.007102936040610075, 0.004671447444707155, 0.002239959081634879,] },
1421
+ /* 58 */ { startIndex: 83, weights: [0.0007392280967906117, 0.0030791081953793764, 0.005418988410383463, 0.007397185545414686, 0.005145462695509195, 0.002893739379942417, 0.0006420162972062826,] },
1422
+ /* 59 */ { startIndex: 86, weights: [0.00017068670422304422, 0.0023375742603093386, 0.004504461772739887, 0.0066713495180010796, 0.005798479542136192, 0.003713231300935149, 0.0016279831761494279,] },
1423
+ /* 60 */ { startIndex: 90, weights: [0.0014345343224704266, 0.0034412189852446318, 0.005447904113680124, 0.006591092795133591, 0.004660011734813452, 0.002728930441662669, 0.0007978491485118866,] },
1424
+ /* 61 */ { startIndex: 93, weights: [0.0004075043834745884, 0.002265830524265766, 0.004124156199395657, 0.005982482805848122, 0.005700822453945875, 0.003912510350346565, 0.0021241982467472553, 0.0003358862304594368,] },
1425
+ /* 62 */ { startIndex: 97, weights: [0.0010099108330905437, 0.002730846870690584, 0.004451782442629337, 0.006172718480229378, 0.005150905344635248, 0.0034948070533573627, 0.0018387088784947991, 0.0001826105872169137,] },
1426
+ /* 63 */ { startIndex: 101, weights: [0.0012943691108375788, 0.002888072282075882, 0.004481775686144829, 0.006075479090213776, 0.0048866597935557365, 0.003353001084178686, 0.0018193417927250266, 0.00028568264679051936,] },
1427
+ /* 64 */ { startIndex: 105, weights: [0.0013131388695910573, 0.0027890161145478487, 0.004264893010258675, 0.0057407706044614315, 0.004859979264438152, 0.0034397069830447435, 0.0020194342359900475, 0.0005991620710119605,] },
1428
+ /* 65 */ { startIndex: 109, weights: [0.0011121684219688177, 0.002478930866345763, 0.0038456933107227087, 0.0052124555222690105, 0.005028639920055866, 0.0037133716978132725, 0.002398103242740035, 0.0010828346712514758,] },
1429
+ /* 66 */ { startIndex: 113, weights: [0.0007317548734135926, 0.0019974694587290287, 0.003263183869421482, 0.004528898745775223, 0.005355686880648136, 0.004137659445405006, 0.0029196315445005894, 0.0017016039928421378, 0.0004835762665607035,] },
1430
+ /* 67 */ { startIndex: 117, weights: [0.00020713974663522094, 0.0013792773243039846, 0.0025514145381748676, 0.003723552217707038, 0.004895689897239208, 0.004680895246565342, 0.0035529187880456448, 0.0024249425623565912, 0.0012969663366675377, 0.00016899015463422984,] },
1431
+ /* 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,] },
1432
+ /* 69 */ { startIndex: 127, weights: [0.000854626705404371, 0.001859853626228869, 0.002865080488845706, 0.003870307235047221, 0.00487553421407938, 0.00408313749358058, 0.003115783678367734, 0.0021484296303242445, 0.001181075582280755, 0.0002137213887181133,] },
1433
+ /* 70 */ { startIndex: 132, weights: [0.0008483415003865957, 0.0017792496364563704, 0.0027101580053567886, 0.0036410661414265633, 0.004571974277496338, 0.004079728852957487, 0.003183893393725157, 0.002288057701662183, 0.0013922222424298525, 0.0004963868414051831,] },
1434
+ /* 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,] },
1435
+ /* 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,] },
1436
+ /* 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,] },
1437
+ /* 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,] },
1438
+ /* 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,] },
1439
+ /* 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,] },
1440
+ /* 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,] },
1441
+ /* 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,] },
1442
+ /* 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,] },
1443
+ ];
1315
1444
  export const defaultWhisperOptions = {
1316
1445
  model: undefined,
1317
1446
  temperature: 0.1,
@@ -1322,6 +1451,28 @@ export const defaultWhisperOptions = {
1322
1451
  maxTokensPerPart: 250,
1323
1452
  suppressRepetition: true,
1324
1453
  decodeTimestampTokens: true,
1454
+ endTokenThreshold: 0.9,
1455
+ includeEndTokenInCandidates: true,
1456
+ encoderProvider: undefined,
1457
+ decoderProvider: undefined,
1325
1458
  seed: undefined,
1326
1459
  };
1460
+ export const defaultWhisperAlignmentOptions = {
1461
+ model: undefined,
1462
+ endTokenThreshold: 0.9,
1463
+ encoderProvider: undefined,
1464
+ decoderProvider: undefined
1465
+ };
1466
+ export const defaultWhisperLanguageDetectionOptions = {
1467
+ model: undefined,
1468
+ temperature: 1.0,
1469
+ encoderProvider: undefined,
1470
+ decoderProvider: undefined,
1471
+ };
1472
+ export const defaultWhisperVADOptions = {
1473
+ model: undefined,
1474
+ temperature: 1.0,
1475
+ encoderProvider: undefined,
1476
+ decoderProvider: undefined,
1477
+ };
1327
1478
  //# sourceMappingURL=WhisperSTT.js.map