fix(batch): add temperature fallback to prevent decoder runaway
Signed-off-by: makaveli10 <vineet.suryan@collabora.com>
This commit is contained in:
@@ -220,6 +220,7 @@ class ServeClientFasterWhisper(ServeClientBase):
|
|||||||
use_vad=self.use_vad,
|
use_vad=self.use_vad,
|
||||||
vad_parameters=self.vad_parameters if self.use_vad else None,
|
vad_parameters=self.vad_parameters if self.use_vad else None,
|
||||||
word_timestamps=self.word_timestamps,
|
word_timestamps=self.word_timestamps,
|
||||||
|
client_uid=self.client_uid,
|
||||||
)
|
)
|
||||||
ServeClientFasterWhisper.BATCH_WORKER.submit(request)
|
ServeClientFasterWhisper.BATCH_WORKER.submit(request)
|
||||||
request.future.wait(timeout=30)
|
request.future.wait(timeout=30)
|
||||||
|
|||||||
@@ -74,6 +74,8 @@ class BatchRequest:
|
|||||||
initial_prompt: Optional[str] = None
|
initial_prompt: Optional[str] = None
|
||||||
use_vad: bool = True
|
use_vad: bool = True
|
||||||
vad_parameters: Optional[Dict] = None
|
vad_parameters: Optional[Dict] = None
|
||||||
|
word_timestamps: bool = False
|
||||||
|
client_uid: Optional[str] = None
|
||||||
# Signaling
|
# Signaling
|
||||||
future: threading.Event = field(default_factory=threading.Event)
|
future: threading.Event = field(default_factory=threading.Event)
|
||||||
# Results (filled by batch worker)
|
# Results (filled by batch worker)
|
||||||
@@ -307,36 +309,87 @@ class BatchInferenceWorker:
|
|||||||
tokenizers_list.append(tokenizer)
|
tokenizers_list.append(tokenizer)
|
||||||
prompts.append(prompt)
|
prompts.append(prompt)
|
||||||
|
|
||||||
# Step 4: Batch GPU generate
|
# Step 4: Batch GPU generate with per-item temperature fallback.
|
||||||
|
# Mirrors faster_whisper.transcribe()'s fallback loop. Items that
|
||||||
|
# pass quality thresholds at lower temperature keep their result;
|
||||||
|
# only failed items are re-decoded at the next temperature.
|
||||||
suppress_tokens = get_suppressed_tokens(tokenizers_list[0], [-1])
|
suppress_tokens = get_suppressed_tokens(tokenizers_list[0], [-1])
|
||||||
|
|
||||||
results = self.transcriber.model.generate(
|
temperatures = [0.0, 0.2, 0.4, 0.6, 0.8, 1.0]
|
||||||
encoder_output,
|
comp_thresh = 2.4
|
||||||
prompts,
|
logprob_thresh = -1.0
|
||||||
beam_size=5,
|
no_speech_thresh = 0.6
|
||||||
patience=1,
|
|
||||||
length_penalty=1,
|
n = len(preprocessed)
|
||||||
max_length=self.transcriber.max_length,
|
final_results = [None] * n # tuples of (gen_result, avg_logprob, used_temp)
|
||||||
suppress_blank=True,
|
pending_indices = list(range(n))
|
||||||
suppress_tokens=suppress_tokens,
|
|
||||||
return_scores=True,
|
for temp in temperatures:
|
||||||
return_no_speech_prob=True,
|
if not pending_indices:
|
||||||
sampling_temperature=0.0,
|
break
|
||||||
repetition_penalty=1,
|
|
||||||
no_repeat_ngram_size=0,
|
if len(pending_indices) == n:
|
||||||
)
|
sub_encoder = encoder_output
|
||||||
|
else:
|
||||||
|
# Re-encode features for just the pending items to get
|
||||||
|
# an encoder_output of the right batch dimension.
|
||||||
|
sub_feature_batch = np.stack(
|
||||||
|
[preprocessed[i][1] for i in pending_indices]
|
||||||
|
)
|
||||||
|
sub_encoder = self.transcriber.encode(sub_feature_batch)
|
||||||
|
sub_prompts = [prompts[i] for i in pending_indices]
|
||||||
|
|
||||||
|
gen_kwargs = dict(
|
||||||
|
beam_size=5 if temp == 0.0 else 1,
|
||||||
|
patience=1,
|
||||||
|
length_penalty=1,
|
||||||
|
max_length=self.transcriber.max_length,
|
||||||
|
suppress_blank=True,
|
||||||
|
suppress_tokens=suppress_tokens,
|
||||||
|
return_scores=True,
|
||||||
|
return_no_speech_prob=True,
|
||||||
|
sampling_temperature=temp,
|
||||||
|
repetition_penalty=1,
|
||||||
|
no_repeat_ngram_size=0,
|
||||||
|
)
|
||||||
|
batch_results = self.transcriber.model.generate(
|
||||||
|
sub_encoder, sub_prompts, **gen_kwargs
|
||||||
|
)
|
||||||
|
|
||||||
|
next_pending = []
|
||||||
|
for j, idx in enumerate(pending_indices):
|
||||||
|
gen_result = batch_results[j]
|
||||||
|
tokens = gen_result.sequences_ids[0]
|
||||||
|
seq_len = len(tokens)
|
||||||
|
cum_logprob = gen_result.scores[0] * seq_len
|
||||||
|
avg_logprob = cum_logprob / (seq_len + 1) if seq_len > 0 else 0.0
|
||||||
|
raw_text = tokenizers_list[idx].decode(tokens).strip()
|
||||||
|
comp_ratio = get_compression_ratio(raw_text) if raw_text else 0.0
|
||||||
|
|
||||||
|
bad = (
|
||||||
|
comp_ratio > comp_thresh
|
||||||
|
or avg_logprob < logprob_thresh
|
||||||
|
)
|
||||||
|
# High no_speech + low logprob -> treat as silence, accept empty.
|
||||||
|
is_silence = (
|
||||||
|
gen_result.no_speech_prob > no_speech_thresh
|
||||||
|
and avg_logprob < logprob_thresh
|
||||||
|
)
|
||||||
|
|
||||||
|
if not bad or is_silence or temp == temperatures[-1]:
|
||||||
|
final_results[idx] = (gen_result, avg_logprob, temp)
|
||||||
|
else:
|
||||||
|
next_pending.append(idx)
|
||||||
|
|
||||||
|
pending_indices = next_pending
|
||||||
|
|
||||||
# Step 5: Per-item segment parsing and result dispatch
|
# Step 5: Per-item segment parsing and result dispatch
|
||||||
for i, (req, features, audio, duration, speech_chunks) in enumerate(preprocessed):
|
for i, (req, features, audio, duration, speech_chunks) in enumerate(preprocessed):
|
||||||
try:
|
try:
|
||||||
tokenizer = tokenizers_list[i]
|
tokenizer = tokenizers_list[i]
|
||||||
gen_result = results[i]
|
gen_result, avg_logprob, used_temp = final_results[i]
|
||||||
|
|
||||||
tokens = gen_result.sequences_ids[0]
|
tokens = gen_result.sequences_ids[0]
|
||||||
seq_len = len(tokens)
|
|
||||||
cum_logprob = gen_result.scores[0] * seq_len
|
|
||||||
avg_logprob = cum_logprob / (seq_len + 1) if seq_len > 0 else 0.0
|
|
||||||
|
|
||||||
segment_size = int(ceil(duration) * self.transcriber.frames_per_second)
|
segment_size = int(ceil(duration) * self.transcriber.frames_per_second)
|
||||||
|
|
||||||
subsegments, _, _ = self.transcriber._split_segments_by_timestamps(
|
subsegments, _, _ = self.transcriber._split_segments_by_timestamps(
|
||||||
@@ -364,7 +417,7 @@ class BatchInferenceWorker:
|
|||||||
compression_ratio=get_compression_ratio(text),
|
compression_ratio=get_compression_ratio(text),
|
||||||
no_speech_prob=gen_result.no_speech_prob,
|
no_speech_prob=gen_result.no_speech_prob,
|
||||||
words=None,
|
words=None,
|
||||||
temperature=0.0,
|
temperature=used_temp,
|
||||||
))
|
))
|
||||||
|
|
||||||
req.result = segments
|
req.result = segments
|
||||||
|
|||||||
Reference in New Issue
Block a user