Spaces:
Paused
Paused
File size: 1,130 Bytes
dd8da9e | 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 | import pytest
from unittest.mock import MagicMock, patch
from transformers import AutoModelForCausalLM, AutoTokenizer
@patch("model.AutoTokenizer")
@patch("model.PydanticOutputParser")
def test_generate_answer_success(self, MockParser, MockTokenizer):
# 1. Setup Mock Tokenizer
mock_tokenizer = MockTokenizer.from_pretrained.return_value
mock_tokenizer.return_value = {"input_ids": torch.tensor([[1, 2, 3]])}
mock_tokenizer.decode.return_value = '{"answer": "Pikachu is electric type", "source": "Pokedex"}'
# 2. Setup Mock Model
mock_model = MagicMock()
mock_model.generate.return_value = torch.tensor([[1, 2, 3, 4, 5]])
# 3. Setup Mock Parser (to return a dummy Pydantic object)
mock_parser = MockParser.return_value
expected_response = MagicMock()
mock_parser.parse.return_value = expected_response
# 4. Execute
prompt = "Tell me about Pikachu"
result = generate_answer(mock_model, mock_tokenizer, prompt)
# 5. Assertions
self.assertEqual(result, expected_response)
mock_model.generate.assert_called_once()
mock_tokenizer.decode.assert_called_once()
|