Skip to content
Draft
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
5 changes: 5 additions & 0 deletions lettucedetect/detectors/llm.py
Original file line number Diff line number Diff line change
Expand Up @@ -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'"
Expand Down
117 changes: 112 additions & 5 deletions lettucedetect/detectors/transformer.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down Expand Up @@ -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
28 changes: 23 additions & 5 deletions scripts/benchmark_detectors.py
Original file line number Diff line number Diff line change
Expand Up @@ -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?"
Expand Down Expand Up @@ -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]:
Expand All @@ -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:
Expand Down
5 changes: 4 additions & 1 deletion tests/test_benchmark_detectors_pytest.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down
120 changes: 116 additions & 4 deletions tests/test_inference_pytest.py
Original file line number Diff line number Diff line change
Expand Up @@ -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."""

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