Skip to content
Open
Show file tree
Hide file tree
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
8 changes: 6 additions & 2 deletions mlx_audio/stt/models/whisper/whisper.py
Original file line number Diff line number Diff line change
Expand Up @@ -183,10 +183,14 @@ def non_speech_tokens(self) -> Tuple[int, ...]:

miscellaneous = set("♩♪♫♬♭♮♯")

result = {self.encode(" -")[0], self.encode(" '")[0]}
result = set()
for seed in (" -", " '"):
tokens = self.encode(seed)
if tokens:
result.add(tokens[0])
for symbol in symbols + list(miscellaneous):
for tokens in [self.encode(symbol), self.encode(" " + symbol)]:
if len(tokens) == 1 or symbol in miscellaneous:
if tokens and (len(tokens) == 1 or symbol in miscellaneous):
result.add(tokens[0])

return tuple(sorted(result))
Expand Down
29 changes: 29 additions & 0 deletions tests/test_whisper_decode_options.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@
from mlx_audio.stt.models.whisper.decoding import DecodingOptions
from mlx_audio.stt.models.whisper.whisper import (
_DECODING_OPTION_NAMES,
HFTokenizerWrapper,
_filter_decode_options,
)

Expand Down Expand Up @@ -54,5 +55,33 @@ def test_input_is_not_mutated(self):
self.assertEqual(options, {"language": "en", "frame_threshold": 25})


class TestHFTokenizerWrapper(unittest.TestCase):
def test_non_speech_tokens_skips_empty_seed_encodings(self):
class EmptySeedTokenizer:
def encode(self, text, add_special_tokens=False):
if text in {" -", " '"}:
return []
return [101]

tokenizer = HFTokenizerWrapper(EmptySeedTokenizer(), multilingual=False)

self.assertEqual(tokenizer.non_speech_tokens, (101,))

def test_non_speech_tokens_skips_empty_miscellaneous_encodings(self):
miscellaneous = set("♩♪♫♬♭♮♯")

class EmptyMiscellaneousTokenizer:
def encode(self, text, add_special_tokens=False):
if text.strip() in miscellaneous:
return []
return [101]

tokenizer = HFTokenizerWrapper(
EmptyMiscellaneousTokenizer(), multilingual=False
)

self.assertEqual(tokenizer.non_speech_tokens, (101,))


if __name__ == "__main__":
unittest.main()
Loading