Files
Awni Hannun 55bb9471b8 Batch generation (#443)
* initial batch generation

* more in batch generate

* concatenation

* use batch API in eval

* unique max tokens per prompt

* basic continuous batching

* simplify

* better perf by ensuring everything in same stream

* use data class for response

* check cache type
2025-09-15 16:02:45 -07:00

60 lines
2.2 KiB
Python

# Copyright © 2024 Apple Inc.
import unittest
from unittest.mock import MagicMock, patch
import mlx.core as mx
from mlx_lm.evaluate import MLXLM
class TestMLXLM(unittest.TestCase):
def setUp(self):
# Mock the load function to avoid loading actual models
self.mock_model = MagicMock()
self.mock_tokenizer = MagicMock()
self.mock_tokenizer.model_max_length = 2048
self.mock_tokenizer.chat_template = None
self.mock_tokenizer.encode = MagicMock(return_value=[1, 2, 3, 4])
with patch("mlx_lm.evaluate.load") as mock_load:
mock_load.return_value = (self.mock_model, self.mock_tokenizer)
self.mlx_lm = MLXLM("test_model", max_tokens=128)
def test_loglikelihood_rolling_processes_all_inputs(self):
"""Test that loglikelihood_rolling processes all inputs correctly when batching."""
# Create 5 mock requests to test batching with batch_size=2
mock_requests = [MagicMock(args=(f"text {i}",)) for i in range(5)]
# Mock inputs
test_inputs = [(i, i + 1, i + 2) for i in range(5)]
self.mlx_lm._tokenize = MagicMock(return_value=test_inputs)
# Mock _score_fn to return different scores for each batch
def mock_score_fn(batch):
batch_size = len(batch)
scores = mx.array([[0.1] * 3 for _ in range(batch_size)])
lengths = mx.array([3] * batch_size)
return scores, lengths, None
self.mlx_lm._score_fn = MagicMock(side_effect=mock_score_fn)
self.mlx_lm._batch_size = 2
result = self.mlx_lm.loglikelihood_rolling(mock_requests)
# Should return 5 results (one per request)
self.assertEqual(len(result), 5)
# Should have called _score_fn 3 times (batches of 2, 2, 1)
self.assertEqual(self.mlx_lm._score_fn.call_count, 3)
# Verify the batches were correct sizes
call_args_list = self.mlx_lm._score_fn.call_args_list
self.assertEqual(len(call_args_list[0][0][0]), 2) # First batch: 2 items
self.assertEqual(len(call_args_list[1][0][0]), 2) # Second batch: 2 items
self.assertEqual(len(call_args_list[2][0][0]), 1) # Third batch: 1 item
if __name__ == "__main__":
unittest.main()