diff --git a/tests/test_bugfix_encoder_window_optional.py b/tests/test_bugfix_encoder_window_optional.py new file mode 100644 index 0000000..ce4744a --- /dev/null +++ b/tests/test_bugfix_encoder_window_optional.py @@ -0,0 +1,69 @@ +import importlib.util +import os +import unittest + + +class EncoderWindowBugfixOptionalTests(unittest.TestCase): + @classmethod + def setUpClass(cls): + enabled = os.getenv("VOXMLX_ENABLE_MLX_RUNTIME_TESTS", "").strip().lower() in { + "1", + "true", + "yes", + "on", + } + if not enabled: + raise unittest.SkipTest( + "Set VOXMLX_ENABLE_MLX_RUNTIME_TESTS=1 to run MLX runtime optional tests" + ) + if importlib.util.find_spec("mlx") is None: + raise unittest.SkipTest("mlx is required for runtime optional tests") + + def test_encode_step_uses_encoder_sliding_window(self): + import mlx.core as mx + + from voxmlx.model import VoxtralRealtime + + sliding_window = 7 + config = { + "dim": 16, + "ada_rms_norm_t_cond_dim": 16, + "n_layers": 1, + "n_heads": 2, + "n_kv_heads": 1, + "head_dim": 8, + "hidden_dim": 32, + "vocab_size": 256, + "rope_theta": 1e6, + "multimodal": { + "whisper_model_args": { + "encoder_args": { + "audio_encoding_args": {"num_mel_bins": 128}, + "dim": 16, + "n_layers": 2, + "n_heads": 2, + "head_dim": 8, + "hidden_dim": 32, + "rope_theta": 1e6, + "sliding_window": sliding_window, + }, + "downsample_args": {"downsample_factor": 4}, + } + }, + } + model = VoxtralRealtime(config) + mel_chunk = mx.zeros((128, 16), dtype=mx.float32) + _, _, _, encoder_cache, _ = model.encode_step( + mel_chunk, + conv1_tail=None, + conv2_tail=None, + encoder_cache=None, + ds_buf=None, + ) + self.assertEqual(len(encoder_cache), 2) + for layer_cache in encoder_cache: + self.assertEqual(layer_cache.max_size, sliding_window) + + +if __name__ == "__main__": + unittest.main() diff --git a/voxmlx/model.py b/voxmlx/model.py index e2fd37e..1c712ab 100644 --- a/voxmlx/model.py +++ b/voxmlx/model.py @@ -118,8 +118,9 @@ def encode_step(self, new_mel, conv1_tail, conv2_tail, encoder_cache, ds_buf): # Create encoder cache on first call if encoder_cache is None: + window = int(self.encoder.sliding_window) encoder_cache = [ - RotatingKVCache(100_000) + RotatingKVCache(window) for _ in range(len(self.encoder.layers)) ]