From 9494dcaf4f91c83356c91e768da79dd22680651f Mon Sep 17 00:00:00 2001 From: "Guaya K. Oh" Date: Thu, 17 Sep 2026 22:57:13 +0200 Subject: [PATCH 1/3] perf: true batched inference for TransformerDetector (fixes #23) TransformerDetector.predict_prompt_batch now performs padded batch tokenization and a single model forward pass per configured batch_size, preserving input order and correctly trimming prompt/padding per sample for both tokens and spans. Add strict len(prompts) == len(answers) validation (no silent zip truncation), including in the LLM detector request path. Tests cover uneven sequence lengths, batch_size 1 and >1, output order, tokens, spans, min_confidence filtering, empty input, and mismatch errors; a spy/stub verifies one forward call per transformer batch without downloading a model. --- lettucedetect/detectors/llm.py | 3 + lettucedetect/detectors/transformer.py | 215 +++++++++++---- tests/test_inference_pytest.py | 358 ++++++++++++++++++++++++- tests/test_llm_detector_pytest.py | 39 +++ 4 files changed, 568 insertions(+), 47 deletions(-) diff --git a/lettucedetect/detectors/llm.py b/lettucedetect/detectors/llm.py index 572d09b..ef560b4 100644 --- a/lettucedetect/detectors/llm.py +++ b/lettucedetect/detectors/llm.py @@ -601,7 +601,10 @@ def predict_prompt_batch( on top of the constructor-level ``min_confidence``. Not applied to token output. :returns: One result per input pair: spans, or per-token dicts when ``output_format="tokens"``. + :raises ValueError: If ``len(prompts) != len(answers)``. """ + if len(prompts) != len(answers): + raise ValueError("Number of prompts must match number of 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..80de563 100644 --- a/lettucedetect/detectors/transformer.py +++ b/lettucedetect/detectors/transformer.py @@ -33,6 +33,7 @@ def __init__( lang: Lang = "en", taxonomy_head: str | None = None, include_taxonomy: bool | list | dict = True, + batch_size: int | None = None, **tok_kwargs: object, ) -> None: """Initialize the transformer detector. @@ -48,11 +49,15 @@ def __init__( enter only as text); ``{"categories": {...}, "subcategories": {...}}`` controls both sets; a list of names selects a subset of the trained categories. Only meaningful together with ``taxonomy_head``. + :param batch_size: Default number of (prompt, answer) pairs scored together in one pass + by :meth:`predict_prompt_batch`. ``None`` (default) scores the whole input list as a + single batch. :param tok_kwargs: Additional keyword arguments for the tokenizer. """ if lang not in LANG_TO_PASSAGE: raise ValueError(f"Invalid language. Choose from {', '.join(LANG_TO_PASSAGE)}") self.lang, self.max_length = lang, max_length + self.batch_size = batch_size self.tokenizer = AutoTokenizer.from_pretrained(model_path, **tok_kwargs) self.model = AutoModelForTokenClassification.from_pretrained(model_path, **tok_kwargs) self.device = device or ( @@ -151,49 +156,43 @@ def _group_passages_into_chunks( return groups if groups else [context] # ------------------------------------------------------------------ - # Single-chunk prediction (the original _predict logic) + # Shared decoding helper for single-example and batched prediction # ------------------------------------------------------------------ - def _predict_single(self, prompt: str, answer: str, output_format: str) -> list: - """Run prediction on a single (prompt, answer) pair that fits in ``max_length``. + def _decode_sample( + self, + input_ids: torch.Tensor, + probabilities: torch.Tensor, + offsets: torch.Tensor, + answer_start_token: int, + answer: str, + output_format: str, + ) -> list: + """Decode one sample's model outputs into tokens or spans. - :param prompt: The prompt string. - :param answer: The answer string. + Shared by the single-example and batched prediction paths so both produce + identical output for the same inputs. All tensors must already be trimmed + to this sample's true (unpadded) sequence length. + + :param input_ids: Token ids for the sample, shape ``(seq_len,)``. + :param probabilities: Softmax probabilities, shape ``(seq_len, num_labels)``. + :param offsets: Character offset mapping, shape ``(seq_len, 2)``. + :param answer_start_token: Index of the first answer-side token. + :param answer: The original answer string. :param output_format: ``"tokens"`` or ``"spans"``. """ - encoding, _, offsets, answer_start_token = HallucinationDataset.prepare_tokenized_input( - self.tokenizer, prompt, answer, self.max_length - ) - - labels = torch.full_like(encoding.input_ids[0], -100, device=self.device) - labels[answer_start_token:] = 0 - - encoding = { - key: value.to(self.device) - for key, value in encoding.items() - if key in ["input_ids", "attention_mask", "labels"] - } - - with torch.no_grad(): - outputs = self.model(**encoding) - logits = outputs.logits - token_preds = torch.argmax(logits, dim=-1)[0] - probabilities = torch.softmax(logits, dim=-1)[0] - - token_preds = torch.where(labels == -100, labels, token_preds) + token_preds = torch.argmax(probabilities, dim=-1) 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( - { - "token": self.tokenizer.decode([token]), - "pred": pred.item(), - "prob": prob[1].item(), - } - ) + for i in range(answer_start_token, input_ids.size(0)): + token_probs.append( + { + "token": self.tokenizer.decode([input_ids[i]]), + "pred": token_preds[i].item(), + "prob": probabilities[i, 1].item(), + } + ) return token_probs # output_format == "spans" @@ -206,9 +205,6 @@ def _predict_single(self, prompt: str, answer: str, output_format: str) -> list: current_span: dict | None = None for i in range(answer_start_token, token_preds.size(0)): - if labels[i].item() == -100: - continue - token_start, token_end = offsets[i].tolist() if token_start == token_end: continue @@ -221,11 +217,7 @@ def _predict_single(self, prompt: str, answer: str, output_format: str) -> list: if is_hallucination: if current_span is None: - current_span = { - "start": rel_start, - "end": rel_end, - "confidence": confidence, - } + current_span = {"start": rel_start, "end": rel_end, "confidence": confidence} else: current_span["end"] = rel_end current_span["confidence"] = max(current_span["confidence"], confidence) @@ -241,6 +233,107 @@ def _predict_single(self, prompt: str, answer: str, output_format: str) -> list: return spans + # ------------------------------------------------------------------ + # Single-chunk prediction (the original _predict logic) + # ------------------------------------------------------------------ + + def _predict_single(self, prompt: str, answer: str, output_format: str) -> list: + """Run prediction on a single (prompt, answer) pair that fits in ``max_length``. + + :param prompt: The prompt string. + :param answer: The answer string. + :param output_format: ``"tokens"`` or ``"spans"``. + """ + encoding, _, offsets, answer_start_token = HallucinationDataset.prepare_tokenized_input( + self.tokenizer, prompt, answer, self.max_length + ) + + encoding = { + key: value.to(self.device) + for key, value in encoding.items() + if key in ["input_ids", "attention_mask", "labels"] + } + + with torch.no_grad(): + outputs = self.model(**encoding) + probabilities = torch.softmax(outputs.logits, dim=-1)[0] + input_ids = encoding["input_ids"][0] + + return self._decode_sample( + input_ids, probabilities, offsets, answer_start_token, answer, output_format + ) + + # ------------------------------------------------------------------ + # Batched prediction with one forward pass for many pairs + # ------------------------------------------------------------------ + + def _predict_batch( + self, prompts: list[str], answers: list[str], output_format: str + ) -> list[list]: + """Tokenize ``prompts``/``answers`` as one padded batch and score them in a + single forward pass. + + :param prompts: Prompt strings for this batch (already length-validated + against ``answers`` by the caller). + :param answers: Answer strings for this batch. + :param output_format: ``"tokens"`` or ``"spans"``. + :returns: One prediction list per (prompt, answer) pair, in input order. + """ + if self.tokenizer.padding_side != "right": + raise ValueError( + "TransformerDetector batched inference requires a right-padding" + ) + + batch = 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 = batch.pop("offset_mapping") + seq_lens = batch["attention_mask"].sum(dim=1) + + model_inputs = { + key: value.to(self.device) + for key, value in batch.items() + if key in ["input_ids", "attention_mask"] + } + + with torch.no_grad(): + outputs = self.model(**model_inputs) + probabilities = torch.softmax(outputs.logits, dim=-1) + + results: list[list] = [] + for i, answer in enumerate(answers): + seq_len = int(seq_lens[i].item()) + sequence_ids = batch.sequence_ids(i) + answer_start_token = next( + (idx for idx, seq_id in enumerate(sequence_ids) if seq_id == 1), + seq_len -1, + ) + + if seq_len >= self.max_length: + logger.warning( + f"predict_prompt_batch: item {i} ({seq_len} tokens) reached " + f"max length ({self.max_length}) and may have been truncated." + ) + + results.append( + self._decode_sample( + batch["input_ids"][i, :seq_len], + probabilities[i, :seq_len], + offsets[i, :seq_len], + answer_start_token, + answer, + output_format, + ) + ) + return results + # ------------------------------------------------------------------ # Multi-chunk prediction with max() aggregation # ------------------------------------------------------------------ @@ -438,6 +531,7 @@ def predict_prompt_batch( answers: list[str], output_format: str = "tokens", min_confidence: float = 0.0, + batch_size: int | None = None, ) -> list: """Predict hallucination tokens or spans from the provided prompts and answers. @@ -445,9 +539,40 @@ def predict_prompt_batch( :param answers: List of answer strings. :param output_format: ``"tokens"`` or ``"spans"``. :param min_confidence: Drop ``"spans"`` below this confidence threshold (``[0, 1]``). + :param batch_size: Max number of pairs tokenized and scored together in one forward + pass. Defaults to ``self.batch_size`` or the whole input list if that is also + unset. Every sample in a batch is padded to the longest sequence in that batch, + so memory scales with ``batch_size x longest_sequence_in_batch``. :returns: List of prediction lists, one per input pair. + :raises ValueError: If ``len(prompts) != len(answers)``. """ + if len(prompts) != len(answers): + raise ValueError("Number of prompts must match number of answers") + if not prompts: + return [] + 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) + + effective_batch_size = batch_size or self.batch_size or len(prompts) + + results: list[list] = [] + for start in range(0, len(prompts), effective_batch_size): + end = start + effective_batch_size + results.extend( + self._predict_batch(prompts[start:end], answers[start:end], output_format) + ) + + if output_format == "spans" and self.typer is not None: + results = [ + self.typer.type_spans(answer, prompt, spans) + for prompt, answer, spans in zip(prompts, answers, results) + ] + return [ - self.predict_prompt(p, a, output_format, min_confidence) - for p, a in zip(prompts, answers) + self._filter_spans_by_confidence(result, output_format, min_confidence) + for result in results ] diff --git a/tests/test_inference_pytest.py b/tests/test_inference_pytest.py index c29a572..d804836 100644 --- a/tests/test_inference_pytest.py +++ b/tests/test_inference_pytest.py @@ -524,8 +524,8 @@ def test_predict_prompt_batch_respects_min_confidence(self): {"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.return_value = {"input_ids": torch.zeros(1, 4, dtype=torch.long)} + with patch.object(TransformerDetector, "_predict_batch", return_value=[spans]): results = self.detector.predict_prompt_batch( ["p1"], ["foo bar"], output_format="spans", min_confidence=0.5 ) @@ -565,3 +565,357 @@ def test_predict_prompt_batch_passes_min_confidence(self): mock_detector.predict_prompt_batch.return_value = [] detector.predict_prompt_batch(["p"], ["a"], output_format="spans", min_confidence=0.3) assert mock_detector.predict_prompt_batch.call_args.kwargs["min_confidence"] == 0.3 + + +class TestPredictPromptBatchLengthValidation: + """Fail if predict_prompt_batch gets mismatched input lengths.""" + + @pytest.fixture(autouse=True) + def setup(self): + """Build a TransformerDetector with mocked model/tokenizer.""" + with ( + patch( + "lettucedetect.detectors.transformer.AutoTokenizer.from_pretrained", + return_value=MagicMock(), + ), + patch( + "lettucedetect.detectors.transformer.AutoModelForTokenClassification.from_pretrained", + return_value=MagicMock(), + ) + ): + self.detector = TransformerDetector(model_path="dummy_path") + self.detector.tokenizer.return_value = { + "input_ids": torch.zeros(1, 4, dtype=torch.long) + } + yield + + @pytest.mark.parametrize( + ("prompts", "answers"), + [ + (["p1", "p2"], ["a1"]), + (["p1"], ["a1", "a2"]), + (["p1"], []), + ([], ["a1"]), + ], + ) + def test_mismatched_lengths_raise_value_error(self, prompts, answers): + """Raise ValueError on any length mismatch.""" + with pytest.raises(ValueError, match="Number of prompts must match number of answers"): + self.detector.predict_prompt_batch(prompts, answers) + + def test_mismatch_raises_before_inference(self): + """Validation stops tokenization or forward pass.""" + with patch.object(TransformerDetector, "_predict_single") as spy: + with pytest.raises(ValueError): + self.detector.predict_prompt_batch(["p1", "p2", "p3"], ["a1"]) + spy.assert_not_called() + self.detector.model.assert_not_called() + + +class TestPredictPromptBatchTrueBatching: + """predict_prompt_batch() must do one model forward pass per batch.""" + + @pytest.fixture(autouse=True) + def setup(self, local_wordpiece_tokenizer): + """Real tiny tokenizer plus model that records every forward call.""" + + def fake_forward(input_ids=None, attention_mask=None, **kwargs): + batch, seq_len = input_ids.shape + logits = torch.zeros(batch, seq_len, 2) + logits[..., 0] = 1.0 + output = MagicMock() + output.logits = logits + return output + + self.spy_model = MagicMock(side_effect=fake_forward) + + with ( + patch( + "lettucedetect.detectors.transformer.AutoTokenizer.from_pretrained", + return_value=local_wordpiece_tokenizer, + ), + patch( + "lettucedetect.detectors.transformer.AutoModelForTokenClassification.from_pretrained", + return_value=self.spy_model, + ), + ): + self.detector = TransformerDetector(model_path="dummy_path", max_length=64) + + def test_one_forward_call_for_whole_batch(self): + """Multiple inputs in one batch => exactly one model() call.""" + prompts = ["the capital of france is paris .", "the capital of france is paris ."] + answers = ["paris", "short answer"] + + self.detector.predict_prompt_batch(prompts, answers) + + assert self.spy_model.call_count == 1 + called_input_ids = self.spy_model.call_args.kwargs["input_ids"] + assert called_input_ids.shape[0] == len(prompts) + + def test_batch_size_controls_number_of_forward_calls(self): + """5 inputs with batch_size=2 => 3 forward calls (2 + 2 + 1).""" + prompts = ["the capital of france is paris ."] * 5 + answers = ["paris"] * 5 + + self.detector.predict_prompt_batch(prompts, answers, batch_size=2) + + assert self.spy_model.call_count == 3 + batch_sizes = [ + call.kwargs["input_ids"].shape[0] for call in self.spy_model.call_args_list + ] + assert batch_sizes == [2, 2, 1] + + def test_empty_input_makes_no_forward_call(self): + """No inputs => zero forward calls (and no ValueError).""" + result = self.detector.predict_prompt_batch([], []) + + assert result == [] + self.spy_model.assert_not_called() + + +class TestPredictPromptBatchContentParity: + """Batched output must match per-sample predict_prompt output exactly.""" + + @pytest.fixture(autouse=True) + def setup(self, local_wordpiece_tokenizer): + """Real tiny tokenizer + a modal whose logits depend on token id. + + Flags the "paris" token as a hallucination and everthing else as + supported, so tests can assert on real content. + """ + paris_id = local_wordpiece_tokenizer.convert_tokens_to_ids("paris") + + def fake_forward(input_ids=None, attention_mask=None, **kwargs): + is_paris = (input_ids == paris_id).float() + logits = torch.zeros(*input_ids.shape, 2) + logits[..., 1] = is_paris * 10.0 - (1 - is_paris) * 10 + logits[..., 0] = -logits[..., 1] + output = MagicMock() + output.logits = logits + return output + + self.spy_model = MagicMock(side_effect=fake_forward) + with ( + patch( + "lettucedetect.detectors.transformer.AutoTokenizer.from_pretrained", + return_value=local_wordpiece_tokenizer, + ), + patch( + "lettucedetect.detectors.transformer.AutoModelForTokenClassification.from_pretrained", + return_value=self.spy_model, + ), + ): + self.detector = TransformerDetector(model_path="dummy_path", max_length=64) + + @pytest.mark.parametrize("output_format", ["tokens", "spans"]) + def test_batch_matches_single_example_path(self, output_format): + """Every sample's batched result must equal its predict_prompt() result.""" + prompts = [ + "the capital of france is paris .", + "the capital of france is paris .", + "short answer word", + ] + answers = [ + "paris is the capital", + "the capital of france is paris and paris", + "word", + ] + + expected = [ + self.detector.predict_prompt(p, a, output_format=output_format) + for p, a in zip(prompts, answers) + ] + actual = self.detector.predict_prompt_batch(prompts, answers, output_format=output_format) + + assert actual == expected + + def test_batch_size_one_matches_whole_batch(self): + """batch_size is a pure performance knob and must not change output.""" + prompts = ["the capital of france is paris ."] * 3 + answers = ["paris", "the capital of france", "word"] + + whole = self.detector.predict_prompt_batch(prompts, answers, output_format="spans") + one_at_a_time = self.detector.predict_prompt_batch( + prompts, answers, output_format="spans", batch_size=1 + ) + + assert whole == one_at_a_time + + def test_order_is_preserved_with_uneven_lengths(self): + """Results must align with input order even with very different sequence lengths.""" + prompts = ["the capital of france is paris ."] * 3 + answers = ["word", "paris", "the capital of france is paris and paris again"] + + results = self.detector.predict_prompt_batch(prompts, answers, output_format="tokens") + + assert not any(tok["pred"] == 1 for tok in results[0]) + assert any(tok["pred"] == 1 for tok in results[1]) + assert any(tok["pred"] == 1 for tok in results[2]) + assert len(results[0]) != len(results[2]) + + def test_confidence_filtering_applies_per_sample_in_batch(self): + """min_confidence filters each sample independently within one batch.""" + prompts = ["the capital of france is paris ."] * 2 + answers = ["paris", "the capital of france"] + + filtered = self.detector.predict_prompt_batch( + prompts, answers, output_format="spans", min_confidence=0.99 + ) + unfiltered = self.detector.predict_prompt_batch( + prompts, answers, output_format="spans", min_confidence=0.0 + ) + + assert unfiltered[0] + assert filtered[0] == unfiltered[0] + assert filtered[1] == unfiltered[1] == [] + + +class TestPredictPromptBatchPaddingAndOrder: + """Per-sample padding/prompt-token stripping + order.""" + + @pytest.fixture(autouse=True) + def setup(self, local_wordpiece_tokenizer): + """Model flags [PAD] token ids as 'hallucination' so leaked padding is visible.""" + pad_id = local_wordpiece_tokenizer.pad_token_id + + def fake_forward(input_ids=None, attention_mask=None, **kwargs): + is_pad = (input_ids == pad_id).float() + logits = torch.zeros(*input_ids.shape, 2) + logits[..., 1] = is_pad * 10.0 - (1 - is_pad) * 10.0 + logits[..., 0] = -logits[..., 1] + output = MagicMock() + output.logits = logits + return output + + self.spy_model = MagicMock(side_effect=fake_forward) + with ( + patch( + "lettucedetect.detectors.transformer.AutoTokenizer.from_pretrained", + return_value=local_wordpiece_tokenizer, + ), + patch( + "lettucedetect.detectors.transformer.AutoModelForTokenClassification.from_pretrained", + return_value=self.spy_model, + ), + ): + self.detector = TransformerDetector(model_path="dummy_path", max_length=64) + + def test_padding_and_prompt_tokens_stripped_order_preserved(self): + """A short sample sharing a batch with a long one must not surface [PAD] + tokens/predictions, must return exactly its own answer-token count (no + prompt tokens), and results must align with input order.""" + prompts = [ + "word", + "the capital of france is paris .", + ] + answers = ["word", "the capital of france is paris and word again"] + + results = self.detector.predict_prompt_batch(prompts, answers, output_format="tokens") + + assert len(results) == 2 + for tokens in results: + # No decoded [PAD] tokens, and our fake model only flags real pad + # positions as pred==1, so nothing should be flagged if trimmed correctly. + assert all(tok["token"] != "[PAD]" for tok in tokens) + assert all(tok["pred"] == 0 for tok in tokens) + + # Token count == answer length + trailing [SEP], never the batch's padded max + # length, and never includes any prompt-side tokens. + for i, answer in enumerate(answers): + answer_ids = self.detector.tokenizer(answer, add_special_tokens=False)["input_ids"] + assert len(results[i]) == len(answer_ids) + 1 + + # results[0] (short answer) must stay first, results[1] (long) stays second. + assert len(results[0]) < len(results[1]) + + +class TestPredictPromptBatchTaxonomyTyping: + """predict_prompt_batch() must route 'spans' output through the taxonomy typer, + once per sample, without downloading a taxonomy-head model.""" + + @pytest.fixture(autouse=True) + def setup(self, local_wordpiece_tokenizer): + """Real tiny tokenizer + a model that flags the 'paris' token as hallucination.""" + paris_id = local_wordpiece_tokenizer.convert_tokens_to_ids("paris") + + def fake_forward(input_ids=None, attention_mask=None, **kwargs): + is_paris = (input_ids == paris_id).float() + logits = torch.zeros(*input_ids.shape, 2) + logits[..., 1] = is_paris * 10.0 - (1 - is_paris) * 10.0 + logits[..., 0] = -logits[..., 1] + output = MagicMock() + output.logits = logits + return output + + spy_model = MagicMock(side_effect=fake_forward) + with ( + patch( + "lettucedetect.detectors.transformer.AutoTokenizer.from_pretrained", + return_value=local_wordpiece_tokenizer, + ), + patch( + "lettucedetect.detectors.transformer.AutoModelForTokenClassification.from_pretrained", + return_value=spy_model, + ), + ): + self.detector = TransformerDetector(model_path="dummy_path", max_length=64) + + # Swap in a fake typer so no taxonomy-head model has to be loaded; it tags + # every span with a category derived from its own text. + self.fake_typer = MagicMock() + + def fake_type_spans(answer, prompt, spans): + for span in spans: + span["category"] = f"CAT[{span['text']}]" + return spans + + self.fake_typer.type_spans.side_effect = fake_type_spans + self.detector.typer = self.fake_typer + + def test_typer_called_once_per_sample_with_matching_prompt_and_answer(self): + """type_spans() must be called once per input pair, in input order, with that + sample's own (answer, prompt, spans) -- never another sample's.""" + prompts = [ + "the capital of france is paris .", + "the capital of france is paris .", + ] + answers = ["paris", "the capital of france is paris and paris"] + + self.detector.predict_prompt_batch(prompts, answers, output_format="spans") + + assert self.fake_typer.type_spans.call_count == 2 + for i, call in enumerate(self.fake_typer.type_spans.call_args_list): + called_answer, called_prompt, _ = call.args + assert called_answer == answers[i] + assert called_prompt == prompts[i] + + def test_typed_fields_appear_in_batch_output(self): + """Fields the typer attaches must be present in the final spans result.""" + prompts = ["the capital of france is paris ."] + answers = ["paris is the capital"] + + results = self.detector.predict_prompt_batch(prompts, answers, output_format="spans") + + assert results[0] + assert all("category" in span for span in results[0]) + + def test_typer_not_called_for_token_output(self): + """Typing only applies to 'spans' output, never 'tokens'.""" + self.detector.predict_prompt_batch( + ["the capital of france is paris ."], ["paris"], output_format="tokens" + ) + self.fake_typer.type_spans.assert_not_called() + + def test_typing_runs_before_confidence_filtering(self): + """A span that survives min_confidence filtering must still carry the + category the typer attached (typing must not run after/be skipped by it).""" + results = self.detector.predict_prompt_batch( + ["the capital of france is paris ."], + ["paris"], + output_format="spans", + min_confidence=0.5, + ) + + assert results[0] + assert results[0][0]["category"] == "CAT[paris]" + diff --git a/tests/test_llm_detector_pytest.py b/tests/test_llm_detector_pytest.py index 1003b86..ae9b016 100644 --- a/tests/test_llm_detector_pytest.py +++ b/tests/test_llm_detector_pytest.py @@ -29,6 +29,18 @@ def complete(self, system, user, model, temperature, schema) -> str: return self.response +class CountingClient(FakeClient): + """FakeClient that records how many completions are requested.""" + + def __init__(self, response: str) -> None: + super().__init__(response) + self.calls = 0 + + def complete(self, system, user, model, temperature, schema) -> str: + self.calls += 1 + return super().complete(system, user, model, temperature, schema) + + @pytest.fixture def cache_file(tmp_path): """Temp cache path so the default on-disk cache is never touched.""" @@ -236,3 +248,30 @@ def test_reasoning_response_format_requests_one_item_per_occurrence(self, cache_ assert "items must be listed in answer order" in block assert "return at most one item per distinct occurrence" in block + + +class TestPredictPromptBatchLengthValidationLLM: + """Fail if predict_prompt_batch gets mismatched input lengths.""" + + @pytest.mark.parametrize( + ("prompts", "answers"), + [ + (["p1", "p2"], ["a1"]), + (["p1"], ["a1", "a2"]), + (["p1"], []), + ([], ["a1"]), + ], + ) + def test_mismatched_lengths_raise_value_error(self, prompts, answers, cache_file): + """Raise ValueError on any length mismatch.""" + detector = make_detector('{"hallucination_list": []}', cache_file) + with pytest.raises(ValueError, match="Number of prompts must match number of answers"): + detector.predict_prompt_batch(prompts, answers) + + def test_mismatch_raises_before_inference(self, cache_file): + """Validation stops any request from reaching the client.""" + client = CountingClient('{"hallucination_list": []}') + detector = LLMDetector(client=client, cache_file=cache_file) + with pytest.raises(ValueError): + detector.predict_prompt_batch(["p1", "p2", "p3"], ["a1"]) + assert client.calls == 0 From bdd14abe7985f07b2e17c2c47cffd53ab3282dba Mon Sep 17 00:00:00 2001 From: "Guaya K. Oh" Date: Tue, 29 Sep 2026 13:04:31 +0200 Subject: [PATCH 2/3] chore: ruff/docstring fixes for #23 --- lettucedetect/detectors/transformer.py | 13 +++++-------- tests/test_inference_pytest.py | 22 +++++++--------------- tests/test_llm_detector_pytest.py | 2 ++ 3 files changed, 14 insertions(+), 23 deletions(-) diff --git a/lettucedetect/detectors/transformer.py b/lettucedetect/detectors/transformer.py index 80de563..15a3f76 100644 --- a/lettucedetect/detectors/transformer.py +++ b/lettucedetect/detectors/transformer.py @@ -270,8 +270,7 @@ def _predict_single(self, prompt: str, answer: str, output_format: str) -> list: def _predict_batch( self, prompts: list[str], answers: list[str], output_format: str ) -> list[list]: - """Tokenize ``prompts``/``answers`` as one padded batch and score them in a - single forward pass. + """Tokenize ``prompts``/``answers`` as one padded batch and score them in one pass. :param prompts: Prompt strings for this batch (already length-validated against ``answers`` by the caller). @@ -280,9 +279,7 @@ def _predict_batch( :returns: One prediction list per (prompt, answer) pair, in input order. """ if self.tokenizer.padding_side != "right": - raise ValueError( - "TransformerDetector batched inference requires a right-padding" - ) + raise ValueError("TransformerDetector batched inference requires a right-padding") batch = self.tokenizer( prompts, @@ -313,7 +310,7 @@ def _predict_batch( sequence_ids = batch.sequence_ids(i) answer_start_token = next( (idx for idx, seq_id in enumerate(sequence_ids) if seq_id == 1), - seq_len -1, + seq_len - 1, ) if seq_len >= self.max_length: @@ -333,7 +330,7 @@ def _predict_batch( ) ) return results - + # ------------------------------------------------------------------ # Multi-chunk prediction with max() aggregation # ------------------------------------------------------------------ @@ -571,7 +568,7 @@ def predict_prompt_batch( self.typer.type_spans(answer, prompt, spans) for prompt, answer, spans in zip(prompts, answers, results) ] - + return [ self._filter_spans_by_confidence(result, output_format, min_confidence) for result in results diff --git a/tests/test_inference_pytest.py b/tests/test_inference_pytest.py index d804836..db24b82 100644 --- a/tests/test_inference_pytest.py +++ b/tests/test_inference_pytest.py @@ -581,7 +581,7 @@ def setup(self): patch( "lettucedetect.detectors.transformer.AutoModelForTokenClassification.from_pretrained", return_value=MagicMock(), - ) + ), ): self.detector = TransformerDetector(model_path="dummy_path") self.detector.tokenizer.return_value = { @@ -660,9 +660,7 @@ def test_batch_size_controls_number_of_forward_calls(self): self.detector.predict_prompt_batch(prompts, answers, batch_size=2) assert self.spy_model.call_count == 3 - batch_sizes = [ - call.kwargs["input_ids"].shape[0] for call in self.spy_model.call_args_list - ] + batch_sizes = [call.kwargs["input_ids"].shape[0] for call in self.spy_model.call_args_list] assert batch_sizes == [2, 2, 1] def test_empty_input_makes_no_forward_call(self): @@ -801,9 +799,7 @@ def fake_forward(input_ids=None, attention_mask=None, **kwargs): self.detector = TransformerDetector(model_path="dummy_path", max_length=64) def test_padding_and_prompt_tokens_stripped_order_preserved(self): - """A short sample sharing a batch with a long one must not surface [PAD] - tokens/predictions, must return exactly its own answer-token count (no - prompt tokens), and results must align with input order.""" + """Padding/prompt tokens are stripped per sample and order is preserved.""" prompts = [ "word", "the capital of france is paris .", @@ -816,7 +812,7 @@ def test_padding_and_prompt_tokens_stripped_order_preserved(self): for tokens in results: # No decoded [PAD] tokens, and our fake model only flags real pad # positions as pred==1, so nothing should be flagged if trimmed correctly. - assert all(tok["token"] != "[PAD]" for tok in tokens) + assert all(tok["token"] != "[PAD]" for tok in tokens) # noqa: S105 assert all(tok["pred"] == 0 for tok in tokens) # Token count == answer length + trailing [SEP], never the batch's padded max @@ -830,8 +826,7 @@ def test_padding_and_prompt_tokens_stripped_order_preserved(self): class TestPredictPromptBatchTaxonomyTyping: - """predict_prompt_batch() must route 'spans' output through the taxonomy typer, - once per sample, without downloading a taxonomy-head model.""" + """predict_prompt_batch() routes spans output through the taxonomy typer.""" @pytest.fixture(autouse=True) def setup(self, local_wordpiece_tokenizer): @@ -873,8 +868,7 @@ def fake_type_spans(answer, prompt, spans): self.detector.typer = self.fake_typer def test_typer_called_once_per_sample_with_matching_prompt_and_answer(self): - """type_spans() must be called once per input pair, in input order, with that - sample's own (answer, prompt, spans) -- never another sample's.""" + """type_spans() is called once per sample, with that sample's prompt/answer.""" prompts = [ "the capital of france is paris .", "the capital of france is paris .", @@ -907,8 +901,7 @@ def test_typer_not_called_for_token_output(self): self.fake_typer.type_spans.assert_not_called() def test_typing_runs_before_confidence_filtering(self): - """A span that survives min_confidence filtering must still carry the - category the typer attached (typing must not run after/be skipped by it).""" + """Typing runs before min_confidence filtering, surviving spans keep attached fields.""" results = self.detector.predict_prompt_batch( ["the capital of france is paris ."], ["paris"], @@ -918,4 +911,3 @@ def test_typing_runs_before_confidence_filtering(self): assert results[0] assert results[0][0]["category"] == "CAT[paris]" - diff --git a/tests/test_llm_detector_pytest.py b/tests/test_llm_detector_pytest.py index ae9b016..e031fa1 100644 --- a/tests/test_llm_detector_pytest.py +++ b/tests/test_llm_detector_pytest.py @@ -33,10 +33,12 @@ class CountingClient(FakeClient): """FakeClient that records how many completions are requested.""" def __init__(self, response: str) -> None: + """Initialize the client and reset the call counter.""" super().__init__(response) self.calls = 0 def complete(self, system, user, model, temperature, schema) -> str: + """Increment the call counter and return the canned response.""" self.calls += 1 return super().complete(system, user, model, temperature, schema) From 3d2a38071ca0182125b1f8b6315ec1c2df9c0b56 Mon Sep 17 00:00:00 2001 From: "Guaya K. Oh" Date: Tue, 29 Sep 2026 13:41:47 +0200 Subject: [PATCH 3/3] fix: add batch_size default and checks for #23 --- lettucedetect/detectors/transformer.py | 21 ++++++++++++--------- tests/test_inference_pytest.py | 8 +++++++- 2 files changed, 19 insertions(+), 10 deletions(-) diff --git a/lettucedetect/detectors/transformer.py b/lettucedetect/detectors/transformer.py index 15a3f76..c488ddf 100644 --- a/lettucedetect/detectors/transformer.py +++ b/lettucedetect/detectors/transformer.py @@ -33,7 +33,7 @@ def __init__( lang: Lang = "en", taxonomy_head: str | None = None, include_taxonomy: bool | list | dict = True, - batch_size: int | None = None, + batch_size: int = 16, **tok_kwargs: object, ) -> None: """Initialize the transformer detector. @@ -50,12 +50,13 @@ def __init__( sets; a list of names selects a subset of the trained categories. Only meaningful together with ``taxonomy_head``. :param batch_size: Default number of (prompt, answer) pairs scored together in one pass - by :meth:`predict_prompt_batch`. ``None`` (default) scores the whole input list as a - single batch. + by :meth:`predict_prompt_batch`. Must be >= 1. :param tok_kwargs: Additional keyword arguments for the tokenizer. """ if lang not in LANG_TO_PASSAGE: raise ValueError(f"Invalid language. Choose from {', '.join(LANG_TO_PASSAGE)}") + if batch_size < 1: + raise ValueError("batch_size must be >= 1") self.lang, self.max_length = lang, max_length self.batch_size = batch_size self.tokenizer = AutoTokenizer.from_pretrained(model_path, **tok_kwargs) @@ -279,7 +280,9 @@ def _predict_batch( :returns: One prediction list per (prompt, answer) pair, in input order. """ if self.tokenizer.padding_side != "right": - raise ValueError("TransformerDetector batched inference requires a right-padding") + raise ValueError( + "TransformerDetector batched inference requires a right-padding tokenizer." + ) batch = self.tokenizer( prompts, @@ -552,9 +555,11 @@ def predict_prompt_batch( f"TransformerDetector doesn't support '{output_format}' format." " Use 'tokens' or 'spans'" ) + if batch_size is not None and batch_size < 1: + raise ValueError("batch_size must be >= 1") self._validate_min_confidence(min_confidence) - effective_batch_size = batch_size or self.batch_size or len(prompts) + effective_batch_size = self.batch_size if batch_size is None else batch_size results: list[list] = [] for start in range(0, len(prompts), effective_batch_size): @@ -564,10 +569,8 @@ def predict_prompt_batch( ) if output_format == "spans" and self.typer is not None: - results = [ - self.typer.type_spans(answer, prompt, spans) - for prompt, answer, spans in zip(prompts, answers, results) - ] + for i in range(len(results)): + results[i] = self.typer.type_spans(answers[i], prompts[i], results[i]) return [ self._filter_spans_by_confidence(result, output_format, min_confidence) diff --git a/tests/test_inference_pytest.py b/tests/test_inference_pytest.py index db24b82..f16eb6e 100644 --- a/tests/test_inference_pytest.py +++ b/tests/test_inference_pytest.py @@ -670,6 +670,12 @@ def test_empty_input_makes_no_forward_call(self): assert result == [] self.spy_model.assert_not_called() + @pytest.mark.parametrize("bad_batch_size", [0, -1]) + def test_invalid_batch_size_raises(self, bad_batch_size): + """predict_prompt_batch() rejects invalid batch_size values (< 1).""" + with pytest.raises(ValueError, match="batch_size"): + self.detector.predict_prompt_batch(["p"], ["a"], batch_size=bad_batch_size) + class TestPredictPromptBatchContentParity: """Batched output must match per-sample predict_prompt output exactly.""" @@ -812,7 +818,7 @@ def test_padding_and_prompt_tokens_stripped_order_preserved(self): for tokens in results: # No decoded [PAD] tokens, and our fake model only flags real pad # positions as pred==1, so nothing should be flagged if trimmed correctly. - assert all(tok["token"] != "[PAD]" for tok in tokens) # noqa: S105 + assert all(tok["token"] != "[PAD]" for tok in tokens) # noqa: S105 assert all(tok["pred"] == 0 for tok in tokens) # Token count == answer length + trailing [SEP], never the batch's padded max