Skip to content
Open
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: 3 additions & 0 deletions lettucedetect/detectors/llm.py
Original file line number Diff line number Diff line change
Expand Up @@ -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'"
Expand Down
215 changes: 170 additions & 45 deletions lettucedetect/detectors/transformer.py
Original file line number Diff line number Diff line change
Expand Up @@ -33,6 +33,7 @@ def __init__(
lang: Lang = "en",
taxonomy_head: str | None = None,
include_taxonomy: bool | list | dict = True,
batch_size: int = 16,
**tok_kwargs: object,
) -> None:
"""Initialize the transformer detector.
Expand All @@ -48,11 +49,16 @@ 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`. 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)
self.model = AutoModelForTokenClassification.from_pretrained(model_path, **tok_kwargs)
self.device = device or (
Expand Down Expand Up @@ -151,49 +157,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"
Expand All @@ -206,9 +206,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
Expand All @@ -221,11 +218,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)
Expand All @@ -241,6 +234,106 @@ 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 one 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 tokenizer."
)

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
# ------------------------------------------------------------------
Expand Down Expand Up @@ -438,16 +531,48 @@ 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.

: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]``).
: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'"
)
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 = self.batch_size if batch_size is None else batch_size

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:
for i in range(len(results)):
results[i] = self.typer.type_spans(answers[i], prompts[i], results[i])

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
]
Loading
Loading