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
88 changes: 88 additions & 0 deletions scripts/evaluate_lfm_span_model.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,88 @@
#!/usr/bin/env python3
"""Per-source evaluation for the LFM2.5 span tagger.

``LfmForSpanTagging`` is a custom module (bidirectional LFM backbone + linear
head), so it cannot be loaded through ``HallucinationDetector``'s
``AutoModelForTokenClassification`` path. This script loads it directly and
reuses ``TransformerDetector``'s prediction/span-decoding plus the shared
char-overlap metrics, so numbers are directly comparable to
``evaluate_span_model.py`` outputs.

Usage:
python scripts/evaluate_lfm_span_model.py \
--model-path /mnt/workspace/users/adamko/lfm25_encoder_binary \
--dataset KRLabsOrg/lettucedetect-code-hallucination \
--dataset KRLabsOrg/lettucedetect-prose-hallucination \
--split test [--by dataset]
"""

from __future__ import annotations

import argparse
import sys
from pathlib import Path

import torch
from safetensors.torch import load_file
from transformers import AutoTokenizer

SCRIPTS = Path(__file__).resolve().parent
sys.path.insert(0, str(SCRIPTS))
sys.path.insert(0, str(SCRIPTS.parent))

from evaluate_span_model import load_samples # noqa: E402
from span_eval_metrics import print_metrics_table # noqa: E402
from train_lfm_span_detector import LfmForSpanTagging # noqa: E402

from lettucedetect.detectors.transformer import TransformerDetector # noqa: E402


def build_detector(
model_path: str, backbone: str, max_length: int, device: str
) -> TransformerDetector:
"""Duck-type a TransformerDetector around the custom LFM tagger."""
model = LfmForSpanTagging(backbone)
state = load_file(Path(model_path) / "model.safetensors")
model.load_state_dict(state)
model.to(device).eval()

det = TransformerDetector.__new__(TransformerDetector)
det.model = model
det.tokenizer = AutoTokenizer.from_pretrained(model_path)
det.device = torch.device(device)
det.max_length = max_length
det.lang = "en"
det.typer = None
return det


def main() -> None:
"""CLI entry point."""
ap = argparse.ArgumentParser(description="Per-source LFM span-tagger evaluation.")
ap.add_argument("--model-path", required=True)
ap.add_argument("--backbone", default="LiquidAI/LFM2.5-Encoder-350M")
ap.add_argument("--dataset", action="append", default=[], required=True)
ap.add_argument("--split", default="test")
ap.add_argument("--by", choices=["dataset", "language"], default="dataset")
ap.add_argument("--only", default="", help="Keep only rows whose `dataset` field == this.")
ap.add_argument("--limit", type=int, default=0)
ap.add_argument("--max-length", type=int, default=8192)
ap.add_argument("--device", default="cuda" if torch.cuda.is_available() else "cpu")
args = ap.parse_args()

detector = build_detector(args.model_path, args.backbone, args.max_length, args.device)
samples = load_samples(args.dataset, args.split, args.limit, args.only)

from tqdm import tqdm

rows = []
with torch.no_grad():
for s in tqdm(samples, desc="predict"):
pred = detector.predict_prompt(s.prompt, s.answer, output_format="spans")
rows.append((getattr(s, args.by), s.labels, pred))

print_metrics_table(rows)


if __name__ == "__main__":
main()
20 changes: 16 additions & 4 deletions scripts/train_lfm_span_detector.py
100644 → 100755
Original file line number Diff line number Diff line change
Expand Up @@ -42,9 +42,20 @@ class LfmForSpanTagging(nn.Module):
def __init__(self, model_name: str, num_labels: int = 2, dropout: float = 0.1) -> None:
"""Load the remote bidirectional backbone and attach a tagger head."""
super().__init__()
self.backbone = AutoModel.from_pretrained(
model_name, trust_remote_code=True, torch_dtype=torch.float32
)
# Some LFM2.5 encoder checkpoints (e.g. LFM2.5-Encoder-350M) store weights
# under the MaskedLM wrapper prefix (lfm2.*); loading the bare AutoModel
# there silently random-initializes everything. Load the wrapper and take
# its backbone, falling back to AutoModel for bare-backbone checkpoints.
try:
from transformers import AutoModelForMaskedLM
mlm = AutoModelForMaskedLM.from_pretrained(
model_name, trust_remote_code=True, torch_dtype=torch.float32
)
self.backbone = getattr(mlm, "lfm2")
except (ValueError, AttributeError, OSError):
self.backbone = AutoModel.from_pretrained(
model_name, trust_remote_code=True, torch_dtype=torch.float32
)
self.config = self.backbone.config
hidden = self.config.hidden_size
# LayerNorm the backbone features before the head: the retriever's raw
Expand Down Expand Up @@ -165,7 +176,8 @@ def _selfcheck() -> None:
loss = nn.functional.cross_entropy(
logits.view(-1, 2).float(), labels.view(-1), ignore_index=-100
)
assert loss.requires_grad and loss.item() > 0
if not (loss.requires_grad and loss.item() > 0):
raise RuntimeError("selfcheck failed: loss not differentiable or not positive")
print("selfcheck ok")


Expand Down
Loading