Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 2 additions & 1 deletion mlx_lm/evaluate.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)

Expand All @@ -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())
Expand Down
34 changes: 34 additions & 0 deletions tests/test_evaluate.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
):
Expand Down
Loading