Skip to content
Merged
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
24 changes: 24 additions & 0 deletions mlx_audio/stt/models/whisper/decoding.py
Original file line number Diff line number Diff line change
Expand Up @@ -127,6 +127,10 @@ class DecodingOptions:
beam_size: Optional[int] = None # number of beams in beam search, if t == 0
patience: Optional[float] = None # patience in beam search (arxiv:2204.05424)

# penalize logits of tokens already present in the generated sequence, to
# discourage degenerate repetition on temperature-ladder escalation
repetition_penalty: Optional[float] = None

# "alpha" in Google NMT, or None for length norm, when ranking generations
# to select which to return among the beams or best-of-N samples
length_penalty: Optional[float] = None
Expand Down Expand Up @@ -369,6 +373,22 @@ def apply(self, logits: mx.array, tokens: mx.array) -> mx.array:
return logits + self.mask


class RepetitionPenalty(LogitFilter):
def __init__(self, penalty: float):
self.penalty = penalty

def apply(self, logits: mx.array, tokens: mx.array) -> mx.array:
n_vocab = logits.shape[-1]
seen = np.zeros((tokens.shape[0], n_vocab), dtype=bool)
for i, seq in enumerate(tokens.tolist()):
for token in seq:
if token < n_vocab:
seen[i, token] = True
seen = mx.array(seen)
penalized = mx.where(logits > 0, logits / self.penalty, logits * self.penalty)
return mx.where(seen, penalized, logits)


class ApplyTimestampRules(LogitFilter):
def __init__(
self,
Expand Down Expand Up @@ -493,6 +513,10 @@ def __init__(self, model: "Whisper", options: DecodingOptions):
model.dims.n_vocab,
)
)
if self.options.repetition_penalty is not None:
self.logit_filters.append(
RepetitionPenalty(self.options.repetition_penalty)
)

if not options.without_timestamps:
precision = CHUNK_LENGTH / model.dims.n_audio_ctx # usually 0.02 seconds
Expand Down
Loading