from __future__ import annotations import importlib.util import unittest from pathlib import Path import numpy as np MODEL_PATH = Path(__file__).resolve().parents[1] / "minimax_mlx_model.py" SPEC = importlib.util.spec_from_file_location("pocketai_minimax_mlx_model", MODEL_PATH) assert SPEC is not None and SPEC.loader is not None MODEL = importlib.util.module_from_spec(SPEC) SPEC.loader.exec_module(MODEL) class MiniMaxMlxModelTests(unittest.TestCase): def test_prompt_matches_pinned_comfyui_normalization(self) -> None: prompt = MODEL.build_prompt( "## Dream pop\n*warm* synths <|bpm 92|>", "[Verse] City lights ^ dissolve in rain", ) self.assertEqual( prompt, "<|im_start|><|caption_start|>Dream pop\nwarm synths bpm is 92" "<|caption_end|><|lyrics_start|>[start]\n[verse]\nCity lights\ndissolve in rain" "<|lyrics_end|><|im_end|><|audio_start|>", ) def test_latent_length_and_flow_schedule_match_comfyui(self) -> None: self.assertEqual(MODEL.latent_length(250), 861) self.assertEqual(MODEL.simple_flow_schedule(4), [1.0, 0.75, 0.5, 0.25, 0.0]) self.assertEqual(len(MODEL.simple_flow_schedule(30)), 31) def test_seed_derivation_is_repeatable_and_domain_separated(self) -> None: self.assertEqual(MODEL.derive_seed(17, "ar"), MODEL.derive_seed(17, "ar")) self.assertNotEqual(MODEL.derive_seed(17, "ar"), MODEL.derive_seed(17, "dit")) def test_dav_decode_slices_cover_long_latent_without_unsafe_chunks(self) -> None: frames = 5_167 slices = MODEL.dav_decode_slices(frames) retained = [] for context_start, context_end, local_start, local_end in slices: self.assertLessEqual(context_end - context_start, MODEL.DAV_MAX_CHUNK_FRAMES) self.assertGreaterEqual(local_start, 0) self.assertLessEqual(local_end, context_end - context_start) retained.extend(range(context_start + local_start, context_start + local_end)) self.assertEqual(retained, list(range(frames))) def test_dav_decode_slices_leave_short_latent_unchanged(self) -> None: self.assertEqual(MODEL.dav_decode_slices(861), [(0, 861, 0, 861)]) def test_stereo_collapse_fraction_detects_missing_long_channel(self) -> None: healthy = np.stack((np.full(100, 0.2), np.full(100, 0.15))) collapsed = np.stack((np.full(100, 0.2), np.full(100, 0.001))) self.assertEqual(MODEL.stereo_collapse_fraction(healthy, sample_rate=10), 0.0) self.assertEqual(MODEL.stereo_collapse_fraction(collapsed, sample_rate=10), 1.0) if __name__ == "__main__": unittest.main()