diff --git a/packages/transformers/src/generation/logits_process.js b/packages/transformers/src/generation/logits_process.js index 647a30806..473c4f798 100644 --- a/packages/transformers/src/generation/logits_process.js +++ b/packages/transformers/src/generation/logits_process.js @@ -304,7 +304,6 @@ export class WhisperTimeStampLogitsProcessor extends LogitsProcessor { if (input_ids[i].length === this.begin_index) { batch_logits_data.subarray(0, this.timestamp_begin).fill(-Infinity); - continue; } // timestamps have to appear in pairs, except directly before eos_token; mask logits accordingly diff --git a/packages/transformers/src/models/whisper/modeling_whisper.js b/packages/transformers/src/models/whisper/modeling_whisper.js index d0d6f0b5a..9b4f9689a 100644 --- a/packages/transformers/src/models/whisper/modeling_whisper.js +++ b/packages/transformers/src/models/whisper/modeling_whisper.js @@ -201,7 +201,9 @@ export class WhisperForConditionalGeneration extends WhisperPreTrainedModel { // input_features shape: [batch=1, n_mels, total_frames] const input_features = inputs; - const total_frames = input_features.dims[2]; + const total_frames = generation_config.num_frames + ? Math.min(input_features.dims[2], generation_config.num_frames) + : input_features.dims[2]; // The encoder downsamples by input_stride (=2 for whisper), so: // num_segment_frames = input_stride * max_source_positions = 3000 mel frames per segment @@ -343,6 +345,7 @@ export class WhisperForConditionalGeneration extends WhisperPreTrainedModel { if (seek_token_timestamps) { allTokenTimestamps.push(...seek_token_timestamps.slice(0, tokens_to_keep)); } + if (segment_offset <= 0) break; seek += segment_offset; }