From a6492e4496c76691e6f7801476ac603ae6acba16 Mon Sep 17 00:00:00 2001 From: Mayukh Das Date: Thu, 24 Sep 2026 03:17:51 +0200 Subject: [PATCH] fix: add real transformer batch inference --- lettucedetect/detectors/llm.py | 5 + lettucedetect/detectors/transformer.py | 117 +++++++++++++++++++++- scripts/benchmark_detectors.py | 28 +++++- tests/test_benchmark_detectors_pytest.py | 5 +- tests/test_inference_pytest.py | 120 ++++++++++++++++++++++- tests/test_llm_detector_pytest.py | 7 ++ 6 files changed, 267 insertions(+), 15 deletions(-) diff --git a/lettucedetect/detectors/llm.py b/lettucedetect/detectors/llm.py index 572d09b..1b13e74 100644 --- a/lettucedetect/detectors/llm.py +++ b/lettucedetect/detectors/llm.py @@ -602,6 +602,11 @@ def predict_prompt_batch( :returns: One result per input pair: spans, or per-token dicts when ``output_format="tokens"``. """ + if len(prompts) != len(answers): + raise ValueError( + "prompts and answers must contain the same number of items " + f"(got {len(prompts)} and {len(answers)})" + ) if output_format not in ["tokens", "spans"]: raise ValueError( f"LLMDetector doesn't support '{output_format}' format. Use 'tokens' or 'spans'" diff --git a/lettucedetect/detectors/transformer.py b/lettucedetect/detectors/transformer.py index 600046d..ff42eb3 100644 --- a/lettucedetect/detectors/transformer.py +++ b/lettucedetect/detectors/transformer.py @@ -182,9 +182,31 @@ def _predict_single(self, prompt: str, answer: str, output_format: str) -> list: token_preds = torch.where(labels == -100, labels, token_preds) + return self._format_prediction( + encoding["input_ids"][0], + labels, + offsets, + token_preds, + probabilities, + answer_start_token, + answer, + output_format, + ) + + def _format_prediction( + self, + input_ids: torch.Tensor, + labels: torch.Tensor, + offsets: torch.Tensor, + token_preds: torch.Tensor, + probabilities: torch.Tensor, + answer_start_token: int, + answer: str, + output_format: str, + ) -> list: + """Convert one model-output row to token or span predictions.""" if output_format == "tokens": token_probs: list[dict] = [] - input_ids = encoding["input_ids"][0] for i, (token, pred, prob) in enumerate(zip(input_ids, token_preds, probabilities)): if labels[i].item() != -100: token_probs.append( @@ -441,13 +463,98 @@ def predict_prompt_batch( ) -> list: """Predict hallucination tokens or spans from the provided prompts and answers. + The complete input lists are padded into one batch and evaluated in one + model forward pass. Callers handling very large lists should split them + into appropriately sized batches for the available accelerator memory. + :param prompts: List of prompt strings. :param answers: List of answer strings. :param output_format: ``"tokens"`` or ``"spans"``. :param min_confidence: Drop ``"spans"`` below this confidence threshold (``[0, 1]``). :returns: List of prediction lists, one per input pair. """ - return [ - self.predict_prompt(p, a, output_format, min_confidence) - for p, a in zip(prompts, answers) - ] + if len(prompts) != len(answers): + raise ValueError( + "prompts and answers must contain the same number of items " + f"(got {len(prompts)} and {len(answers)})" + ) + if output_format not in ("tokens", "spans"): + raise ValueError( + f"TransformerDetector doesn't support '{output_format}' format. " + "Use 'tokens' or 'spans'" + ) + self._validate_min_confidence(min_confidence) + if not prompts: + return [] + + untruncated = self.tokenizer( + prompts, + answers, + add_special_tokens=True, + truncation=False, + ) + for index, input_ids in enumerate(untruncated["input_ids"]): + if len(input_ids) > self.max_length: + logger.warning( + "predict_prompt_batch: input %d (%d tokens) exceeds max_length (%d). " + "The prompt will be truncated. Use predict() with structured " + "passages for automatic chunking.", + index, + len(input_ids), + self.max_length, + ) + + encoding = self.tokenizer( + prompts, + answers, + truncation="only_first", + max_length=self.max_length, + padding=True, + return_offsets_mapping=True, + return_tensors="pt", + add_special_tokens=True, + ) + offsets = encoding.pop("offset_mapping") + attention_mask = encoding["attention_mask"] + labels = torch.full_like(encoding["input_ids"], -100) + answer_starts: list[int] = [] + + for row in range(len(prompts)): + sequence_ids = encoding.sequence_ids(row) + try: + answer_start = sequence_ids.index(1) + except ValueError: + answer_start = int(attention_mask[row].sum().item()) - 1 + sequence_length = int(attention_mask[row].sum().item()) + labels[row, answer_start:sequence_length] = 0 + answer_starts.append(answer_start) + + model_encoding = { + key: value.to(self.device) + for key, value in encoding.items() + if key in ["input_ids", "attention_mask"] + } + model_labels = labels.to(self.device) + + with torch.no_grad(): + logits = self.model(**model_encoding).logits + token_preds = torch.argmax(logits, dim=-1) + probabilities = torch.softmax(logits, dim=-1) + token_preds = torch.where(model_labels == -100, model_labels, token_preds) + + results: list[list] = [] + for row, (prompt, answer, answer_start) in enumerate(zip(prompts, answers, answer_starts)): + result = self._format_prediction( + model_encoding["input_ids"][row], + model_labels[row], + offsets[row], + token_preds[row], + probabilities[row], + answer_start, + answer, + output_format, + ) + if output_format == "spans" and self.typer is not None: + result = self.typer.type_spans(answer, prompt, result) + results.append(self._filter_spans_by_confidence(result, output_format, min_confidence)) + return results diff --git a/scripts/benchmark_detectors.py b/scripts/benchmark_detectors.py index c9409fb..d6fd6d4 100644 --- a/scripts/benchmark_detectors.py +++ b/scripts/benchmark_detectors.py @@ -5,11 +5,12 @@ import argparse import json import platform -import resource import statistics import time +import tracemalloc from dataclasses import asdict, dataclass from pathlib import Path +from types import ModuleType from typing import NotRequired, Protocol, TypedDict, cast DEFAULT_QUESTION = "What is the capital of France?" @@ -174,11 +175,22 @@ def import_torch() -> TorchModule | None: return cast(TorchModule, torch) +def import_resource() -> ModuleType | None: + """Import the Unix resource module when the platform provides it.""" + try: + import resource + except ImportError: + return None + return resource + + def reset_peak_memory(device: str) -> None: """Reset GPU peak-memory counters before the measured loop when available.""" torch = import_torch() if torch is not None and device.startswith("cuda") and torch.cuda.is_available(): torch.cuda.reset_peak_memory_stats(device) + elif import_resource() is None: + tracemalloc.start() def peak_memory_bytes(device: str) -> tuple[int, str]: @@ -187,10 +199,16 @@ def peak_memory_bytes(device: str) -> tuple[int, str]: if torch is not None and device.startswith("cuda") and torch.cuda.is_available(): return torch.cuda.max_memory_allocated(device), "torch.cuda.max_memory_allocated" - usage = resource.getrusage(resource.RUSAGE_SELF).ru_maxrss - if platform.system() == "Darwin": - return int(usage), "resource.getrusage(RUSAGE_SELF).ru_maxrss" - return int(usage) * 1024, "resource.getrusage(RUSAGE_SELF).ru_maxrss" + resource = import_resource() + if resource is not None: + usage = resource.getrusage(resource.RUSAGE_SELF).ru_maxrss + if platform.system() == "Darwin": + return int(usage), "resource.getrusage(RUSAGE_SELF).ru_maxrss" + return int(usage) * 1024, "resource.getrusage(RUSAGE_SELF).ru_maxrss" + + _, peak = tracemalloc.get_traced_memory() + tracemalloc.stop() + return peak, "tracemalloc.get_traced_memory" def percentile_95(values: list[float]) -> float: diff --git a/tests/test_benchmark_detectors_pytest.py b/tests/test_benchmark_detectors_pytest.py index 00307e4..e4f04ca 100644 --- a/tests/test_benchmark_detectors_pytest.py +++ b/tests/test_benchmark_detectors_pytest.py @@ -66,7 +66,10 @@ def test_run_benchmark_reports_latency_throughput_device_and_memory(self): assert result.latency_mean_ms >= 0 assert result.throughput_cases_per_second > 0 assert result.peak_memory_bytes > 0 - assert result.peak_memory_source == "resource.getrusage(RUSAGE_SELF).ru_maxrss" + if module.import_resource() is None: + assert result.peak_memory_source == "tracemalloc.get_traced_memory" + else: + assert result.peak_memory_source == "resource.getrusage(RUSAGE_SELF).ru_maxrss" assert [case.name for case in result.case_results] == ["short", "medium"] def test_run_benchmark_resets_and_reports_cuda_peak_memory(self): diff --git a/tests/test_inference_pytest.py b/tests/test_inference_pytest.py index c29a572..020846e 100644 --- a/tests/test_inference_pytest.py +++ b/tests/test_inference_pytest.py @@ -233,6 +233,112 @@ def test_form_prompt_without_question(self): assert "Summarize" in prompt +class TestTransformerPromptBatch: + """Tests for padded transformer batch inference.""" + + @pytest.fixture(autouse=True) + def setup(self, local_wordpiece_tokenizer): + """Build a detector around a deterministic model stub.""" + self.detector = TransformerDetector.__new__(TransformerDetector) + self.detector.tokenizer = local_wordpiece_tokenizer + self.detector.max_length = 32 + self.detector.device = torch.device("cpu") + self.detector.typer = None + + hallucinated_ids = { + local_wordpiece_tokenizer.convert_tokens_to_ids("paris"), + local_wordpiece_tokenizer.convert_tokens_to_ids("answer"), + } + + def forward(input_ids, attention_mask, labels=None): + del attention_mask, labels + logits = torch.empty((*input_ids.shape, 2), dtype=torch.float) + logits[..., 0] = 4.0 + logits[..., 1] = 0.0 + for token_id in hallucinated_ids: + mask = input_ids == token_id + logits[mask, 0] = 0.0 + logits[mask, 1] = 4.0 + output = MagicMock() + output.logits = logits + return output + + self.detector.model = MagicMock(side_effect=forward) + + def test_batch_matches_single_predictions_and_calls_model_once(self): + """Uneven inputs retain order and use one model forward pass.""" + prompts = ["The capital of France is", "short"] + answers = ["paris.", "short answer"] + + expected = [ + self.detector.predict_prompt(prompt, answer, output_format="tokens") + for prompt, answer in zip(prompts, answers) + ] + self.detector.model.reset_mock() + + actual = self.detector.predict_prompt_batch(prompts, answers, output_format="tokens") + + assert actual == expected + self.detector.model.assert_called_once() + assert [item[0]["token"].strip() for item in actual] == ["paris", "short"] + assert all( + item["token"] != "[PAD]" # noqa: S105 - tokenizer special token + for result in actual + for item in result + ) + + def test_single_item_batch_uses_one_forward_call(self): + """A batch of size one still uses the padded batch path.""" + result = self.detector.predict_prompt_batch(["short"], ["paris."], "tokens") + + assert result[0][0]["pred"] == 1 + self.detector.model.assert_called_once() + + def test_spans_confidence_filtering_and_taxonomy_typing(self): + """Span output is filtered after optional taxonomy typing.""" + typer = MagicMock() + typer.type_spans.side_effect = lambda answer, prompt, spans: [ + {**span, "category": f"{prompt}:{answer}"} for span in spans + ] + self.detector.typer = typer + + results = self.detector.predict_prompt_batch( + ["capital", "short"], + ["paris.", "short answer"], + output_format="spans", + min_confidence=0.9, + ) + + assert [[span["text"] for span in result] for result in results] == [ + ["paris"], + ["answer"], + ] + assert results[0][0]["category"] == "capital:paris." + assert results[1][0]["category"] == "short:short answer" + assert typer.type_spans.call_count == 2 + + filtered = self.detector.predict_prompt_batch( + ["capital"], ["paris."], output_format="spans", min_confidence=0.99 + ) + assert filtered == [[]] + + def test_empty_batch_skips_tokenizer_and_model(self): + """An empty input produces an empty output without inference.""" + tokenizer = MagicMock(wraps=self.detector.tokenizer) + self.detector.tokenizer = tokenizer + + assert self.detector.predict_prompt_batch([], []) == [] + tokenizer.assert_not_called() + self.detector.model.assert_not_called() + + def test_mismatched_lengths_raise_before_inference(self): + """Trailing prompts or answers are rejected instead of truncated by zip.""" + with pytest.raises(ValueError, match="same number"): + self.detector.predict_prompt_batch(["one", "two"], ["answer"]) + + self.detector.model.assert_not_called() + + class TestChunking: """Tests for automatic context chunking when input exceeds max_length.""" @@ -517,15 +623,21 @@ def test_invalid_min_confidence_raises(self, method_name, bad_value): *args, output_format="spans", min_confidence=bad_value ) - def test_predict_prompt_batch_respects_min_confidence(self): + def test_predict_prompt_batch_respects_min_confidence(self, local_wordpiece_tokenizer): """The batch path applies the threshold to each item's spans.""" spans = [ {"start": 0, "end": 3, "confidence": 0.40, "text": "foo"}, {"start": 4, "end": 7, "confidence": 0.90, "text": "bar"}, ] - # predict_prompt measures token length first; return a small fixed count. - self.detector.tokenizer.return_value = {"input_ids": torch.zeros(1, 4, dtype=torch.long)} - with patch.object(TransformerDetector, "_predict_single", return_value=spans): + self.detector.tokenizer = local_wordpiece_tokenizer + + def forward(**kwargs): + output = MagicMock() + output.logits = torch.zeros((*kwargs["input_ids"].shape, 2)) + return output + + self.detector.model.side_effect = forward + with patch.object(TransformerDetector, "_format_prediction", return_value=spans): results = self.detector.predict_prompt_batch( ["p1"], ["foo bar"], output_format="spans", min_confidence=0.5 ) diff --git a/tests/test_llm_detector_pytest.py b/tests/test_llm_detector_pytest.py index 1003b86..bfa1c43 100644 --- a/tests/test_llm_detector_pytest.py +++ b/tests/test_llm_detector_pytest.py @@ -89,6 +89,13 @@ def test_predict_prompt_batch_returns_token_lists(self, cache_file): assert any(t["pred"] == 1 for t in results[0]) assert all(t["pred"] == 0 for t in results[1]) + def test_predict_prompt_batch_rejects_mismatched_lengths(self, cache_file): + """Batch prediction does not silently discard an unmatched input.""" + detector = make_detector('{"hallucination_list": []}', cache_file) + + with pytest.raises(ValueError, match="same number"): + detector.predict_prompt_batch(["p1", "p2"], ["a1"]) + def test_supported_tokens_use_low_constant_prob(self, cache_file): """Tokens outside any span get pred=0 and the supported-prob constant.""" answer = "Everything here is fine."