diff --git a/mlx_lm/evaluate.py b/mlx_lm/evaluate.py index 04616c047..ff1904eaf 100644 --- a/mlx_lm/evaluate.py +++ b/mlx_lm/evaluate.py @@ -195,6 +195,7 @@ def loglikelihood(self, requests) -> list[tuple[float, bool]]: scores, is_greedy = [], [] for q, rs in tqdm(zip(questions, responses), total=len(questions)): prefix = self._tokenize([q])[0] + completion_start = len(prefix) full_sequences = self._tokenize([q + r for r in rs]) max_completed_l = max(len(s) for s in full_sequences) @@ -216,7 +217,7 @@ def loglikelihood(self, requests) -> list[tuple[float, bool]]: max_idx = mx.argmax(logprobs).item() for s in full_sequences: - inputs = s[len(prefix) :] + inputs = s[completion_start:] # The logprobs from the last token of the prompt are # for the first input token scores.append(logprobs[0, inputs[0]].item()) diff --git a/tests/test_evaluate.py b/tests/test_evaluate.py index 402bacbd1..4e270e8ff 100644 --- a/tests/test_evaluate.py +++ b/tests/test_evaluate.py @@ -104,6 +104,40 @@ def test_generate_strip_until_then_strip_thinking(self, mock_batch_generate): self.assertEqual(result, ["answer"]) + def test_loglikelihood_scores_only_continuation_after_context_truncation(self): + def mock_score_fn(inputs, cache=None): + targets = inputs[:, 1:] + return ( + -targets.astype(mx.float32), + None, + mx.ones(targets.shape, dtype=mx.bool_), + ) + + request = MagicMock(args=("context", " continuation")) + + for max_tokens, prefix in ((4, [1, 2, 3]), (3, [2, 3]), (2, [3])): + with self.subTest(max_tokens=max_tokens): + self.mlx_lm._max_tokens = max_tokens + self.mlx_lm._tokenize = MagicMock( + side_effect=[ + [[1, 2, 3]], # Context tokens. + [[1, 2, 3, 4, 5]], # Context + continuation tokens. + ] + ) + self.mlx_lm._process_prompt = MagicMock( + return_value=(-mx.arange(6, dtype=mx.float32)[None, :], []) + ) + self.mlx_lm._score_fn = MagicMock(side_effect=mock_score_fn) + + result = self.mlx_lm.loglikelihood([request]) + + self.mlx_lm._process_prompt.assert_called_once_with(prefix) + self.mlx_lm._score_fn.assert_called_once() + inputs = self.mlx_lm._score_fn.call_args[0][0] + self.assertEqual(inputs.tolist(), [[4, 5]]) + # The total also checks the first target scored from the prompt. + self.assertEqual(result, [(-9.0, False)]) + def test_loglikelihood_returns_negative_infinity_when_context_is_fully_truncated( self, ):