File size: 3,071 Bytes
f1b0176 06feb89 f1b0176 06feb89 f1b0176 06feb89 f1b0176 06feb89 f1b0176 06feb89 f1b0176 06feb89 f1b0176 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 | import unittest
from unittest.mock import patch
from training_coach.parser_runtime import (
DEFAULT_TRANSFORMERS_MODEL,
PARSER_MODEL,
MODEL_CACHE_ENV_VAR,
ParserRuntimeUnavailableError,
_load_transformers,
generate_parser_response,
load_parser_model,
parse_check_in_with_model,
)
class ParserRuntimeTest(unittest.TestCase):
def tearDown(self):
load_parser_model.cache_clear()
def test_runtime_exports_model_name(self):
self.assertEqual(PARSER_MODEL, "Qwen/Qwen2.5-1.5B-Instruct")
self.assertEqual(DEFAULT_TRANSFORMERS_MODEL, "Qwen/Qwen3-1.7B")
def test_runtime_uses_default_huggingface_cache_by_default(self):
self.assertEqual(MODEL_CACHE_ENV_VAR, "PARSER_MODEL_CACHE_DIR")
def test_load_parser_model_uses_default_hf_cache_without_override(self):
class FakeTokenizer:
@classmethod
def from_pretrained(cls, model_name, cache_dir=None):
return {"model_name": model_name, "cache_dir": cache_dir}
class FakeModel:
@classmethod
def from_pretrained(cls, model_name, **kwargs):
return {"model_name": model_name, **kwargs}
load_parser_model.cache_clear()
with patch.dict("os.environ", {}, clear=True), patch(
"training_coach.parser_runtime._load_transformers",
return_value=(FakeModel, FakeTokenizer),
):
tokenizer, model = load_parser_model("test/model")
self.assertIsNone(tokenizer["cache_dir"])
self.assertIsNone(model["cache_dir"])
def test_load_parser_model_allows_explicit_cache_override(self):
class FakeTokenizer:
@classmethod
def from_pretrained(cls, model_name, cache_dir=None):
return {"model_name": model_name, "cache_dir": cache_dir}
class FakeModel:
@classmethod
def from_pretrained(cls, model_name, **kwargs):
return {"model_name": model_name, **kwargs}
load_parser_model.cache_clear()
with patch.dict(
"os.environ",
{MODEL_CACHE_ENV_VAR: "/tmp/parser-cache"},
clear=True,
), patch(
"training_coach.parser_runtime._load_transformers",
return_value=(FakeModel, FakeTokenizer),
):
tokenizer, model = load_parser_model("test/model")
self.assertEqual(tokenizer["cache_dir"], "/tmp/parser-cache")
self.assertEqual(model["cache_dir"], "/tmp/parser-cache")
def test_runtime_can_import_transformers_or_reports_clear_error(self):
try:
loaded = _load_transformers()
except ParserRuntimeUnavailableError as error:
self.assertIn("Install transformers", str(error))
else:
self.assertEqual(len(loaded), 2)
def test_runtime_functions_are_callable(self):
self.assertTrue(callable(generate_parser_response))
self.assertTrue(callable(parse_check_in_with_model))
if __name__ == "__main__":
unittest.main()
|