muse-glimmer-30b / tests /test_muse_core.py
ssdataanalysis's picture
Replace api_name=False with explicit private endpoints to avoid FnIndex errors
4d0d04c verified
Raw
History Blame Contribute Delete
4.14 kB
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("&lt;/think&gt;", 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()