| import unittest |
|
|
| from muse_core import ( |
| META_SAMPLING, |
| NATIVE_GREEDY, |
| choose_seed, |
| coerce_parsed_reply, |
| estimate_gpu_duration, |
| friendly_error, |
| generation_kwargs, |
| preset_values, |
| render_reply, |
| validate_controls, |
| ) |
|
|
|
|
| class PresetTests(unittest.TestCase): |
| def test_native_preset_is_greedy_with_documented_values_ready(self): |
| self.assertEqual(preset_values(NATIVE_GREEDY), (False, 1.0, 0.95, 64)) |
|
|
| def test_meta_preset_enables_sampling(self): |
| self.assertEqual(preset_values(META_SAMPLING), (True, 1.0, 0.95, 64)) |
|
|
| def test_greedy_kwargs_omit_inert_sampling_parameters(self): |
| kwargs = generation_kwargs( |
| do_sample=False, |
| max_new_tokens=512, |
| temperature=1.0, |
| top_p=0.95, |
| top_k=64, |
| repetition_penalty=1.0, |
| ) |
| self.assertNotIn("temperature", kwargs) |
| self.assertNotIn("top_p", kwargs) |
| self.assertNotIn("top_k", kwargs) |
| self.assertEqual(kwargs["eos_token_id"], [200001, 200008]) |
| self.assertEqual(kwargs["pad_token_id"], 200018) |
|
|
| def test_sampling_kwargs_match_model_card(self): |
| kwargs = generation_kwargs( |
| do_sample=True, |
| max_new_tokens=512, |
| temperature=1.0, |
| top_p=0.95, |
| top_k=64, |
| repetition_penalty=1.0, |
| ) |
| self.assertEqual(kwargs["temperature"], 1.0) |
| self.assertEqual(kwargs["top_p"], 0.95) |
| self.assertEqual(kwargs["top_k"], 64) |
|
|
|
|
| class ValidationTests(unittest.TestCase): |
| def test_documented_controls_are_valid(self): |
| validate_controls( |
| max_new_tokens=512, |
| temperature=1.0, |
| top_p=0.95, |
| top_k=64, |
| repetition_penalty=1.0, |
| reasoning_strength="high", |
| ) |
|
|
| def test_invalid_reasoning_is_rejected(self): |
| with self.assertRaises(ValueError): |
| validate_controls( |
| max_new_tokens=512, |
| temperature=1.0, |
| top_p=0.95, |
| top_k=64, |
| repetition_penalty=1.0, |
| reasoning_strength="extreme", |
| ) |
|
|
| def test_seed_resolution(self): |
| self.assertEqual(choose_seed(42, False), 42) |
| self.assertEqual(choose_seed(42, True, randbelow=lambda _limit: 123), 123) |
|
|
| def test_duration_is_bounded_and_image_aware(self): |
| text_duration = estimate_gpu_duration(512, False) |
| image_duration = estimate_gpu_duration(512, True) |
| self.assertGreater(image_duration, text_duration) |
| self.assertGreaterEqual(estimate_gpu_duration(1, False), 60) |
| self.assertLessEqual(estimate_gpu_duration(100_000, True), 240) |
|
|
|
|
| class ResponseTests(unittest.TestCase): |
| def test_parser_fields_are_normalized(self): |
| reply = coerce_parsed_reply( |
| {"reasoning_content": " think ", "content": " answer ", "tool_calls": [{"x": 1}]} |
| ) |
| self.assertEqual(reply.reasoning, "think") |
| self.assertEqual(reply.content, "answer") |
| self.assertEqual(reply.tool_calls, [{"x": 1}]) |
|
|
| def test_reasoning_renders_in_collapsible_region(self): |
| rendered = render_reply("steps", "answer", show_reasoning=True) |
| self.assertIn("<think>\nsteps\n</think>", rendered) |
| self.assertTrue(rendered.endswith("answer")) |
|
|
| def test_reasoning_can_be_hidden_without_hiding_answer(self): |
| rendered = render_reply("secret chain", "answer", show_reasoning=False) |
| self.assertNotIn("secret chain", rendered) |
| self.assertEqual(rendered, "answer") |
|
|
| def test_model_tags_cannot_break_reasoning_wrapper(self): |
| rendered = render_reply("a </think> b", "ok", show_reasoning=True) |
| self.assertIn("</think>", rendered) |
| self.assertEqual(rendered.count("</think>"), 1) |
|
|
| def test_errors_do_not_echo_arbitrary_details(self): |
| message = friendly_error(RuntimeError("secret prompt at /private/path")) |
| self.assertNotIn("secret", message) |
| self.assertNotIn("/private/path", message) |
|
|
|
|
| if __name__ == "__main__": |
| unittest.main() |
|
|
|
|