#!/usr/bin/env python3
"""Compute teacher-forced SFT loss for a Hugging Face model, optionally with a PEFT adapter."""
from __future__ import annotations
import argparse
import json
import math
import re
from pathlib import Path
import pandas as pd
import torch
from transformers import AutoModelForCausalLM, AutoTokenizer
THINK_BLOCK_RE = re.compile(r"(?is)(.*?)")
ANSWER_BLOCK_RE = re.compile(r"(?is)(.*?)")
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(
description=(
"Compute assistant-token negative log likelihood on a VERL SFT parquet file. "
"The model must be a merged Hugging Face model directory or HF model id."
)
)
parser.add_argument("--model", required=True, help="Merged HF model directory or HF model id.")
parser.add_argument("--adapter", help="Optional PEFT/LoRA adapter directory.")
parser.add_argument("--tokenizer", help="Optional tokenizer path. Defaults to --model.")
parser.add_argument("--data", default="data/hep_sft/val.parquet", help="Validation/test parquet.")
parser.add_argument("--output", required=True, help="Summary JSON output path.")
parser.add_argument("--per-example-output", help="Optional per-example JSONL output path.")
parser.add_argument("--generations-output", help="Optional generated-answer JSONL output path.")
parser.add_argument("--limit", type=int, default=-1, help="Maximum examples to score; <=0 means all.")
parser.add_argument("--max-length", type=int, default=2048, help="Reject examples longer than this.")
parser.add_argument("--max-new-tokens", type=int, default=512, help="Maximum tokens for generation.")
parser.add_argument(
"--ignore-loss-between",
nargs=2,
action="append",
default=[("", "")],
metavar=("START", "END"),
help="Mask assistant loss for text inside START/END markers. Defaults to ....",
)
parser.add_argument(
"--no-ignore-think-loss",
action="store_true",
help="Do not apply the default ... loss mask.",
)
parser.add_argument("--trust-remote-code", action="store_true")
return parser.parse_args()
def to_messages(value: object) -> list[dict[str, str]]:
if isinstance(value, list):
return value
if hasattr(value, "tolist"):
converted = value.tolist()
if isinstance(converted, list):
return converted
raise TypeError(f"Unsupported messages value: {type(value)!r}")
def split_tagged_response(text: str) -> dict[str, object]:
think_match = THINK_BLOCK_RE.search(text)
answer_match = ANSWER_BLOCK_RE.search(text)
return {
"think": think_match.group(1).strip() if think_match else "",
"answer": answer_match.group(1).strip() if answer_match else "",
"has_think": think_match is not None,
"has_answer": answer_match is not None,
}
def score_example(
model: AutoModelForCausalLM,
tokenizer: AutoTokenizer,
messages: list[dict[str, str]],
max_length: int,
ignored_loss_spans: list[tuple[str, str]],
) -> tuple[float, int, int]:
prompt_messages = [message for message in messages if message.get("role") != "assistant"]
prompt_ids = tokenizer.apply_chat_template(
prompt_messages,
add_generation_prompt=True,
tokenize=True,
return_tensors="pt",
)
full_ids = tokenizer.apply_chat_template(
messages,
add_generation_prompt=False,
tokenize=True,
return_tensors="pt",
)
sequence_length = int(full_ids.shape[-1])
prompt_length = int(prompt_ids.shape[-1])
if sequence_length > max_length:
raise ValueError(f"{sequence_length=} exceeds {max_length=}")
if prompt_length >= sequence_length:
raise ValueError(f"{prompt_length=} is not shorter than {sequence_length=}")
labels = full_ids.clone()
labels[:, :prompt_length] = -100
mask_ignored_assistant_spans(
tokenizer=tokenizer,
input_ids=full_ids,
labels=labels,
assistant_start=prompt_length,
ignored_loss_spans=ignored_loss_spans,
)
assistant_tokens = int((labels != -100).sum().item())
if assistant_tokens <= 0:
raise ValueError("No assistant tokens to score")
device = next(model.parameters()).device
full_ids = full_ids.to(device)
labels = labels.to(device)
with torch.no_grad():
outputs = model(input_ids=full_ids, labels=labels)
loss = float(outputs.loss.detach().cpu())
return loss, assistant_tokens, sequence_length
def decode_token_ids(tokenizer: AutoTokenizer, token_ids: list[int]) -> str:
return tokenizer.decode(
token_ids,
skip_special_tokens=False,
clean_up_tokenization_spaces=False,
)
def decoded_prefix_lengths(tokenizer: AutoTokenizer, token_ids: list[int]) -> list[int]:
return [len(decode_token_ids(tokenizer, token_ids[:index])) for index in range(len(token_ids) + 1)]
def mask_ignored_assistant_spans(
tokenizer: AutoTokenizer,
input_ids: torch.Tensor,
labels: torch.Tensor,
assistant_start: int,
ignored_loss_spans: list[tuple[str, str]],
) -> None:
if not ignored_loss_spans:
return
token_ids = input_ids[0].tolist() if input_ids.ndim == 2 else input_ids.tolist()
decoded = decode_token_ids(tokenizer, token_ids)
prefix_lengths = decoded_prefix_lengths(tokenizer, token_ids)
assistant_start_char = prefix_lengths[assistant_start]
for start_marker, end_marker in ignored_loss_spans:
search_start = assistant_start_char
while True:
marker_start = decoded.find(start_marker, search_start)
if marker_start < 0:
break
inner_start = marker_start + len(start_marker)
marker_end = decoded.find(end_marker, inner_start)
if marker_end < 0:
break
while inner_start < marker_end and decoded[inner_start].isspace():
inner_start += 1
inner_end = marker_end
while inner_end > inner_start and decoded[inner_end - 1].isspace():
inner_end -= 1
for token_index in range(assistant_start, len(token_ids)):
token_start = prefix_lengths[token_index]
token_end = prefix_lengths[token_index + 1]
if token_end > inner_start and token_start < inner_end:
labels[:, token_index] = -100
search_start = marker_end + len(end_marker)
def generate_answer(
model: AutoModelForCausalLM,
tokenizer: AutoTokenizer,
messages: list[dict[str, str]],
max_new_tokens: int,
) -> tuple[str, str, str]:
prompt_messages = [message for message in messages if message.get("role") != "assistant"]
reference = next(
(str(message.get("content", "")) for message in messages if message.get("role") == "assistant"),
"",
)
prompt = str(prompt_messages[-1].get("content", "")) if prompt_messages else ""
inputs = tokenizer.apply_chat_template(
prompt_messages,
add_generation_prompt=True,
tokenize=True,
return_tensors="pt",
).to(next(model.parameters()).device)
with torch.no_grad():
generated = model.generate(
inputs,
max_new_tokens=max_new_tokens,
do_sample=False,
temperature=None,
top_p=None,
pad_token_id=tokenizer.eos_token_id,
)
output_ids = generated[0, inputs.shape[-1] :]
prediction = tokenizer.decode(output_ids, skip_special_tokens=True).strip()
return prompt, prediction, reference
def main() -> None:
args = parse_args()
ignored_loss_spans = [] if args.no_ignore_think_loss else [(str(a), str(b)) for a, b in args.ignore_loss_between]
output_path = Path(args.output)
output_path.parent.mkdir(parents=True, exist_ok=True)
per_example_path = Path(args.per_example_output) if args.per_example_output else None
if per_example_path:
per_example_path.parent.mkdir(parents=True, exist_ok=True)
generations_path = Path(args.generations_output) if args.generations_output else None
if generations_path:
generations_path.parent.mkdir(parents=True, exist_ok=True)
tokenizer = AutoTokenizer.from_pretrained(
args.tokenizer or args.model,
trust_remote_code=args.trust_remote_code,
)
model = AutoModelForCausalLM.from_pretrained(
args.model,
torch_dtype=torch.bfloat16,
device_map="auto",
trust_remote_code=args.trust_remote_code,
)
if args.adapter:
from peft import PeftModel
model = PeftModel.from_pretrained(model, args.adapter)
model.eval()
df = pd.read_parquet(args.data)
if args.limit > 0:
df = df.head(args.limit)
total_nll = 0.0
total_tokens = 0
scored = 0
skipped = 0
per_example_rows: list[dict[str, object]] = []
generation_rows: list[dict[str, object]] = []
for row in df.to_dict("records"):
row_id = row.get("id")
messages = to_messages(row["messages"])
try:
loss, assistant_tokens, sequence_length = score_example(
model=model,
tokenizer=tokenizer,
messages=messages,
max_length=args.max_length,
ignored_loss_spans=ignored_loss_spans,
)
except Exception as exc:
skipped += 1
record = {
"id": row_id,
"status": "skipped",
"error": str(exc),
}
per_example_rows.append(record)
print(f"skipped {row_id}: {exc}")
continue
if generations_path:
prompt, prediction, reference = generate_answer(
model=model,
tokenizer=tokenizer,
messages=messages,
max_new_tokens=args.max_new_tokens,
)
prediction_parts = split_tagged_response(prediction)
reference_parts = split_tagged_response(reference)
generation_rows.append(
{
"id": row_id,
"arxiv_id": row.get("arxiv_id"),
"task_type": row.get("task_type"),
"target_type": row.get("target_type"),
"target_category": row.get("target_category"),
"target_process_id": row.get("target_process_id"),
"target_background": row.get("target_background"),
"loss": loss,
"assistant_tokens": assistant_tokens,
"prompt": prompt,
"prediction": prediction,
"reference": reference,
"prediction_think": prediction_parts["think"],
"prediction_answer": prediction_parts["answer"],
"prediction_has_think": prediction_parts["has_think"],
"prediction_has_answer": prediction_parts["has_answer"],
"reference_think": reference_parts["think"],
"reference_answer": reference_parts["answer"],
"reference_has_think": reference_parts["has_think"],
"reference_has_answer": reference_parts["has_answer"],
}
)
total_nll += loss * assistant_tokens
total_tokens += assistant_tokens
scored += 1
record = {
"id": row_id,
"status": "scored",
"target_type": row.get("target_type"),
"target_category": row.get("target_category"),
"target_process_id": row.get("target_process_id"),
"target_background": row.get("target_background"),
"loss": loss,
"perplexity": math.exp(loss) if loss < 100 else float("inf"),
"assistant_tokens": assistant_tokens,
"sequence_length": sequence_length,
}
per_example_rows.append(record)
print(f"scored {row_id}: loss={loss:.6f} assistant_tokens={assistant_tokens}")
if total_tokens <= 0:
raise SystemExit("No assistant tokens were scored")
mean_loss = total_nll / total_tokens
summary = {
"adapter": args.adapter,
"model": args.model,
"data": args.data,
"examples_scored": scored,
"examples_skipped": skipped,
"assistant_tokens": total_tokens,
"loss": mean_loss,
"perplexity": math.exp(mean_loss) if mean_loss < 100 else float("inf"),
"max_length": args.max_length,
"ignored_loss_spans": ignored_loss_spans,
}
output_path.write_text(json.dumps(summary, indent=2, sort_keys=True) + "\n")
if per_example_path:
with per_example_path.open("w") as handle:
for record in per_example_rows:
handle.write(json.dumps(record, ensure_ascii=False) + "\n")
if generations_path:
with generations_path.open("w") as handle:
for record in generation_rows:
handle.write(json.dumps(record, ensure_ascii=False) + "\n")
print(json.dumps(summary, indent=2, sort_keys=True))
if __name__ == "__main__":
main()