echogarden 1.0.5 → 1.1.1

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 +7 -11
  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 +639 -478
  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 +7 -5
  81. package/docs/Engines.md +43 -32
  82. package/docs/Licenses.md +3 -4
  83. package/docs/Options.md +49 -13
  84. package/docs/Releases.md +7 -7
  85. package/docs/Server.md +8 -6
  86. package/docs/Tasklist.md +49 -62
  87. package/docs/Technical.md +1 -1
  88. package/package.json +7 -11
  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 +14 -16
  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 +959 -669
  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,99 @@ 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 (disabled for now)
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);
252
+ // Add tokens to output
246
253
  allDecodedTokens.push(...partTokens);
247
254
  timeline.push(...partTimeline);
255
+ // Update previous text tokens
248
256
  previousPartTextTokens = partTokens.filter(token => this.isTextToken(token));
257
+ audioOffset = audioEndOffset;
249
258
  logger.end();
250
259
  }
260
+ // Convert token timeline to word timeline
251
261
  timeline = this.tokenTimelineToWordTimeline(timeline, language);
262
+ // Convert tokens to transcript
252
263
  const transcript = this.tokensToText(allDecodedTokens).trim();
253
264
  logger.end();
254
265
  return { transcript, timeline };
255
266
  }
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);
267
+ async align(rawAudio, transcript, sourceLanguage, task, whisperAlignmentOptions) {
268
+ await this.initializeTokenizerIfNeeded();
269
+ whisperAlignmentOptions = extendDeep(defaultWhisperAlignmentOptions, whisperAlignmentOptions);
270
+ const shouldSplitToSentences = false;
271
+ const targetLanguage = task === 'transcribe' ? sourceLanguage : 'en';
272
+ let simplifiedTranscript = '';
273
+ if (shouldSplitToSentences) {
274
+ const sentences = splitToSentences(transcript, targetLanguage);
275
+ for (const sentence of sentences) {
276
+ let sentenceWords = await splitToWords(sentence, 'en');
277
+ sentenceWords = sentenceWords.filter(word => isWord(word));
278
+ simplifiedTranscript += sentenceWords.join(' ');
279
+ simplifiedTranscript += ' ';
280
+ }
281
+ }
282
+ else {
283
+ let words = await splitToWords(transcript, targetLanguage);
284
+ words = words.map(word => word.trim());
285
+ words = words.filter(word => isWord(word));
286
+ simplifiedTranscript = words.join(' ');
287
+ }
288
+ // Tokenize the transcript
289
+ const simplifiedTranscriptTokens = this.textToTokens(simplifiedTranscript);
290
+ // Initialize custom logit filter that allows only the transcript tokens to be decoded
291
+ // in order.
264
292
  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;
293
+ const logitFilter = (logits, decodedTokens, isFirstPart, isFinalPart) => {
294
+ const decodedTextTokens = decodedTokens.filter(token => this.isTextToken(token));
295
+ const nextTokenToDecode = simplifiedTranscriptTokens[decodedTextTokens.length] ?? endOfTextToken;
296
+ const newLogits = logits.map((logit, index) => {
297
+ if (index === nextTokenToDecode) {
298
+ return logit;
299
+ }
300
+ // If it's the final part, the ent-of-text token logit is set to -Infinity.
301
+ // This will force to force all transcript tokens to be decoded even if the model doesn't
302
+ // recognize them.
303
+ if (!isFinalPart && index === endOfTextToken) {
304
+ return logit;
305
+ }
306
+ return -Infinity;
307
+ });
308
+ return newLogits;
309
+ };
310
+ // Set options for alignment
311
+ const options = {
312
+ model: this.modelName,
313
+ temperature: 0.0,
314
+ prompt: undefined,
315
+ topCandidateCount: 1,
316
+ punctuationThreshold: Infinity,
317
+ autoPromptParts: false,
318
+ maxTokensPerPart: Infinity,
319
+ suppressRepetition: false,
320
+ decodeTimestampTokens: true,
321
+ endTokenThreshold: whisperAlignmentOptions.endTokenThreshold,
322
+ includeEndTokenInCandidates: false,
323
+ seed: undefined,
324
+ };
325
+ // Recognize
326
+ const { timeline } = await this.recognize(rawAudio, task, sourceLanguage, options, logitFilter);
327
+ return timeline;
278
328
  }
279
329
  async detectLanguage(audioFeatures, temperature) {
280
330
  if (!this.isMultiligualModel) {
281
331
  throw new Error('Language detection is only supported with multilingual models');
282
332
  }
283
- await this.initializeTokenizerIfNeeded();
284
333
  await this.initializeDecoderSessionIfNeeded();
285
334
  // Prepare and run decoder
286
335
  const logger = new Logger();
@@ -317,38 +366,76 @@ export class Whisper {
317
366
  logger.end();
318
367
  return results;
319
368
  }
320
- async decodeTokens(audioFeatures, initialTokens, audioDuration, isFirstPart, isFinalPart, options) {
369
+ async detectVoiceActivity(audioFeatures, temperature) {
370
+ await this.initializeDecoderSessionIfNeeded();
371
+ // Prepare and run decoder
372
+ const logger = new Logger();
373
+ await logger.startAsync('Detect voice activity with Whisper model');
374
+ const sotToken = this.tokenConfig.startOfTextToken;
375
+ const initialTokens = [sotToken];
376
+ const offset = 0;
377
+ const Onnx = await import('onnxruntime-node');
378
+ const initialKvDimensions = this.getKvDimensions(1, initialTokens.length);
379
+ const kvCacheTensor = new Onnx.Tensor('float32', new Float32Array(initialKvDimensions[0] * initialKvDimensions[1] * initialKvDimensions[2] * initialKvDimensions[3]), initialKvDimensions);
380
+ const tokensTensor = new Onnx.Tensor('int64', new BigInt64Array(initialTokens.map(token => BigInt(token))), [1, initialTokens.length]);
381
+ const offsetTensor = new Onnx.Tensor('int64', new BigInt64Array([BigInt(offset)]), []);
382
+ const decoderInputs = {
383
+ tokens: tokensTensor,
384
+ audio_features: audioFeatures,
385
+ kv_cache: kvCacheTensor,
386
+ offset: offsetTensor
387
+ };
388
+ const decoderOutputs = await this.textDecoder.run(decoderInputs);
389
+ const logitsBuffer = decoderOutputs['logits'].data;
390
+ const tokenConfig = this.tokenConfig;
391
+ const logits = Array.from(logitsBuffer);
392
+ const probabilities = softmax(logits, temperature);
393
+ const noSpeechProbability = probabilities[tokenConfig.nonSpeechToken];
394
+ return 1.0 - noSpeechProbability;
395
+ }
396
+ // Decode tokens using the decoder model
397
+ async decodeTokens(audioFeatures, initialTokens, audioDuration, isFirstPart, isFinalPart, options, logitFilter) {
398
+ // Initialize
321
399
  await this.initializeTokenizerIfNeeded();
322
400
  await this.initializeDecoderSessionIfNeeded();
323
401
  const logger = new Logger();
324
- const allowedPunctuationMarks = this.getAllowedPunctuationMarks();
325
402
  await logger.startAsync('Decode text tokens with Whisper decoder model');
326
403
  options = extendDeep(defaultWhisperOptions, options);
327
404
  const Onnx = await import('onnxruntime-node');
405
+ // Get token information
328
406
  const endOfTextToken = this.tokenConfig.endOfTextToken;
329
407
  const timestampTokensStart = this.tokenConfig.timestampTokensStart;
330
- const suppressedTokens = new Set(this.getSuppressedTokens());
408
+ const suppressedTextTokens = this.getSuppressedTextTokens();
409
+ const suppressedMetadataTokens = this.getSuppressedMetadataTokens();
410
+ const allowedPunctuationMarks = this.getAllowedPunctuationMarks();
331
411
  const spaceToken = this.textToTokens(' ')[0];
332
- const maxDecodedTokenCount = options.maxTokensPerPart;
412
+ // Initialize variables for decoding loop
333
413
  let decodedTokens = initialTokens.slice();
334
414
  const initialKvDimensions = this.getKvDimensions(1, decodedTokens.length);
335
415
  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 = [];
416
+ let decodedTokensTimestampLogits = [];
417
+ let decodedTokensConfidence = [];
340
418
  let decodedTokensCrossAttentionQKs = [];
341
419
  for (let i = 0; i < decodedTokens.length; i++) {
420
+ decodedTokensTimestampLogits.push(new Array(1501));
421
+ decodedTokensConfidence.push(1.0);
342
422
  decodedTokensCrossAttentionQKs.push(undefined);
343
423
  }
424
+ let lastTimestampTokenIndex = -1;
425
+ let timestampTokenSeenCount = 0;
344
426
  let bufferedTokensToPrint = [];
427
+ // Define method to add a token to output
428
+ function addToken(tokenToAdd, timestampLogits, confidence, crossAttentionQKs) {
429
+ decodedTokens.push(tokenToAdd);
430
+ decodedTokensTimestampLogits.push(timestampLogits);
431
+ decodedTokensConfidence.push(confidence);
432
+ decodedTokensCrossAttentionQKs.push(crossAttentionQKs);
433
+ }
345
434
  // Start decoding loop
346
- for (let decodedTokenCount = 0; decodedTokenCount < maxDecodedTokenCount; decodedTokenCount++) {
435
+ for (let decodedTokenCount = 0; decodedTokenCount < options.maxTokensPerPart; decodedTokenCount++) {
347
436
  const isInitialState = decodedTokens.length == initialTokens.length;
348
- const tokensToDecode = isInitialState ? decodedTokens : [decodedTokens[decodedTokens.length - 1]];
349
- const offset = isInitialState ? 0 : decodedTokens.length;
437
+ // If not in initial state, reshape KV Cache tensor to accomodate a new output token
350
438
  if (!isInitialState) {
351
- // Reshape KV Cache tensor
352
439
  const dims = kvCacheTensor.dims;
353
440
  const currentKvCacheGroups = splitFloat32Array(kvCacheTensor.data, dims[2] * dims[3]);
354
441
  const reshapedKvCacheTensor = new Onnx.Tensor('float32', new Float32Array(dims[0] * dims[1] * (decodedTokens.length) * dims[3]), [dims[0], dims[1], decodedTokens.length, dims[3]]);
@@ -358,199 +445,255 @@ export class Whisper {
358
445
  }
359
446
  kvCacheTensor = reshapedKvCacheTensor;
360
447
  }
361
- // Prepare and run decoder
448
+ // Prepare values for decoder
449
+ const tokensToDecode = isInitialState ? decodedTokens : [decodedTokens[decodedTokens.length - 1]];
450
+ const offset = isInitialState ? 0 : decodedTokens.length;
362
451
  const tokensTensor = new Onnx.Tensor('int64', new BigInt64Array(tokensToDecode.map(token => BigInt(token))), [1, tokensToDecode.length]);
363
452
  const offsetTensor = new Onnx.Tensor('int64', new BigInt64Array([BigInt(offset)]), []);
364
- const decoderInputs = { tokens: tokensTensor, audio_features: audioFeatures, kv_cache: kvCacheTensor, offset: offsetTensor };
453
+ const decoderInputs = {
454
+ tokens: tokensTensor,
455
+ audio_features: audioFeatures,
456
+ kv_cache: kvCacheTensor,
457
+ offset: offsetTensor
458
+ };
459
+ // Run decoder model
365
460
  const decoderOutputs = await this.textDecoder.run(decoderInputs);
461
+ // Extract decoder model results
366
462
  const logitsBuffer = decoderOutputs['logits'].data;
367
463
  kvCacheTensor = decoderOutputs['output_kv_cache'];
464
+ const crossAttentionQKsForTokenOnnx = decoderOutputs['cross_attention_qks'];
465
+ const crossAttentionQKsForToken = makeOnnxLikeFloat32Tensor(crossAttentionQKsForTokenOnnx);
466
+ crossAttentionQKsForTokenOnnx.dispose();
368
467
  // Compute logits
369
- const resultLogits = splitFloat32Array(logitsBuffer, logitsBuffer.length / decoderOutputs['logits'].dims[1]);
370
- const allTokenLogits = Array.from(resultLogits[resultLogits.length - 1]);
371
- const timestampTokenLogits = allTokenLogits.slice(timestampTokensStart);
372
- // Suppress logits for tokens in the suppressed set
373
- for (let logitIndex = 0; logitIndex < allTokenLogits.length; logitIndex++) {
374
- const isWrongTokenForInitialState = isInitialState &&
375
- (logitIndex === spaceToken || logitIndex === endOfTextToken);
376
- const isInSuppressedList = suppressedTokens.has(logitIndex);
377
- const shouldSuppressToken = isWrongTokenForInitialState || isInSuppressedList;
378
- if (shouldSuppressToken) {
379
- allTokenLogits[logitIndex] = -Infinity;
380
- }
468
+ const resultLogitsFloatArrays = splitFloat32Array(logitsBuffer, logitsBuffer.length / decoderOutputs['logits'].dims[1]);
469
+ const allTokenLogits = Array.from(resultLogitsFloatArrays[resultLogitsFloatArrays.length - 1]);
470
+ // Suppress metadata tokens in the suppression set
471
+ for (const suppressedTokenIndex of suppressedMetadataTokens) {
472
+ allTokenLogits[suppressedTokenIndex] = -Infinity;
381
473
  }
382
- // Add best token
383
- function addToken(tokenToAdd, timestampLogits, confidence) {
384
- decodedTokens.push(tokenToAdd);
385
- decodedTokensTimestampLogits.push(timestampLogits);
386
- decodedTokensCrossAttentionQKs.push(decoderOutputs['cross_attention_qks']);
387
- decodedTokensConfidence.push(confidence);
474
+ if (isInitialState) {
475
+ // If in initial state, suppress end-of-text token
476
+ allTokenLogits[endOfTextToken] = -Infinity;
388
477
  }
389
- let shouldDecodeNonTimestampToken = true;
390
- if (options.decodeTimestampTokens) {
478
+ const timestampTokenLogits = allTokenLogits.slice(timestampTokensStart);
479
+ const decodeTimestampTokenIfNeeded = () => {
480
+ // Try to decode a timestamp token, if needed
481
+ // If timestamp tokens is disabled in options, don't decode a timestamp
482
+ if (!options.decodeTimestampTokens) {
483
+ return false;
484
+ }
485
+ const previousTokenWasTimestamp = this.isTimestampToken(decodedTokens[decodedTokens.length - 1]);
486
+ const secondPreviousTokenWasTimestamp = this.isTimestampToken(decodedTokens[decodedTokens.length - 2]);
487
+ // If there are two successive timestamp tokens decoded, or the previous timestamp was the first token,
488
+ // don't decode a timestamp
489
+ if (previousTokenWasTimestamp &&
490
+ (decodedTokens.length === initialTokens.length + 1) || secondPreviousTokenWasTimestamp) {
491
+ return false;
492
+ }
391
493
  // Derive token probabilities
392
494
  const probabilities = softmax(allTokenLogits, 1.0);
393
495
  const logProbabilities = logOfVector(probabilities);
394
496
  const nonTimestampTokenLogProbs = logProbabilities.slice(0, timestampTokensStart);
497
+ // Find highest non-timestamp token
395
498
  const indexOfMaxNonTimestampLogProb = indexOfMax(nonTimestampTokenLogProbs);
396
499
  const valueOfMaxNonTimestampLogProb = nonTimestampTokenLogProbs[indexOfMaxNonTimestampLogProb];
500
+ // Find highest timestamp token
397
501
  const timestampTokenLogProbs = logProbabilities.slice(timestampTokensStart);
398
502
  const indexOfMaxTimestampLogProb = indexOfMax(timestampTokenLogProbs);
503
+ // Compute the log of the sum of exponentials of the log probabilities
504
+ // of the timestamp tokens
399
505
  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;
506
+ // If the sum isn't greater than the log probability of the highest non-timestamp token,
507
+ // don't decode a timestamp
508
+ if (logSumExpOfTimestampTokenLogProbs <= valueOfMaxNonTimestampLogProb) {
509
+ return false;
405
510
  }
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;
511
+ // Decode a timestamp token
512
+ timestampTokenSeenCount += 1;
513
+ if (previousTokenWasTimestamp) {
514
+ // If previously decoded token was a timestamp token, repeat it
515
+ const previousToken = decodedTokens[decodedTokens.length - 1];
516
+ const previousTokenTimestampLogits = decodedTokensTimestampLogits[decodedTokensTimestampLogits.length - 1];
517
+ const previousTokenConfidence = decodedTokensConfidence[decodedTokensConfidence.length - 1];
518
+ addToken(previousToken, previousTokenTimestampLogits, previousTokenConfidence, crossAttentionQKsForToken);
519
+ lastTimestampTokenIndex = decodedTokens.length;
520
+ }
521
+ else {
522
+ // Otherwise decode the highest probability timestamp
523
+ const timestampToken = timestampTokensStart + indexOfMaxTimestampLogProb;
524
+ const confidence = probabilities[timestampToken];
525
+ addToken(timestampToken, timestampTokenLogits, confidence, crossAttentionQKsForToken);
424
526
  }
527
+ return true;
528
+ };
529
+ // Call the method to decode timestamp token if needed
530
+ const timestampTokenDecoded = decodeTimestampTokenIfNeeded();
531
+ if (timestampTokenDecoded) {
532
+ await yieldToEventLoop();
533
+ continue;
425
534
  }
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
- }
535
+ // Decode a non-timestamp token
536
+ let nonTimestampTokenLogits = allTokenLogits.slice(0, timestampTokensStart);
537
+ let shouldDecodeEndfOfTextToken = false;
538
+ // If not in initial state, and the end-of-text token's probability is sufficiently higher than
539
+ // the second highest ranked token, then accept end-of-text
540
+ if (!isInitialState) {
541
+ const endOfTextTokenLogit = nonTimestampTokenLogits[endOfTextToken];
542
+ const otherTokensLogits = nonTimestampTokenLogits.slice();
543
+ otherTokensLogits[endOfTextToken] = -Infinity;
544
+ const indexOfMaximumOtherTokenLogit = indexOfMax(otherTokensLogits);
545
+ const maximumOtherTokenLogit = nonTimestampTokenLogits[indexOfMaximumOtherTokenLogit];
546
+ const endProbabilities = softmax([endOfTextTokenLogit, maximumOtherTokenLogit], 1.0);
547
+ if (endProbabilities[0] > options.endTokenThreshold) {
548
+ shouldDecodeEndfOfTextToken = true;
452
549
  }
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
- });
550
+ }
551
+ if (logitFilter) {
552
+ // Apply custom logit filter function if given
553
+ nonTimestampTokenLogits = logitFilter(nonTimestampTokenLogits, decodedTokens, isFirstPart, isFinalPart);
554
+ // If the custom filter set the end-of-text token to be Infinity, or -Infinity,
555
+ // then override any previous decision and accept or reject it, respectively
556
+ if (nonTimestampTokenLogits[endOfTextToken] === Infinity) {
557
+ shouldDecodeEndfOfTextToken = true;
463
558
  }
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;
559
+ else if (nonTimestampTokenLogits[endOfTextToken] === -Infinity) {
560
+ shouldDecodeEndfOfTextToken = false;
477
561
  }
478
- let chosenCandidateRank;
479
- if (rankOfPromisingPunctuationToken >= 0 &&
480
- rankOfPromisingPunctuationToken < rankOfSpaceToken) {
481
- chosenCandidateRank = rankOfPromisingPunctuationToken;
562
+ // If filter caused all word token logits to be -Infinity, then there is no
563
+ // other token to decode. Fall back to accept end-of-text
564
+ if (nonTimestampTokenLogits.slice(0, endOfTextToken).every(logit => logit === -Infinity)) {
565
+ shouldDecodeEndfOfTextToken = true;
482
566
  }
483
- else {
484
- chosenCandidateRank = this.randomGen.selectRandomIndexFromDistribution(topCandidateProbabilities);
567
+ }
568
+ else {
569
+ // Otherwise, suppress text tokens in the suppression set
570
+ for (const suppressedTokenIndex of suppressedTextTokens) {
571
+ nonTimestampTokenLogits[suppressedTokenIndex] = -Infinity;
485
572
  }
486
- const chosenToken = topCandidates[chosenCandidateRank].token;
487
- const chosenTokenConfidence = topCandidateProbabilities[chosenCandidateRank];
488
- addToken(chosenToken, timestampTokenLogits, chosenTokenConfidence);
489
- if (chosenToken === endOfTextToken) {
490
- break;
573
+ // Suppress the space token if at initial state
574
+ if (isInitialState) {
575
+ nonTimestampTokenLogits[spaceToken] = -Infinity;
491
576
  }
492
- if (this.isTextToken(chosenToken)) {
493
- bufferedTokensToPrint.push(chosenToken);
494
- let textToPrint = this.tokensToText(bufferedTokensToPrint);
495
- if (textToPrint.codePointAt(0) !== 65533) {
496
- if (isFirstPart && decodedTokens.every(token => this.isMetadataToken(token))) {
497
- textToPrint = textToPrint.trimStart();
498
- }
499
- logger.write(textToPrint);
500
- bufferedTokensToPrint = [];
577
+ }
578
+ // If end-of-text token should be decoded, then add it and break
579
+ // out of the loop
580
+ if (shouldDecodeEndfOfTextToken) {
581
+ addToken(endOfTextToken, timestampTokenLogits, 1.0, crossAttentionQKsForToken);
582
+ break;
583
+ }
584
+ // Suppress end-of-text token if it shouldn't be included in candidates
585
+ if (!options.includeEndTokenInCandidates) {
586
+ nonTimestampTokenLogits[endOfTextToken] = -Infinity;
587
+ }
588
+ // Find top candidates
589
+ const sortedNonTimestampLogitsWithIndexes = Array.from(nonTimestampTokenLogits).map((logit, index) => ({ token: index, logit }));
590
+ sortedNonTimestampLogitsWithIndexes.sort((a, b) => b.logit - a.logit);
591
+ let topCandidates = sortedNonTimestampLogitsWithIndexes.slice(0, options.topCandidateCount)
592
+ .map(entry => ({
593
+ token: entry.token,
594
+ logit: entry.logit,
595
+ text: this.tokenToText(entry.token, true)
596
+ }));
597
+ // Apply repetition suppression if enabled
598
+ if (options.suppressRepetition) {
599
+ // Using some hardcoded constants, for now
600
+ const tokenWindowSize = 30;
601
+ const thresholdMatchLength = 6;
602
+ const thresholdCycleRepetition = 2.0;
603
+ const filteredCandidates = [];
604
+ for (const candidate of topCandidates) {
605
+ const lastDecodedTextTokens = decodedTokens
606
+ .filter(token => this.isTextToken(token))
607
+ .reverse()
608
+ .slice(0, tokenWindowSize);
609
+ const { longestMatch, longestCycleRepetition } = getTokenRepetitionScore([candidate.token, ...lastDecodedTextTokens]);
610
+ if (longestMatch >= thresholdMatchLength || longestCycleRepetition >= thresholdCycleRepetition) {
611
+ continue;
501
612
  }
613
+ filteredCandidates.push(candidate);
614
+ }
615
+ // If all candidates have been filtered out, accept an end-of-text token
616
+ if (filteredCandidates.length === 0) {
617
+ filteredCandidates.push({
618
+ token: endOfTextToken,
619
+ logit: Infinity,
620
+ text: this.tokenToText(endOfTextToken, true)
621
+ });
622
+ }
623
+ topCandidates = filteredCandidates;
624
+ }
625
+ // Compute top candidate probabilities
626
+ const topCandidateProbabilities = softmax(topCandidates.map(a => a.logit), options.temperature);
627
+ // Find highest ranking punctuation token
628
+ const rankOfPromisingPunctuationToken = topCandidates.findIndex((entry, index) => {
629
+ const tokenText = this.tokenToText(entry.token).trim();
630
+ const isPunctuationToken = allowedPunctuationMarks.includes(tokenText);
631
+ if (!isPunctuationToken) {
632
+ return false;
633
+ }
634
+ const tokenProb = topCandidateProbabilities[index];
635
+ return tokenProb >= options.punctuationThreshold;
636
+ });
637
+ // Find rank of space token
638
+ let rankOfSpaceToken = topCandidates.findIndex(candidate => candidate.token === spaceToken);
639
+ if (rankOfSpaceToken < 0) {
640
+ rankOfSpaceToken = Infinity;
641
+ }
642
+ // Choose token
643
+ let chosenCandidateRank;
644
+ // Select a high-ranking punctuation token if found, and it has
645
+ // a rank higher than the space token,
646
+ if (rankOfPromisingPunctuationToken >= 0 &&
647
+ rankOfPromisingPunctuationToken < rankOfSpaceToken) {
648
+ chosenCandidateRank = rankOfPromisingPunctuationToken;
649
+ }
650
+ else {
651
+ // Otherwise, select randomly from top k distribution
652
+ chosenCandidateRank = this.randomGen.selectRandomIndexFromDistribution(topCandidateProbabilities);
653
+ }
654
+ // Add chosen token
655
+ const chosenToken = topCandidates[chosenCandidateRank].token;
656
+ const chosenTokenConfidence = topCandidateProbabilities[chosenCandidateRank];
657
+ addToken(chosenToken, timestampTokenLogits, chosenTokenConfidence, crossAttentionQKsForToken);
658
+ // If chosen token is the end-of-text token, break
659
+ if (chosenToken === endOfTextToken) {
660
+ break;
661
+ }
662
+ // Print token if needed
663
+ if (this.isTextToken(chosenToken)) {
664
+ bufferedTokensToPrint.push(chosenToken);
665
+ let textToPrint = this.tokensToText(bufferedTokensToPrint);
666
+ // If the decoded text is valid, print it
667
+ if (!containsInvalidCodepoint(textToPrint)) {
668
+ if (isFirstPart && decodedTokens.every(token => this.isMetadataToken(token))) {
669
+ textToPrint = textToPrint.trimStart();
670
+ }
671
+ logger.write(textToPrint);
672
+ bufferedTokensToPrint = [];
502
673
  }
503
674
  }
504
675
  await yieldToEventLoop();
505
676
  }
506
- if (timestampsSeenCount >= 2 && !isFinalPart) {
507
- decodedTokens = decodedTokens.slice(0, lastTimestampTokenIndex);
508
- decodedTokensTimestampLogits = decodedTokensTimestampLogits.slice(0, lastTimestampTokenIndex);
509
- decodedTokensCrossAttentionQKs = decodedTokensCrossAttentionQKs.slice(0, lastTimestampTokenIndex);
677
+ // If at least two timestamp tokens were decoded and it's not the final part,
678
+ // truncate up to the last timestamp token
679
+ if (timestampTokenSeenCount >= 2 && !isFinalPart) {
680
+ const sliceEndTokenIndex = lastTimestampTokenIndex;
681
+ decodedTokens = decodedTokens.slice(0, sliceEndTokenIndex);
682
+ decodedTokensTimestampLogits = decodedTokensTimestampLogits.slice(0, sliceEndTokenIndex);
683
+ decodedTokensCrossAttentionQKs = decodedTokensCrossAttentionQKs.slice(0, sliceEndTokenIndex);
684
+ decodedTokensConfidence = decodedTokensConfidence.slice(0, sliceEndTokenIndex);
510
685
  }
511
686
  logger.write('\n');
512
687
  logger.end();
513
- // Return the tokens
688
+ // Return the decoded tokens
514
689
  return {
515
690
  decodedTokens,
516
691
  decodedTokensTimestampLogits,
517
- crossAttentionQKs: decodedTokensCrossAttentionQKs,
518
- decodedTokensConfidence
692
+ decodedTokensConfidence,
693
+ decodedTokensCrossAttentionQKs,
519
694
  };
520
695
  }
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
- }
696
+ // Encode audio using the encoder model
554
697
  async encodeAudio(rawAudio) {
555
698
  await this.initializeEncoderSessionIfNeeded();
556
699
  const Onnx = await import('onnxruntime-node');
@@ -562,13 +705,19 @@ export class Whisper {
562
705
  const filterbankCount = 80;
563
706
  const maxAudioSamples = sampleRate * 30;
564
707
  const maxAudioFrames = 3000;
708
+ if (audioSamples.length > maxAudioSamples) {
709
+ throw new Error(`Audio part is longer than 30 seconds`);
710
+ }
711
+ // Compute a mel spectogram
565
712
  await logger.startAsync('Extract mel spectogram from audio part');
713
+ // Pad audio samples to ensure that have a duration of 30 seconds
566
714
  const paddedAudioSamples = new Float32Array(maxAudioSamples);
567
- paddedAudioSamples.set(audioSamples.subarray(0, maxAudioSamples), 0);
715
+ paddedAudioSamples.set(audioSamples, 0);
568
716
  const rawAudioPart = { audioChannels: [paddedAudioSamples], sampleRate };
569
717
  const { melSpectogram } = await computeMelSpectogramUsingFilterbanks(rawAudioPart, fftOrder, fftOrder, hopLength, filterbanks);
570
718
  await logger.startAsync('Normalize mel spectogram');
571
719
  const logMelSpectogram = melSpectogram.map(spectrum => spectrum.map(mel => Math.log10(Math.max(mel, 1e-10))));
720
+ // Find maximum log mel value in the spectrum
572
721
  let maxLogMel = -Infinity;
573
722
  for (const spectrum of logMelSpectogram) {
574
723
  for (const mel of spectrum) {
@@ -577,13 +726,16 @@ export class Whisper {
577
726
  }
578
727
  }
579
728
  }
729
+ // Normalize log mel spectogram (based on Python reference code)
580
730
  const normalizedLogMelSpectogram = logMelSpectogram.map(spectrum => spectrum.map(logMel => (Math.max(logMel, maxLogMel - 8) + 4) / 4));
731
+ // Flatten the normalized log mel spectogram
581
732
  const flattenedNormalizedLogMelSpectogram = new Float32Array(maxAudioFrames * filterbankCount);
582
733
  for (let i = 0; i < filterbankCount; i++) {
583
734
  for (let j = 0; j < maxAudioFrames; j++) {
584
735
  flattenedNormalizedLogMelSpectogram[(i * maxAudioFrames) + j] = normalizedLogMelSpectogram[j][i];
585
736
  }
586
737
  }
738
+ // Run the encoder model
587
739
  await logger.startAsync('Encode mel spectogram with Whisper encoder model');
588
740
  const inputTensor = new Onnx.Tensor('float32', flattenedNormalizedLogMelSpectogram, [1, filterbankCount, maxAudioFrames]);
589
741
  const encoderInputs = { mel: inputTensor };
@@ -592,99 +744,6 @@ export class Whisper {
592
744
  logger.end();
593
745
  return encodedAudioFeatures;
594
746
  }
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
747
  tokenTimelineToWordTimeline(tokenTimeline, language) {
689
748
  function isSeparatorCharacter(char) {
690
749
  const nonSeparatingPunctuation = [`'`, `-`, `.`, `·`, `•`];
@@ -699,6 +758,9 @@ export class Whisper {
699
758
  function endsWithSeparatorCharacter(text) {
700
759
  return isSeparatorCharacter(text[text.length - 1]);
701
760
  }
761
+ if (language != 'zh' && language != 'ja') {
762
+ tokenTimeline = tokenTimeline.filter(entry => this.isTextToken(entry.id));
763
+ }
702
764
  const resultTimeline = [];
703
765
  let groups = [];
704
766
  for (let tokenIndex = 0; tokenIndex < tokenTimeline.length; tokenIndex++) {
@@ -821,35 +883,35 @@ export class Whisper {
821
883
  const applySoftmax = true;
822
884
  const normalize = true;
823
885
  const applyMedianFilter = true;
824
- const fixateTimestampTokens = false;
886
+ const anchorTimestampTokens = false;
825
887
  const softmaxTemperature = 1.0;
826
- const medianFilterWidth = 7;
888
+ // Apply softmax to each token's frames, if enabled
827
889
  if (applySoftmax) {
828
- // Apply softmax to each token's frames
829
890
  for (const head of attentionHeads) {
830
891
  for (let tokenIndex = 0; tokenIndex < tokenCount; tokenIndex++) {
831
892
  head[tokenIndex] = softmax(head[tokenIndex], softmaxTemperature);
832
893
  }
833
894
  }
834
895
  }
896
+ // Normalize all weights in each individual head, if enabled
835
897
  if (normalize) {
836
- // Normalize all weights in each individual head
837
898
  for (const head of attentionHeads) {
838
899
  const allWeightsForHead = head.flatMap(tokenFrames => tokenFrames);
839
- const meanOfAllWeights = meanOfVector(allWeightsForHead);
840
- const stdDeviationOfAllWeights = stdDeviationOfVector(allWeightsForHead) + 1e-10;
900
+ const meanOfAllWeightsForHead = meanOfVector(allWeightsForHead);
901
+ const stdDeviationOfAllWeightsForHead = stdDeviationOfVector(allWeightsForHead, 'population', meanOfAllWeightsForHead) + 1e-10;
902
+ const stdDeviationReciprocal = 1.0 / (stdDeviationOfAllWeightsForHead + 1e-10);
841
903
  for (const tokenFrames of head) {
842
904
  for (let frameIndex = 0; frameIndex < tokenFrames.length; frameIndex++) {
843
- tokenFrames[frameIndex] = (tokenFrames[frameIndex] - meanOfAllWeights) / stdDeviationOfAllWeights;
905
+ tokenFrames[frameIndex] = (tokenFrames[frameIndex] - meanOfAllWeightsForHead) * stdDeviationReciprocal;
844
906
  }
845
907
  }
846
908
  }
847
909
  }
910
+ // Apply median filter to each token's frames, if enabled
848
911
  if (applyMedianFilter) {
849
- // Apply median filter to each token's frames
850
912
  for (const head of attentionHeads) {
851
913
  for (let tokenIndex = 0; tokenIndex < tokenCount; tokenIndex++) {
852
- head[tokenIndex] = medianFilter(head[tokenIndex], medianFilterWidth);
914
+ head[tokenIndex] = medianOf5Filter(head[tokenIndex]);
853
915
  }
854
916
  }
855
917
  }
@@ -869,8 +931,8 @@ export class Whisper {
869
931
  frameMeansForToken[tokenIndex][frameIndex] = frameMean;
870
932
  }
871
933
  }
872
- if (fixateTimestampTokens) {
873
- // Fixate timestamp tokens to the original ones detected
934
+ // Anchor timestamp tokens timestamps to their original values, if enabled
935
+ if (anchorTimestampTokens) {
874
936
  const timestampTokensStart = this.tokenConfig.timestampTokensStart;
875
937
  for (let tokenIndex = 0; tokenIndex < tokens.length; tokenIndex++) {
876
938
  const token = tokens[tokenIndex];
@@ -882,14 +944,86 @@ export class Whisper {
882
944
  }
883
945
  }
884
946
  // Perform DTW
885
- const tokenIndexes = [...Array(tokenCount).keys()];
886
- const frameIndexes = [...Array(segmentFrameCount).keys()];
947
+ const tokenIndexes = getIntegerRange(0, tokenCount);
948
+ const frameIndexes = getIntegerRange(0, segmentFrameCount);
887
949
  let { path } = alignDTWWindowed(tokenIndexes, frameIndexes, (tokenIndex, frameIndex) => {
888
950
  return -frameMeansForToken[tokenIndex][frameIndex];
889
951
  }, segmentFrameCount);
890
952
  path = path.map(entry => ({ source: entry.source, dest: segmentStartFrame + entry.dest }));
891
953
  return path;
892
954
  }
955
+ async initializeIfNeeded() {
956
+ await this.initializeTokenizerIfNeeded();
957
+ await this.initializeEncoderSessionIfNeeded();
958
+ await this.initializeDecoderSessionIfNeeded();
959
+ }
960
+ async initializeTokenizerIfNeeded() {
961
+ if (this.tiktoken) {
962
+ return;
963
+ }
964
+ const logger = new Logger();
965
+ await logger.startAsync('Load tokenizer data');
966
+ const tiktokenModulePackagePath = await loadPackage('whisper-tiktoken-data');
967
+ const tiktokenDataFilePath = path.join(tiktokenModulePackagePath, this.isMultiligualModel ? 'multilingual.tiktoken' : 'gpt2.tiktoken');
968
+ let tiktokenData = await readFile(tiktokenDataFilePath, { encoding: 'utf8' });
969
+ const tokenConfig = this.tokenConfig;
970
+ const metadataTokens = {
971
+ [tokenConfig.endOfTextToken]: '[EndOfText]',
972
+ [tokenConfig.startOfTextToken]: '[StartOfText]',
973
+ [tokenConfig.translateTaskToken]: '[TranslateTask]',
974
+ [tokenConfig.transcribeTaskToken]: '[TranscribeTask]',
975
+ [tokenConfig.startOfPromptToken]: '[StartOfPrompt]',
976
+ [tokenConfig.nonSpeechToken]: '[NonSpeech]',
977
+ [tokenConfig.noTimestampsToken]: '[NoTimestamps]',
978
+ };
979
+ if (this.isMultiligualModel) {
980
+ metadataTokens[50256] = '[Unused_50256]';
981
+ metadataTokens[50360] = '[Unused_50360]';
982
+ }
983
+ const languageTokenCount = tokenConfig.languageTokensEnd - tokenConfig.languageTokensStart;
984
+ for (let i = 0; i < languageTokenCount; i++) {
985
+ const tokenIndex = this.tokenConfig.languageTokensStart + i;
986
+ metadataTokens[tokenIndex] = `[Language_${i}]`;
987
+ }
988
+ const timestampTokensCount = 1501;
989
+ for (let i = 0; i < timestampTokensCount; i++) {
990
+ const tokenIndex = this.tokenConfig.timestampTokensStart + i;
991
+ const tokenTime = this.timestampTokenToSeconds(tokenIndex);
992
+ metadataTokens[tokenIndex] = `[Timestamp_${tokenTime.toFixed(2)}]`;
993
+ }
994
+ const inverseMetadataTokensLookup = {};
995
+ for (const [key, value] of Object.entries(metadataTokens)) {
996
+ inverseMetadataTokensLookup[value] = parseInt(key);
997
+ }
998
+ const patternString = `'s|'t|'re|'ve|'m|'ll|'d| ?\\p{L}+| ?\\p{N}+| ?[^\\s\\p{L}\\p{N}]+|\\s+(?!\\S)|\\s+`;
999
+ const { Tiktoken } = await import('tiktoken/lite');
1000
+ this.tiktoken = new Tiktoken(tiktokenData, inverseMetadataTokensLookup, patternString);
1001
+ logger.end();
1002
+ }
1003
+ async initializeEncoderSessionIfNeeded() {
1004
+ if (this.audioEncoder) {
1005
+ return;
1006
+ }
1007
+ const logger = new Logger();
1008
+ await logger.startAsync(`Create encoder inference session for model '${this.modelName}'`);
1009
+ const encoderFilePath = path.join(this.modelDir, 'encoder.onnx');
1010
+ const Onnx = await import('onnxruntime-node');
1011
+ const onnxSessionOptions = getOnnxSessionOptions({ executionProviders: this.encoderExecutionProviders });
1012
+ this.audioEncoder = await Onnx.InferenceSession.create(encoderFilePath, onnxSessionOptions);
1013
+ logger.end();
1014
+ }
1015
+ async initializeDecoderSessionIfNeeded() {
1016
+ if (this.textDecoder) {
1017
+ return;
1018
+ }
1019
+ const logger = new Logger();
1020
+ await logger.startAsync(`Create decoder inference session for model '${this.modelName}'`);
1021
+ const decoderFilePath = path.join(this.modelDir, 'decoder.onnx');
1022
+ const Onnx = await import('onnxruntime-node');
1023
+ const onnxSessionOptions = getOnnxSessionOptions({ executionProviders: this.decoderExecutionProviders });
1024
+ this.textDecoder = await Onnx.InferenceSession.create(decoderFilePath, onnxSessionOptions);
1025
+ logger.end();
1026
+ }
893
1027
  getKvDimensions(groupCount, length) {
894
1028
  const modelName = this.modelName;
895
1029
  if (modelName == 'tiny' || modelName == 'tiny.en') {
@@ -1010,7 +1144,7 @@ export class Whisper {
1010
1144
  }
1011
1145
  getSuppressedTextTokens() {
1012
1146
  const allowedPunctuationMarks = this.getAllowedPunctuationMarks();
1013
- const nonWordTokensData = this.getNonWordTokenData();
1147
+ const nonWordTokensData = this.getWordTokenData().nonWordTokenData;
1014
1148
  const suppressedTextTokens = nonWordTokensData
1015
1149
  .filter(entry => !allowedPunctuationMarks.includes(entry.text))
1016
1150
  .map(entry => entry.id);
@@ -1039,22 +1173,27 @@ export class Whisper {
1039
1173
  }
1040
1174
  return allowedPunctuation;
1041
1175
  }
1042
- getNonWordTokenData() {
1176
+ getWordTokenData() {
1177
+ const wordTokenData = [];
1043
1178
  const nonWordTokenData = [];
1044
- const invalidUTF8Char = String.fromCharCode(65533);
1045
1179
  for (let i = 0; i < this.tokenConfig.endOfTextToken; i++) {
1046
1180
  const tokenText = this.tokenToText(i, false);
1047
- const tokenTextWithoutWhitespace = tokenText.replaceAll(/\s/g, '');
1048
- const isNonWordToken = /^[\p{Punctuation}\p{Symbol}]+$/u.test(tokenTextWithoutWhitespace);
1049
- const containsInvalidUTF8 = getUTF32Chars(tokenTextWithoutWhitespace).utf32chars.includes(invalidUTF8Char);
1050
- if (isNonWordToken && !containsInvalidUTF8) {
1181
+ const isNonWordToken = /^[\s\p{Punctuation}\p{Symbol}]+$/u.test(tokenText);
1182
+ const containsInvalidUTF8 = containsInvalidCodepoint(tokenText);
1183
+ if (isNonWordToken && (this.isEnglishOnlyModel || !containsInvalidUTF8)) {
1051
1184
  nonWordTokenData.push({
1052
1185
  id: i,
1053
1186
  text: tokenText,
1054
1187
  });
1055
1188
  }
1189
+ else {
1190
+ wordTokenData.push({
1191
+ id: i,
1192
+ text: tokenText,
1193
+ });
1194
+ }
1056
1195
  }
1057
- return nonWordTokenData;
1196
+ return { wordTokenData, nonWordTokenData };
1058
1197
  }
1059
1198
  getTokensData(tokens) {
1060
1199
  const tokensData = [];
@@ -1067,88 +1206,6 @@ export class Whisper {
1067
1206
  return tokensData;
1068
1207
  }
1069
1208
  }
1070
- const filterbanks = [
1071
- /* 0 */ { startIndex: 1, weights: [0.02486259490251541,] },
1072
- /* 1 */ { startIndex: 1, weights: [0.001990821911022067, 0.022871771827340126,] },
1073
- /* 2 */ { startIndex: 2, weights: [0.003981643822044134, 0.02088095061480999,] },
1074
- /* 3 */ { startIndex: 3, weights: [0.0059724655002355576, 0.018890129402279854,] },
1075
- /* 4 */ { startIndex: 4, weights: [0.007963287644088268, 0.01689930632710457,] },
1076
- /* 5 */ { startIndex: 5, weights: [0.009954108856618404, 0.014908484183251858,] },
1077
- /* 6 */ { startIndex: 6, weights: [0.011944931000471115, 0.012917662039399147,] },
1078
- /* 7 */ { startIndex: 7, weights: [0.013935752213001251, 0.010926840826869011,] },
1079
- /* 8 */ { startIndex: 8, weights: [0.015926575288176537, 0.0089360186830163,] },
1080
- /* 9 */ { startIndex: 9, weights: [0.017917396500706673, 0.006945197004824877,] },
1081
- /* 10 */ { startIndex: 10, weights: [0.01990821771323681, 0.004954374860972166,] },
1082
- /* 11 */ { startIndex: 11, weights: [0.021899040788412094, 0.0029635531827807426,] },
1083
- /* 12 */ { startIndex: 12, weights: [0.02388986200094223, 0.0009727313299663365,] },
1084
- /* 13 */ { startIndex: 13, weights: [0.025880683213472366,] },
1085
- /* 14 */ { startIndex: 14, weights: [0.025835324078798294,] },
1086
- /* 15 */ { startIndex: 14, weights: [0.0010180906392633915, 0.023844502866268158,] },
1087
- /* 16 */ { startIndex: 15, weights: [0.003008912317454815, 0.021853681653738022,] },
1088
- /* 17 */ { startIndex: 16, weights: [0.004999734461307526, 0.019862858578562737,] },
1089
- /* 18 */ { startIndex: 17, weights: [0.006990555673837662, 0.0178720373660326,] },
1090
- /* 19 */ { startIndex: 18, weights: [0.008981377817690372, 0.015881216153502464,] },
1091
- /* 20 */ { startIndex: 19, weights: [0.010972199961543083, 0.013890394009649754,] },
1092
- /* 21 */ { startIndex: 20, weights: [0.01296302117407322, 0.011899571865797043,] },
1093
- /* 22 */ { startIndex: 21, weights: [0.01495384331792593, 0.009908749721944332,] },
1094
- /* 23 */ { startIndex: 22, weights: [0.01694466546177864, 0.007917927578091621,] },
1095
- /* 24 */ { startIndex: 23, weights: [0.018935488536953926, 0.005927106365561485,] },
1096
- /* 25 */ { startIndex: 24, weights: [0.020874010398983955, 0.004040425643324852,] },
1097
- /* 26 */ { startIndex: 25, weights: [0.022114217281341553, 0.0033186059445142746,] },
1098
- /* 27 */ { startIndex: 26, weights: [0.02173672430217266, 0.0036109676584601402,] },
1099
- /* 28 */ { startIndex: 27, weights: [0.020497702062129974, 0.004762193653732538,] },
1100
- /* 29 */ { startIndex: 28, weights: [0.018486659973859787, 0.006592618301510811,] },
1101
- /* 30 */ { startIndex: 29, weights: [0.01585603691637516, 0.00896277092397213,] },
1102
- /* 31 */ { startIndex: 30, weights: [0.012738768011331558, 0.011751330457627773,] },
1103
- /* 32 */ { startIndex: 31, weights: [0.009250369854271412, 0.014853144995868206,] },
1104
- /* 33 */ { startIndex: 32, weights: [0.005490840878337622, 0.018177473917603493, 0.0028155462350696325,] },
1105
- /* 34 */ { startIndex: 33, weights: [0.0015463664894923568, 0.01632951945066452, 0.007420188747346401,] },
1106
- /* 35 */ { startIndex: 35, weights: [0.011181050911545753, 0.012018864043056965,] },
1107
- /* 36 */ { startIndex: 36, weights: [0.006065350491553545, 0.016561277210712433, 0.004360878840088844,] },
1108
- /* 37 */ { startIndex: 37, weights: [0.0010297985281795263, 0.012770536355674267, 0.009707189165055752,] },
1109
- /* 38 */ { startIndex: 39, weights: [0.006986402906477451, 0.01485429983586073, 0.004391219466924667,] },
1110
- /* 39 */ { startIndex: 40, weights: [0.001418047584593296, 0.011486922390758991, 0.010089744813740253, 0.00040022286702878773,] },
1111
- /* 40 */ { startIndex: 42, weights: [0.005411104764789343, 0.014735566452145576, 0.006518189795315266,] },
1112
- /* 41 */ { startIndex: 44, weights: [0.00827841367572546, 0.012277561239898205, 0.00396781275048852,] },
1113
- /* 42 */ { startIndex: 45, weights: [0.002187808509916067, 0.010184479877352715, 0.00998187530785799, 0.0022864851634949446,] },
1114
- /* 43 */ { startIndex: 47, weights: [0.00386943481862545, 0.011274894699454308, 0.008466221392154694, 0.0013397691072896123,] },
1115
- /* 44 */ { startIndex: 49, weights: [0.004820294212549925, 0.011678251437842846, 0.007608682848513126, 0.0010091039584949613,] },
1116
- /* 45 */ { startIndex: 51, weights: [0.005156961735337973, 0.011507894843816757, 0.007301822770386934, 0.0011901655234396458,] },
1117
- /* 46 */ { startIndex: 53, weights: [0.004982104524970055, 0.010863498784601688, 0.007451189681887627, 0.001791381393559277,] },
1118
- /* 47 */ { startIndex: 55, weights: [0.004385921638458967, 0.009832492098212242, 0.007973956875503063, 0.002732589840888977,] },
1119
- /* 48 */ { startIndex: 57, weights: [0.0034474546555429697, 0.008491347543895245, 0.008797688409686089, 0.00394382793456316,] },
1120
- /* 49 */ { startIndex: 59, weights: [0.0022357646375894547, 0.0069067515432834625, 0.009859241545200348, 0.005364237818866968, 0.0008692338014952838,] },
1121
- /* 50 */ { startIndex: 61, weights: [0.0008110002381727099, 0.005136650986969471, 0.00946230161935091, 0.0069410777650773525, 0.0027783995028585196,] },
1122
- /* 51 */ { startIndex: 64, weights: [0.003231203882023692, 0.007237049750983715, 0.00862883497029543, 0.004773912951350212, 0.0009189908159896731,] },
1123
- /* 52 */ { startIndex: 66, weights: [0.001233637798577547, 0.0049433219246566296, 0.008653006516397, 0.006818502210080624, 0.003248583758249879,] },
1124
- /* 53 */ { startIndex: 69, weights: [0.0026164355222135782, 0.006051854696124792, 0.008880467154085636, 0.005574479699134827, 0.002268492942675948,] },
1125
- /* 54 */ { startIndex: 71, weights: [0.0002863667905330658, 0.003467798000201583, 0.0066492292098701, 0.00787146482616663, 0.004809896927326918, 0.0017483289120718837,] },
1126
- /* 55 */ { startIndex: 74, weights: [0.0009245910914614797, 0.0038708120118826628, 0.00681703258305788, 0.007283343467861414, 0.004448124207556248, 0.0016129047144204378,] },
1127
- /* 56 */ { startIndex: 77, weights: [0.0011703289346769452, 0.003898728871718049, 0.006627128925174475, 0.0070473202504217625, 0.004421714693307877, 0.0017961094854399562,] },
1128
- /* 57 */ { startIndex: 80, weights: [0.0010892992140725255, 0.003615982597693801, 0.006142666097730398, 0.007102936040610075, 0.004671447444707155, 0.002239959081634879,] },
1129
- /* 58 */ { startIndex: 83, weights: [0.0007392280967906117, 0.0030791081953793764, 0.005418988410383463, 0.007397185545414686, 0.005145462695509195, 0.002893739379942417, 0.0006420162972062826,] },
1130
- /* 59 */ { startIndex: 86, weights: [0.00017068670422304422, 0.0023375742603093386, 0.004504461772739887, 0.0066713495180010796, 0.005798479542136192, 0.003713231300935149, 0.0016279831761494279,] },
1131
- /* 60 */ { startIndex: 90, weights: [0.0014345343224704266, 0.0034412189852446318, 0.005447904113680124, 0.006591092795133591, 0.004660011734813452, 0.002728930441662669, 0.0007978491485118866,] },
1132
- /* 61 */ { startIndex: 93, weights: [0.0004075043834745884, 0.002265830524265766, 0.004124156199395657, 0.005982482805848122, 0.005700822453945875, 0.003912510350346565, 0.0021241982467472553, 0.0003358862304594368,] },
1133
- /* 62 */ { startIndex: 97, weights: [0.0010099108330905437, 0.002730846870690584, 0.004451782442629337, 0.006172718480229378, 0.005150905344635248, 0.0034948070533573627, 0.0018387088784947991, 0.0001826105872169137,] },
1134
- /* 63 */ { startIndex: 101, weights: [0.0012943691108375788, 0.002888072282075882, 0.004481775686144829, 0.006075479090213776, 0.0048866597935557365, 0.003353001084178686, 0.0018193417927250266, 0.00028568264679051936,] },
1135
- /* 64 */ { startIndex: 105, weights: [0.0013131388695910573, 0.0027890161145478487, 0.004264893010258675, 0.0057407706044614315, 0.004859979264438152, 0.0034397069830447435, 0.0020194342359900475, 0.0005991620710119605,] },
1136
- /* 65 */ { startIndex: 109, weights: [0.0011121684219688177, 0.002478930866345763, 0.0038456933107227087, 0.0052124555222690105, 0.005028639920055866, 0.0037133716978132725, 0.002398103242740035, 0.0010828346712514758,] },
1137
- /* 66 */ { startIndex: 113, weights: [0.0007317548734135926, 0.0019974694587290287, 0.003263183869421482, 0.004528898745775223, 0.005355686880648136, 0.004137659445405006, 0.0029196315445005894, 0.0017016039928421378, 0.0004835762665607035,] },
1138
- /* 67 */ { startIndex: 117, weights: [0.00020713974663522094, 0.0013792773243039846, 0.0025514145381748676, 0.003723552217707038, 0.004895689897239208, 0.004680895246565342, 0.0035529187880456448, 0.0024249425623565912, 0.0012969663366675377, 0.00016899015463422984,] },
1139
- /* 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,] },
1140
- /* 69 */ { startIndex: 127, weights: [0.000854626705404371, 0.001859853626228869, 0.002865080488845706, 0.003870307235047221, 0.00487553421407938, 0.00408313749358058, 0.003115783678367734, 0.0021484296303242445, 0.001181075582280755, 0.0002137213887181133,] },
1141
- /* 70 */ { startIndex: 132, weights: [0.0008483415003865957, 0.0017792496364563704, 0.0027101580053567886, 0.0036410661414265633, 0.004571974277496338, 0.004079728852957487, 0.003183893393725157, 0.002288057701662183, 0.0013922222424298525, 0.0004963868414051831,] },
1142
- /* 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,] },
1143
- /* 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,] },
1144
- /* 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,] },
1145
- /* 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,] },
1146
- /* 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,] },
1147
- /* 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,] },
1148
- /* 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,] },
1149
- /* 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,] },
1150
- /* 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,] },
1151
- ];
1152
1209
  export async function loadPackagesAndGetPaths(modelName, languageCode) {
1153
1210
  if (modelName) {
1154
1211
  modelName = normalizeWhisperModelName(modelName, languageCode);
@@ -1314,6 +1371,88 @@ const alignmentHeadsIndexes = {
1314
1371
  '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)
1315
1372
  'large': [212, 277, 331, 332, 333, 355, 356, 364, 371, 379, 391, 422, 423, 443, 449, 452, 465, 467, 473, 505, 521, 532, 555],
1316
1373
  };
1374
+ const filterbanks = [
1375
+ /* 0 */ { startIndex: 1, weights: [0.02486259490251541,] },
1376
+ /* 1 */ { startIndex: 1, weights: [0.001990821911022067, 0.022871771827340126,] },
1377
+ /* 2 */ { startIndex: 2, weights: [0.003981643822044134, 0.02088095061480999,] },
1378
+ /* 3 */ { startIndex: 3, weights: [0.0059724655002355576, 0.018890129402279854,] },
1379
+ /* 4 */ { startIndex: 4, weights: [0.007963287644088268, 0.01689930632710457,] },
1380
+ /* 5 */ { startIndex: 5, weights: [0.009954108856618404, 0.014908484183251858,] },
1381
+ /* 6 */ { startIndex: 6, weights: [0.011944931000471115, 0.012917662039399147,] },
1382
+ /* 7 */ { startIndex: 7, weights: [0.013935752213001251, 0.010926840826869011,] },
1383
+ /* 8 */ { startIndex: 8, weights: [0.015926575288176537, 0.0089360186830163,] },
1384
+ /* 9 */ { startIndex: 9, weights: [0.017917396500706673, 0.006945197004824877,] },
1385
+ /* 10 */ { startIndex: 10, weights: [0.01990821771323681, 0.004954374860972166,] },
1386
+ /* 11 */ { startIndex: 11, weights: [0.021899040788412094, 0.0029635531827807426,] },
1387
+ /* 12 */ { startIndex: 12, weights: [0.02388986200094223, 0.0009727313299663365,] },
1388
+ /* 13 */ { startIndex: 13, weights: [0.025880683213472366,] },
1389
+ /* 14 */ { startIndex: 14, weights: [0.025835324078798294,] },
1390
+ /* 15 */ { startIndex: 14, weights: [0.0010180906392633915, 0.023844502866268158,] },
1391
+ /* 16 */ { startIndex: 15, weights: [0.003008912317454815, 0.021853681653738022,] },
1392
+ /* 17 */ { startIndex: 16, weights: [0.004999734461307526, 0.019862858578562737,] },
1393
+ /* 18 */ { startIndex: 17, weights: [0.006990555673837662, 0.0178720373660326,] },
1394
+ /* 19 */ { startIndex: 18, weights: [0.008981377817690372, 0.015881216153502464,] },
1395
+ /* 20 */ { startIndex: 19, weights: [0.010972199961543083, 0.013890394009649754,] },
1396
+ /* 21 */ { startIndex: 20, weights: [0.01296302117407322, 0.011899571865797043,] },
1397
+ /* 22 */ { startIndex: 21, weights: [0.01495384331792593, 0.009908749721944332,] },
1398
+ /* 23 */ { startIndex: 22, weights: [0.01694466546177864, 0.007917927578091621,] },
1399
+ /* 24 */ { startIndex: 23, weights: [0.018935488536953926, 0.005927106365561485,] },
1400
+ /* 25 */ { startIndex: 24, weights: [0.020874010398983955, 0.004040425643324852,] },
1401
+ /* 26 */ { startIndex: 25, weights: [0.022114217281341553, 0.0033186059445142746,] },
1402
+ /* 27 */ { startIndex: 26, weights: [0.02173672430217266, 0.0036109676584601402,] },
1403
+ /* 28 */ { startIndex: 27, weights: [0.020497702062129974, 0.004762193653732538,] },
1404
+ /* 29 */ { startIndex: 28, weights: [0.018486659973859787, 0.006592618301510811,] },
1405
+ /* 30 */ { startIndex: 29, weights: [0.01585603691637516, 0.00896277092397213,] },
1406
+ /* 31 */ { startIndex: 30, weights: [0.012738768011331558, 0.011751330457627773,] },
1407
+ /* 32 */ { startIndex: 31, weights: [0.009250369854271412, 0.014853144995868206,] },
1408
+ /* 33 */ { startIndex: 32, weights: [0.005490840878337622, 0.018177473917603493, 0.0028155462350696325,] },
1409
+ /* 34 */ { startIndex: 33, weights: [0.0015463664894923568, 0.01632951945066452, 0.007420188747346401,] },
1410
+ /* 35 */ { startIndex: 35, weights: [0.011181050911545753, 0.012018864043056965,] },
1411
+ /* 36 */ { startIndex: 36, weights: [0.006065350491553545, 0.016561277210712433, 0.004360878840088844,] },
1412
+ /* 37 */ { startIndex: 37, weights: [0.0010297985281795263, 0.012770536355674267, 0.009707189165055752,] },
1413
+ /* 38 */ { startIndex: 39, weights: [0.006986402906477451, 0.01485429983586073, 0.004391219466924667,] },
1414
+ /* 39 */ { startIndex: 40, weights: [0.001418047584593296, 0.011486922390758991, 0.010089744813740253, 0.00040022286702878773,] },
1415
+ /* 40 */ { startIndex: 42, weights: [0.005411104764789343, 0.014735566452145576, 0.006518189795315266,] },
1416
+ /* 41 */ { startIndex: 44, weights: [0.00827841367572546, 0.012277561239898205, 0.00396781275048852,] },
1417
+ /* 42 */ { startIndex: 45, weights: [0.002187808509916067, 0.010184479877352715, 0.00998187530785799, 0.0022864851634949446,] },
1418
+ /* 43 */ { startIndex: 47, weights: [0.00386943481862545, 0.011274894699454308, 0.008466221392154694, 0.0013397691072896123,] },
1419
+ /* 44 */ { startIndex: 49, weights: [0.004820294212549925, 0.011678251437842846, 0.007608682848513126, 0.0010091039584949613,] },
1420
+ /* 45 */ { startIndex: 51, weights: [0.005156961735337973, 0.011507894843816757, 0.007301822770386934, 0.0011901655234396458,] },
1421
+ /* 46 */ { startIndex: 53, weights: [0.004982104524970055, 0.010863498784601688, 0.007451189681887627, 0.001791381393559277,] },
1422
+ /* 47 */ { startIndex: 55, weights: [0.004385921638458967, 0.009832492098212242, 0.007973956875503063, 0.002732589840888977,] },
1423
+ /* 48 */ { startIndex: 57, weights: [0.0034474546555429697, 0.008491347543895245, 0.008797688409686089, 0.00394382793456316,] },
1424
+ /* 49 */ { startIndex: 59, weights: [0.0022357646375894547, 0.0069067515432834625, 0.009859241545200348, 0.005364237818866968, 0.0008692338014952838,] },
1425
+ /* 50 */ { startIndex: 61, weights: [0.0008110002381727099, 0.005136650986969471, 0.00946230161935091, 0.0069410777650773525, 0.0027783995028585196,] },
1426
+ /* 51 */ { startIndex: 64, weights: [0.003231203882023692, 0.007237049750983715, 0.00862883497029543, 0.004773912951350212, 0.0009189908159896731,] },
1427
+ /* 52 */ { startIndex: 66, weights: [0.001233637798577547, 0.0049433219246566296, 0.008653006516397, 0.006818502210080624, 0.003248583758249879,] },
1428
+ /* 53 */ { startIndex: 69, weights: [0.0026164355222135782, 0.006051854696124792, 0.008880467154085636, 0.005574479699134827, 0.002268492942675948,] },
1429
+ /* 54 */ { startIndex: 71, weights: [0.0002863667905330658, 0.003467798000201583, 0.0066492292098701, 0.00787146482616663, 0.004809896927326918, 0.0017483289120718837,] },
1430
+ /* 55 */ { startIndex: 74, weights: [0.0009245910914614797, 0.0038708120118826628, 0.00681703258305788, 0.007283343467861414, 0.004448124207556248, 0.0016129047144204378,] },
1431
+ /* 56 */ { startIndex: 77, weights: [0.0011703289346769452, 0.003898728871718049, 0.006627128925174475, 0.0070473202504217625, 0.004421714693307877, 0.0017961094854399562,] },
1432
+ /* 57 */ { startIndex: 80, weights: [0.0010892992140725255, 0.003615982597693801, 0.006142666097730398, 0.007102936040610075, 0.004671447444707155, 0.002239959081634879,] },
1433
+ /* 58 */ { startIndex: 83, weights: [0.0007392280967906117, 0.0030791081953793764, 0.005418988410383463, 0.007397185545414686, 0.005145462695509195, 0.002893739379942417, 0.0006420162972062826,] },
1434
+ /* 59 */ { startIndex: 86, weights: [0.00017068670422304422, 0.0023375742603093386, 0.004504461772739887, 0.0066713495180010796, 0.005798479542136192, 0.003713231300935149, 0.0016279831761494279,] },
1435
+ /* 60 */ { startIndex: 90, weights: [0.0014345343224704266, 0.0034412189852446318, 0.005447904113680124, 0.006591092795133591, 0.004660011734813452, 0.002728930441662669, 0.0007978491485118866,] },
1436
+ /* 61 */ { startIndex: 93, weights: [0.0004075043834745884, 0.002265830524265766, 0.004124156199395657, 0.005982482805848122, 0.005700822453945875, 0.003912510350346565, 0.0021241982467472553, 0.0003358862304594368,] },
1437
+ /* 62 */ { startIndex: 97, weights: [0.0010099108330905437, 0.002730846870690584, 0.004451782442629337, 0.006172718480229378, 0.005150905344635248, 0.0034948070533573627, 0.0018387088784947991, 0.0001826105872169137,] },
1438
+ /* 63 */ { startIndex: 101, weights: [0.0012943691108375788, 0.002888072282075882, 0.004481775686144829, 0.006075479090213776, 0.0048866597935557365, 0.003353001084178686, 0.0018193417927250266, 0.00028568264679051936,] },
1439
+ /* 64 */ { startIndex: 105, weights: [0.0013131388695910573, 0.0027890161145478487, 0.004264893010258675, 0.0057407706044614315, 0.004859979264438152, 0.0034397069830447435, 0.0020194342359900475, 0.0005991620710119605,] },
1440
+ /* 65 */ { startIndex: 109, weights: [0.0011121684219688177, 0.002478930866345763, 0.0038456933107227087, 0.0052124555222690105, 0.005028639920055866, 0.0037133716978132725, 0.002398103242740035, 0.0010828346712514758,] },
1441
+ /* 66 */ { startIndex: 113, weights: [0.0007317548734135926, 0.0019974694587290287, 0.003263183869421482, 0.004528898745775223, 0.005355686880648136, 0.004137659445405006, 0.0029196315445005894, 0.0017016039928421378, 0.0004835762665607035,] },
1442
+ /* 67 */ { startIndex: 117, weights: [0.00020713974663522094, 0.0013792773243039846, 0.0025514145381748676, 0.003723552217707038, 0.004895689897239208, 0.004680895246565342, 0.0035529187880456448, 0.0024249425623565912, 0.0012969663366675377, 0.00016899015463422984,] },
1443
+ /* 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,] },
1444
+ /* 69 */ { startIndex: 127, weights: [0.000854626705404371, 0.001859853626228869, 0.002865080488845706, 0.003870307235047221, 0.00487553421407938, 0.00408313749358058, 0.003115783678367734, 0.0021484296303242445, 0.001181075582280755, 0.0002137213887181133,] },
1445
+ /* 70 */ { startIndex: 132, weights: [0.0008483415003865957, 0.0017792496364563704, 0.0027101580053567886, 0.0036410661414265633, 0.004571974277496338, 0.004079728852957487, 0.003183893393725157, 0.002288057701662183, 0.0013922222424298525, 0.0004963868414051831,] },
1446
+ /* 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,] },
1447
+ /* 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,] },
1448
+ /* 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,] },
1449
+ /* 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,] },
1450
+ /* 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,] },
1451
+ /* 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,] },
1452
+ /* 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,] },
1453
+ /* 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,] },
1454
+ /* 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,] },
1455
+ ];
1317
1456
  export const defaultWhisperOptions = {
1318
1457
  model: undefined,
1319
1458
  temperature: 0.1,
@@ -1324,6 +1463,28 @@ export const defaultWhisperOptions = {
1324
1463
  maxTokensPerPart: 250,
1325
1464
  suppressRepetition: true,
1326
1465
  decodeTimestampTokens: true,
1466
+ endTokenThreshold: 0.9,
1467
+ includeEndTokenInCandidates: true,
1468
+ encoderProvider: undefined,
1469
+ decoderProvider: undefined,
1327
1470
  seed: undefined,
1328
1471
  };
1472
+ export const defaultWhisperAlignmentOptions = {
1473
+ model: undefined,
1474
+ endTokenThreshold: 0.9,
1475
+ encoderProvider: undefined,
1476
+ decoderProvider: undefined
1477
+ };
1478
+ export const defaultWhisperLanguageDetectionOptions = {
1479
+ model: undefined,
1480
+ temperature: 1.0,
1481
+ encoderProvider: undefined,
1482
+ decoderProvider: undefined,
1483
+ };
1484
+ export const defaultWhisperVADOptions = {
1485
+ model: undefined,
1486
+ temperature: 1.0,
1487
+ encoderProvider: undefined,
1488
+ decoderProvider: undefined,
1489
+ };
1329
1490
  //# sourceMappingURL=WhisperSTT.js.map