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
43 changes: 37 additions & 6 deletions src/teich/audit.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,8 @@

from datasets import Dataset

from .protocol import chatml_turn_role_at_start


@dataclass
class SFTAuditReport:
Expand All @@ -31,6 +33,20 @@ def _as_list(value: Any) -> list[Any]:
return list(value)


def _contiguous_supervised_runs(labels: list[int]) -> list[tuple[int, int]]:
runs: list[tuple[int, int]] = []
run_start: int | None = None
for index, label in enumerate(labels):
if label != -100 and run_start is None:
run_start = index
elif label == -100 and run_start is not None:
runs.append((run_start, index))
run_start = None
if run_start is not None:
runs.append((run_start, len(labels)))
return runs


def _audit_training_row(row: dict[str, Any], tokenizer: Any, row_index: int) -> tuple[list[str], list[str], dict[str, Any]]:
errors: list[str] = []
warnings: list[str] = []
Expand Down Expand Up @@ -69,12 +85,24 @@ def _audit_training_row(row: dict[str, Any], tokenizer: Any, row_index: int) ->
supervised_text = _decode(tokenizer, supervised_ids)
sample["supervised_preview"] = supervised_text[:500]

# A non-assistant header at the beginning of a contiguous label run is
# unambiguously context. Protocol-looking strings later in assistant output
# may be quoted transcript examples and cannot be classified from labels
# alone; Teich's formatter validates those against marker-derived spans.
for run_start, run_end in _contiguous_supervised_runs(labels):
run_text = _decode(tokenizer, input_ids[run_start:run_end])
role = chatml_turn_role_at_start(run_text)
if role is not None and role != "assistant":
errors.append(
f"row {row_index}: supervised run begins at masked-context ChatML {role!r} header"
)
break

masked_ids = [token_id for token_id, label in zip(input_ids, labels) if label == -100]
masked_text = _decode(tokenizer, masked_ids[-200:]) if masked_ids else ""
sample["masked_suffix_preview"] = masked_text[-500:]

suspicious_masked_markers = (
"<|im_start|>user",
"<|start_header_id|>user<|end_header_id|>",
"<start_of_turn>user",
"<|start_of_role|>user<|end_of_role|>",
Expand Down Expand Up @@ -111,14 +139,17 @@ def audit_sft_dataset(dataset: Dataset, tokenizer: Any, *, sample_size: int | No
if dataset.num_rows == 0:
return SFTAuditReport(ok=False, errors=["dataset contains no rows"])

limit = dataset.num_rows if sample_size is None else min(max(sample_size, 0), dataset.num_rows)
if limit == 0:
warnings.append("sample_size is 0; no rows audited")
preview_limit = dataset.num_rows if sample_size is None else min(max(sample_size, 0), dataset.num_rows)
if preview_limit == 0:
warnings.append("sample_size is 0; no row previews retained (all rows were still audited)")

for row_index in range(limit):
# Sampling controls report size only. Correctness checks are deliberately
# exhaustive so an audit that returns ok=True is a dataset-wide gate.
for row_index in range(dataset.num_rows):
row_errors, row_warnings, sample = _audit_training_row(dataset[row_index], tokenizer, row_index)
errors.extend(row_errors)
warnings.extend(row_warnings)
samples.append(sample)
if row_index < preview_limit:
samples.append(sample)

return SFTAuditReport(ok=not errors, errors=errors, warnings=warnings, samples=samples)
Loading
Loading