# /// script
# requires-python = ">=3.12,<3.13"
# dependencies = [
# "torch==2.13.0",
# "transformers==5.16.1",
# "huggingface-hub==1.31.0",
# "pyarrow==25.0.1",
# "safetensors==0.8.0",
# "tokenizers==0.23.2",
# ]
# ///
"""Self-contained TabFix joint detector/denoiser. See --help; no repository imports.
Run on HF Jobs: hf jobs uv run --flavor l40sx1 --timeout 100m --secrets HF_TOKEN
train_tabfix.py --repo-id YOUR_NAME/tabfix-preview --train-seconds 3600
Public checkpoint includes this script, tokenizer, weights, optimizer and RNG state.
"""
from __future__ import annotations
import argparse
from collections import Counter
from contextlib import nullcontext
from dataclasses import asdict, dataclass
from decimal import Decimal, InvalidOperation
from datetime import date, datetime
from functools import lru_cache
import hashlib
import html
import json
import math
import os
from pathlib import Path
import random
import re
import shutil
import signal
import time
import xml.etree.ElementTree as ET
from typing import Iterator, Protocol, cast
import pyarrow.parquet as pq
import torch
from torch import Tensor, nn
import torch.nn.functional as F
from huggingface_hub import HfApi, hf_hub_download, snapshot_download # pyright: ignore[reportUnknownVariableType] — third-party typing boundary
from safetensors.torch import load_file, save_file # pyright: ignore[reportUnknownVariableType] — third-party typing boundary
from tokenizers import Tokenizer, pre_tokenizers
from transformers import AutoTokenizer, ModernBertConfig, ModernBertForMaskedLM, PreTrainedTokenizerFast
DATASET = "Antix5/tabular-errors-v1"
REVISION = "602b8a454ca448713d8a5658650d4feaad8d579c"
BASE = "jhu-clsp/mmBERT-base"
BASE_REVISION = "c5955035435e2bf121cde7f3c8863ef52ff35d82"
BUSINESS_CATEGORIES = (
"text.encoding", "text.invisible", "text.spelling", "format.number",
"format.date", "format.boolean", "format.case", "format.whitespace",
"format.unit", "schema.type", "schema.identifier", "schema.category",
"consistency.unit", "consistency.dependency", "consistency.temporal",
"missing.required", "missing.conditional", "structure.shift",
)
CONTROL = ("<|tabfix_repair|>", "<|tabfix_end|>", "<|tabfix_pad|>")
CATEGORIES = ("text.encoding", "text.spelling")
K = len(CATEGORIES)
stop_requested = False
def mapping(value: object) -> dict[str, object]:
if not isinstance(value, dict):
raise TypeError("Expected JSON object")
return cast(dict[str, object], value)
def sequence(value: object) -> list[object]:
if not isinstance(value, list):
raise TypeError("Expected JSON list")
return cast(list[object], value)
def integer(value: object) -> int:
if not isinstance(value, int):
raise TypeError("Expected integer")
return value
def string(value: object) -> str:
if not isinstance(value, str):
raise TypeError("Expected string")
return value
def number(value: object) -> float:
if not isinstance(value, (int, float)):
raise TypeError("Expected number")
return float(value)
def read_json(path: Path) -> dict[str, object]:
return mapping(cast(object, json.loads(path.read_text())))
def write_json(path: Path, value: object) -> None:
path.write_text(json.dumps(value, ensure_ascii=False, indent=2) + "\n")
def emit(event: str, **fields: object) -> None:
print(json.dumps({"event": event, **fields}, ensure_ascii=False), flush=True)
@dataclass(frozen=True)
class RuleFinding:
row: int
columns: tuple[int, ...]
category: str
rule: str
status: str
evidence: str
def xml_json(element: ET.Element, name: str, default: object) -> object:
text = element.findtext(name)
return cast(object, json.loads(text)) if text else default
def column_checks(value: str, column: ET.Element) -> list[tuple[str, str, bool | None]]:
"""Check declared contracts only; None means unsupported/undetermined."""
rules = mapping(xml_json(column, "constraints", {}))
formats = [string(x) for x in sequence(xml_json(column, "accepted_formats", []))]
nulls = sequence(xml_json(column, "null_markers", []))
if not value or value in nulls:
return [("missing.required", "nullability", column.get("nullable") != "forbidden")]
result: list[tuple[str, str, bool | None]] = []
for key in ("enum", "lexicon"):
if rules.get(key):
result.append(("format.boolean" if column.get("type") == "boolean" else "schema.category", key, value in sequence(rules[key])))
if isinstance(rules.get("pattern"), str):
try:
result.append(("schema.identifier", "pattern", re.fullmatch(string(rules["pattern"]), value) is not None))
except re.error:
result.append(("schema.identifier", "pattern", None))
for key in ("upper", "lower"):
if rules.get("case") == key:
result.append(("format.case", "case", value == (value.upper() if key == "upper" else value.lower())))
if rules.get("trimmed") is True:
result.append(("format.whitespace", "trimmed", value == value.strip()))
forbidden = rules.get("forbidden_characters")
if isinstance(forbidden, list):
result.append(("text.invisible", "forbidden_characters", all(string(c) not in value for c in sequence(cast(object, forbidden)))))
kind = column.get("type")
if value in sequence(rules.get("accepted_representations", [])):
return result
if kind in ("integer", "decimal", "numeric_syntax"):
patterns: list[str] = []
if kind == "integer":
patterns = [r"[+-]?[0-9]+"]
elif kind == "numeric_syntax":
patterns = [r"[+-]?(?=[0-9., \u00a0\u202f]*[0-9])[0-9., \u00a0\u202f]+(?:[eE][+-]?[0-9]+)?"]
else:
for fmt in formats:
if fmt.startswith("decimal_"):
sep = "," if "comma" in fmt else r"\."
patterns.append(r"[+-]?[0-9]+" + (sep + r"[0-9]{2}" if "2dp" in fmt else "(?:" + sep + r"[0-9]+)?"))
valid = any(re.fullmatch(p, value) for p in patterns) if patterns else None
result.append(("schema.type" if kind == "integer" else "format.number", "numeric_format", bool(valid) if valid is not None else None))
if valid and kind != "numeric_syntax":
parsed = Decimal(value.replace(",", "."))
for key in ("minimum", "maximum"):
bound = rules.get(key)
if isinstance(bound, (int, float)):
result.append(("schema.type", key, parsed >= Decimal(str(bound)) if key == "minimum" else parsed <= Decimal(str(bound))))
elif kind == "date":
fmts = {"YYYY-MM-DD": "%Y-%m-%d", "DD/MM/YYYY": "%d/%m/%Y", "MM/DD/YYYY": "%m/%d/%Y"}
valid_date = False
supported = bool(formats) and all(f in fmts for f in formats)
for fmt in formats:
if fmt in fmts:
try:
dt = datetime.strptime(value, fmts[fmt])
valid_date |= dt.strftime(fmts[fmt]) == value
except ValueError:
pass
result.append(("format.date", "calendar_format", valid_date if supported else None))
return result
def deterministic_findings(xml: str) -> list[RuleFinding]:
root = ET.fromstring(xml)
columns = root.findall("./schema/column")
schema = root.find("schema")
if schema is None:
return []
keys = {c.get("key", ""): i for i, c in enumerate(columns) if c.get("key")}
relations = sequence(xml_json(schema, "relations", []))
dependencies = mapping(xml_json(schema, "functional_dependencies", {}))
findings: list[RuleFinding] = []
for row in root.findall("./rows/row"):
ri = int(row.attrib["index"])
values = [c.text or "" for c in row.findall("cell")]
if len(values) != len(columns):
continue
checks = [column_checks(v, c) for v, c in zip(values, columns)]
for ci, (value, col, cc) in enumerate(zip(values, columns, checks)):
for cat, rule, valid in cc:
findings.append(RuleFinding(ri, (ci,), cat, rule, "unknown" if valid is None else "valid" if valid else "invalid", "declared column contract"))
rules = mapping(xml_json(col, "constraints", {}))
condition = rules.get("required_when")
if isinstance(condition, dict):
cond = mapping(cast(object, condition))
determinant = cond.get("column")
valid_condition = isinstance(determinant, int) and 0 <= determinant < len(values)
known = valid_condition and bool(values[cast(int, determinant)]) and all(v is True for _, _, v in checks[cast(int, determinant)])
required: bool | None = None
if known:
dv = values[cast(int, determinant)]
if "equals" in cond:
required = dv == cond["equals"]
elif "not_equals" in cond:
required = dv != cond["not_equals"]
present = bool(value) and value not in sequence(xml_json(col, "null_markers", []))
findings.append(RuleFinding(ri, (ci,), "missing.conditional", "required_when", "unknown" if required is None else "valid" if present or not required else "invalid", "declared condition; determinant must be valid"))
for raw in relations:
if not isinstance(raw, list) or not raw:
continue
rel = [string(x) for x in sequence(cast(object, raw))]
kind, args = rel[0], rel[1:]
if any(key not in keys for key in args):
continue
ix = tuple(keys[key] for key in args)
cat = "consistency.unit" if kind in ("unit_equivalence", "pack_mass", "unit_factor") else "consistency.temporal" if kind == "interval" else "consistency.dependency"
if kind == "conditional":
continue # Executed through required_when.
status = "unknown"
if all(values[i] and all(v is True for _, _, v in checks[i]) for i in ix):
vv = [values[i] for i in ix]
try:
valid: bool | None = None
if kind == "lookup" and len(ix) == 2:
reference = mapping(dependencies.get(args[1], {})).get(vv[0])
valid = vv[1] == reference if reference is not None else None
elif kind == "interval" and len(ix) == 3:
start, end = date.fromisoformat(vv[0]), date.fromisoformat(vv[1])
valid = end >= start and Decimal(vv[2]) == (end - start).days
elif kind == "unit_equivalence" and len(ix) == 4:
factors = {"kg": Decimal(1000), "g": Decimal(1), "mg": Decimal(".001")}
if vv[1] in factors and vv[3] in factors:
valid = Decimal(vv[0].replace(",", ".")) * factors[vv[1]] == Decimal(vv[2].replace(",", ".")) * factors[vv[3]]
elif kind in ("difference", "multiply", "balance", "le"):
nums = [Decimal(v.replace(",", ".")) for v in vv]
if kind == "le" and len(nums) == 2:
valid = nums[0] <= nums[1]
elif len(nums) == 3:
expected = nums[0] * nums[1] if kind == "multiply" else nums[0] - nums[1]
# Compare at the explicitly declared output precision.
formats = sequence(xml_json(columns[ix[-1]], "accepted_formats", []))
if any("2dp" in string(f) for f in formats):
expected = expected.quantize(Decimal(".01"))
valid = expected == nums[2]
status = "unknown" if valid is None else "valid" if valid else "invalid"
except (ValueError, InvalidOperation):
status = "unknown"
findings.append(RuleFinding(ri, ix, cat, kind, status, "relation contradiction identifies a group, not its faulty member"))
return findings
def deterministic_repair(value: str, column: ET.Element) -> str | None:
"""Only unambiguous canonicalization, never fill missing information."""
if not value:
return None
rules = mapping(xml_json(column, "constraints", {}))
candidate = value.strip() if rules.get("trimmed") is True else value
if rules.get("case") == "upper":
candidate = candidate.upper()
elif rules.get("case") == "lower":
candidate = candidate.lower()
enum = sequence(rules.get("enum", []))
equivalent = [string(v) for v in enum if string(v).casefold() == candidate.casefold()]
if candidate not in enum and len(equivalent) == 1 and column.get("type") == "boolean":
candidate = equivalent[0]
checks = column_checks(candidate, column)
if candidate != value and checks and all(ok is True for _, _, ok in checks):
return candidate
return None
def replace_cell_xml(xml: str, row: int, column: int, value: str) -> str:
cell = next(c for c in cells(xml) if (c.row, c.column) == (row, column))
escaped = html.escape(value, quote=False).replace("\r", "
").replace("\n", "
").replace("\t", " ") if value else ""
return xml[:cell.anchor] + escaped + xml[cell.close:]
def invalid_cells(xml: str) -> set[tuple[int, int]]:
return {(f.row, c) for f in deterministic_findings(xml) if f.status == "invalid" for c in f.columns}
@dataclass(frozen=True)
class Edit:
row: int
column: int
start: int
end: int
replacement: str
category: str
@classmethod
def parse(cls, item: object) -> Edit:
d = mapping(item)
return cls(integer(d["row"]), integer(d["column"]), integer(d["start"]),
integer(d["end"]), string(d["replacement"]), string(d["category"]))
@dataclass
class Record:
id: str
xml: str
edits: list[Edit]
metadata: dict[str, object]
@classmethod
def parse(cls, item: object) -> Record:
d = mapping(item)
return cls(string(d["id"]), string(d["corrupt_xml"]),
[Edit.parse(e) for e in sequence(cast(object, json.loads(string(d["errors"]))))],
mapping(cast(object, json.loads(string(d["metadata"])))) )
@dataclass
class Cell:
row: int
column: int
value: str
# Each decoded character maps to one interval in the serialized XML.
bounds: list[tuple[int, int]]
anchor: int
close: int
def cells(xml: str) -> list[Cell]:
result: list[Cell] = []
for row in re.finditer(r'(.*?)
', xml, re.S):
for col, cell in enumerate(re.finditer(r']*>(.*?) | ', row[2], re.S)):
raw = cell[1]
origin = row.start(2) + cell.start(1)
bounds: list[tuple[int, int]] = []
chars: list[str] = []
if raw != "":
for part in re.finditer(r'&[^;]+;|[^&]', raw, re.S):
decoded = html.unescape(part[0])
if len(decoded) != 1:
raise ValueError("Unsupported XML entity")
chars.append(decoded)
bounds.append((origin + part.start(), origin + part.end()))
result.append(Cell(int(row[1]), col, "".join(chars), bounds, origin, origin + len(raw)))
if not result:
raise ValueError("No cells in XML")
return result
@dataclass
class Prepared:
record: Record
ids: list[int]
offsets: list[tuple[int, int]]
cells: list[Cell]
# positions carry a category BIO label vector; gaps carry binary labels.
positions: list[int]
labels: list[list[int]]
gaps: list[int]
gap_labels: list[list[int]]
collisions: int = 0
def tokens_for(cell: Cell, offsets: list[tuple[int, int]]) -> list[int]:
if not cell.bounds:
return []
lo, hi = cell.bounds[0][0], cell.bounds[-1][1]
return [i for i, (a, b) in enumerate(offsets) if b > lo and a < hi and b > a]
def interval(cell: Cell, offsets: list[tuple[int, int]], edit: Edit) -> tuple[list[int], int, int]:
"""Token-expanded edit bounds in decoded coordinates; no XML edits."""
tok = tokens_for(cell, offsets)
if edit.start == edit.end:
point = cell.bounds[edit.start][0] if edit.start < len(cell.bounds) else cell.close
hit = [i for i in tok if offsets[i][0] < point < offsets[i][1]]
else:
lo, hi = cell.bounds[edit.start][0], cell.bounds[edit.end - 1][1]
hit = [i for i in tok if offsets[i][1] > lo and offsets[i][0] < hi]
if not hit:
return [], edit.start, edit.end
low, high = offsets[hit[0]][0], offsets[hit[-1]][1]
covered = [i for i, (a, b) in enumerate(cell.bounds) if b > low and a < high]
return hit, min(covered), max(covered) + 1
def anchor_token(cell: Cell, offsets: list[tuple[int, int]], at: int) -> int:
point = cell.bounds[at][0] if at < len(cell.bounds) else cell.close
return next(i for i, (a, b) in enumerate(offsets) if b > a and b > point and a <= point)
class Consumer:
def __init__(self, tokenizer: PreTrainedTokenizerFast, max_length: int = 8192):
self.tokenizer = tokenizer
self.max_length = max_length
self.stats: Counter[str] = Counter()
self.literal_tokenizer = Tokenizer.from_str(tokenizer.backend_tokenizer.to_str())
self.literal_tokenizer.pre_tokenizer = pre_tokenizers.Metaspace(replacement="▁", prepend_scheme="never") # pyright: ignore[reportUnknownMemberType] — third-party typing boundary
def encode(self, text: str) -> list[int]:
return self.tokenizer.encode(text, add_special_tokens=False) # pyright: ignore[reportUnknownMemberType] — third-party typing boundary
def literal(self, text: str) -> list[int]:
ids = self.literal_tokenizer.encode(text, add_special_tokens=False).ids
if any(i in self.tokenizer.all_special_ids for i in ids):
raise ValueError("Literal text contains a reserved special token")
if self.literal_tokenizer.decode(ids, skip_special_tokens=False) != text:
raise ValueError("Replacement tokenization is not lossless")
return ids
def prepare(self, record: Record) -> Prepared:
encoded = self.tokenizer(record.xml, return_offsets_mapping=True, truncation=False)
ids = cast(list[int], encoded["input_ids"])
offsets = cast(list[tuple[int, int]], encoded["offset_mapping"])
if len(ids) > self.max_length:
raise ValueError("Context exceeds token budget; never truncate cells")
cc = cells(record.xml)
targets = {integer(x) for x in sequence(record.metadata["target_rows"])}
columns = {integer(x) for x in sequence(record.metadata["supervised_columns"])}
trusted = record.metadata["label_scope"] == "validated_synthetic_contracts"
# Real examples certify injected labels, not absence of every other defect.
labels: dict[int, list[int]] = {}
gaps: dict[int, list[int]] = {}
for c in cc:
if c.row not in targets or c.column not in columns:
continue
for t in tokens_for(c, offsets):
labels[t] = [0 if trusted else -100] * K
boundaries = {0, len(c.value)}
for t in tokens_for(c, offsets):
boundaries.update(j for j, (a, _) in enumerate(c.bounds) if a == offsets[t][0])
for at in boundaries:
gaps[anchor_token(c, offsets, at)] = [0 if trusted else -100] * K
collisions = 0
neural_indices = record.metadata.get("neural_error_indices")
for ei, e in enumerate(record.edits):
if e.category not in CATEGORIES or (isinstance(neural_indices, list) and ei not in neural_indices):
continue
c = next(c for c in cc if (c.row, c.column) == (e.row, e.column))
hit, _, _ = interval(c, offsets, e)
category = CATEGORIES.index(e.category)
if hit:
for j, t in enumerate(hit):
target = labels.setdefault(t, [-100] * K)
if target[category] > 0:
collisions += 1
target[category] = 1 if j == 0 else 2
else:
gaps.setdefault(anchor_token(c, offsets, e.start), [-100] * K)[category] = 1
# Deterministic cells are not part of the residual neural objective.
blocked: set[tuple[int, int]] = invalid_cells(record.xml) if record.metadata.get("task_version") == 2 else set()
for c in cc:
if (c.row, c.column) in blocked:
for t in tokens_for(c, offsets):
labels.pop(t, None)
for at in range(len(c.value) + 1):
gaps.pop(anchor_token(c, offsets, at), None)
pos = sorted(labels)
gp = sorted(gaps)
return Prepared(record, ids, offsets, cc, pos, [labels[p] for p in pos], gp,
[gaps[p] for p in gp], collisions)
def correction(self, p: Prepared, e: Edit, rng: random.Random, training: bool,
capacity: int | None = None, mask_all: bool = False) -> tuple[list[int], list[int], list[int], list[bool], str]:
c = next(c for c in p.cells if (c.row, c.column) == (e.row, e.column))
# Whole-cell target, original retained locally. The edit is supplied by
# correction_candidates; its replacement is never put into inference input.
target = c.value[:e.start] + e.replacement + c.value[e.end:]
gold = self.literal(target) if training else []
n = len(self.literal(c.value))
budget = capacity or min(256, max(8, n + max(2, math.ceil(.2 * n))))
if training and len(gold) > budget:
raise ValueError("Whole-cell repair exceeds source-derived buffer")
mask_id = integer(self.tokenizer.mask_token_id)
original = p.record.xml[c.anchor:c.close]
local = "" + original + "" + string(self.tokenizer.mask_token) * (budget + 1) + ""
skeleton = p.record.xml[:c.anchor] + local + p.record.xml[c.close:]
ids = self.tokenizer.encode(skeleton, add_special_tokens=True) # pyright: ignore[reportUnknownMemberType] — third-party typing boundary
positions = [i for i, token in enumerate(ids) if token == mask_id]
if len(positions) != budget + 1:
raise ValueError("Reserved mask token appears outside the selected edit")
if len(ids) > self.max_length:
raise ValueError("In-place correction exceeds context budget")
end = integer(self.tokenizer.convert_tokens_to_ids(CONTROL[1]))
pad = integer(self.tokenizer.convert_tokens_to_ids(CONTROL[2]))
answer: list[int]
if not training:
answer = [integer(self.tokenizer.mask_token_id)] * (budget + 1)
return ids, positions, [], [], target
answer = gold + [end] + [pad] * (budget - len(gold))
count = len(answer) if mask_all or rng.random() < .5 else rng.randint(1, len(answer))
chosen = sorted(rng.sample(range(len(answer)), count))
inputs = answer.copy()
for i in chosen:
inputs[i] = integer(self.tokenizer.mask_token_id)
for j, token in enumerate(inputs):
ids[positions[j]] = token
return ids, [positions[i] for i in chosen], [answer[i] for i in chosen], [i > len(gold) for i in chosen], target
def eligible(record: Record) -> list[Edit]:
"""Whole-cell targets, including certified valid copies, stored in release v2."""
candidates = record.metadata.get("correction_candidates")
if isinstance(candidates, list):
return [Edit.parse(e) for e in sequence(cast(object, candidates))]
# Local callers/tests can supply one selected edit directly.
result: list[Edit] = []
for value in sequence(record.metadata.get("repair_policies", [])):
policy = mapping(value)
if policy.get("correction_supervision") in ("reference_repair", "schema_determined", "target_conditioned"):
result.append(record.edits[integer(policy["error_index"])])
return result
class Batch(Protocol):
def to_pylist(self) -> list[object]: ...
class ParquetReader(Protocol):
def iter_batches(self, batch_size: int) -> Iterator[Batch]: ...
class HiddenOutput(Protocol):
last_hidden_state: Tensor
class TabFix(nn.Module):
def __init__(self, mlm: ModernBertForMaskedLM):
super().__init__()
self.mlm = mlm
self.bio = nn.Linear(mlm.config.hidden_size, K * 3)
self.insert = nn.Linear(mlm.config.hidden_size, K)
def hidden_batch(self, sequences: list[list[int]], device: torch.device) -> Tensor:
length = max(map(len, sequences))
x = torch.full((len(sequences), length), integer(self.mlm.config.pad_token_id), dtype=torch.long, device=device)
attention = torch.zeros_like(x)
for i, ids in enumerate(sequences):
x[i, :len(ids)] = torch.tensor(ids, dtype=torch.long, device=device)
attention[i, :len(ids)] = 1
output = cast(HiddenOutput, self.mlm.model(input_ids=x, attention_mask=attention))
return output.last_hidden_state
def hidden(self, ids: list[int], device: torch.device) -> Tensor:
return self.hidden_batch([ids], device)[0]
def vocab(self, hidden: Tensor) -> Tensor:
return cast(Tensor, self.mlm.decoder(self.mlm.head(hidden)))
def mean_parts(losses: Tensor, positive: Tensor, valid: Tensor) -> Tensor:
"""Equal positive/background group weight; no loss if group is absent."""
pieces: list[Tensor] = []
for mask in (valid & positive, valid & ~positive):
if bool(mask.any()):
pieces.append(losses[mask].mean())
return torch.stack(pieces).mean() if pieces else losses.sum() * 0
def detection_loss(model: TabFix, p: Prepared, device: torch.device) -> Tensor:
return detection_from_hidden(model, p, model.hidden(p.ids, device))
def detection_from_hidden(model: TabFix, p: Prepared, h: Tensor) -> Tensor:
device = h.device
labels = torch.tensor(p.labels, device=device, dtype=torch.long).reshape(-1, K)
logits = model.bio.forward(h[p.positions]).reshape(-1, K, 3)
loss = F.cross_entropy(logits.reshape(-1, 3).float(), labels.reshape(-1), ignore_index=-100, reduction="none").reshape(-1, K)
bio = mean_parts(loss, labels > 0, labels != -100)
gl = torch.tensor(p.gap_labels, device=device, dtype=torch.float32).reshape(-1, K)
gi = model.insert.forward(h[p.gaps])
il = F.binary_cross_entropy_with_logits(gi.float(), gl.clamp_min(0), reduction="none")
return bio + .5 * mean_parts(il, gl > 0, gl != -100)
def correction_loss(model: TabFix, view: tuple[list[int], list[int], list[int], list[bool], str], device: torch.device) -> Tensor:
return correction_from_hidden(model, view, model.hidden(view[0], device))
def correction_from_hidden(model: TabFix, view: tuple[list[int], list[int], list[int], list[bool], str], h: Tensor) -> Tensor:
_, positions, targets, tail, _ = view
device = h.device
logits = model.vocab(h[positions])
losses = F.cross_entropy(logits.float(), torch.tensor(targets, device=device), reduction="none")
is_tail = torch.tensor(tail, device=device)
main = losses[~is_tail].mean() if bool((~is_tail).any()) else losses.sum() * 0
padding = losses[is_tail].mean() if bool(is_tail.any()) else losses.sum() * 0
return main + .1 * padding
@dataclass
class RunConfig:
repo_id: str = ""
output: str = "tabfix_run"
local_data: str = ""
tokenizer: str = ""
resume: str = ""
resume_revision: str = "main"
tiny: bool = False
steps: int = 100000
train_seconds: int = 1200
save_seconds: int = 420
lr: float = 2e-5
seed: int = 20260913
validation_examples: int = 252
validation_repairs: int = 7
validation_seconds: int = 600
validation_steps: int = 200
patience: int = 4
min_delta: float = 0.001
warmup_steps: int = 100
schedule_steps: int = 2400
max_length: int = 8192
correction_weight: float = math.log(K) / math.log(256000)
batch_size: int = 2
verify_only: bool = False
def learning_rate(cfg: RunConfig, step: int) -> float:
if step < cfg.warmup_steps:
return cfg.lr * (step + 1) / max(1, cfg.warmup_steps)
progress = min(1., (step - cfg.warmup_steps) / max(1, cfg.schedule_steps - cfg.warmup_steps))
return cfg.lr * (0.1 + 0.9 * (1 + math.cos(math.pi * progress)) / 2)
@dataclass
class Selection:
best_score: float = math.inf
best_step: int = 0
bad_checks: int = 0
def observe(self, score: float, step: int, delta: float) -> bool:
if not math.isfinite(score):
raise FloatingPointError("Nonfinite held-out validation score")
if score < self.best_score - delta:
self.best_score, self.best_step, self.bad_checks = score, step, 0
return True
self.bad_checks += 1
return False
def read_records(split: str, cfg: RunConfig) -> list[Record]:
if split not in ("train", "validation"):
raise ValueError("Test split is intentionally not consumed in development")
if cfg.local_data:
path = Path(cfg.local_data) / "data" / f"{split}.parquet"
else:
path = Path(hf_hub_download(DATASET, f"data/{split}.parquet", repo_type="dataset", revision=REVISION))
manifest = read_json(Path(hf_hub_download(DATASET, "release.json", repo_type="dataset", revision=REVISION)))
assert hashlib.sha256(path.read_bytes()).hexdigest() == mapping(manifest["files"])[f"data/{split}.parquet"]
records: list[Record] = []
for batch in cast(ParquetReader, pq.ParquetFile(path)).iter_batches(batch_size=256):
records.extend(Record.parse(x) for x in batch.to_pylist())
return records
def autocast(device: torch.device):
return torch.autocast("cuda", dtype=torch.bfloat16) if device.type == "cuda" else nullcontext()
def make_model(cfg: RunConfig) -> tuple[TabFix, PreTrainedTokenizerFast]:
source = cfg.tokenizer or BASE
tok = cast(PreTrainedTokenizerFast, AutoTokenizer.from_pretrained(source, revision=None if cfg.tokenizer else BASE_REVISION)) # pyright: ignore[reportUnknownMemberType] — third-party typing boundary
_ = tok.add_special_tokens({"additional_special_tokens": list(CONTROL)})
if cfg.tiny:
config = ModernBertConfig(vocab_size=len(tok), hidden_size=32, intermediate_size=64,
num_hidden_layers=2, num_attention_heads=4, max_position_embeddings=8192,
pad_token_id=integer(tok.pad_token_id), bos_token_id=integer(tok.bos_token_id), eos_token_id=integer(tok.eos_token_id))
config._attn_implementation = "sdpa" # pyright: ignore[reportPrivateUsage] — third-party typing boundary
mlm = ModernBertForMaskedLM(config)
else:
mlm = ModernBertForMaskedLM.from_pretrained(BASE, revision=BASE_REVISION, attn_implementation="sdpa") # pyright: ignore[reportUnknownMemberType] — third-party typing boundary
_ = mlm.resize_token_embeddings(len(tok), mean_resizing=False)
return TabFix(mlm), tok
def load_checkpoint(path: Path, device: torch.device) -> tuple[TabFix, PreTrainedTokenizerFast, dict[str, object]]:
manifest = read_json(path / "checkpoint.json")
for name, digest in mapping(manifest["files"]).items():
if hashlib.sha256((path / name).read_bytes()).hexdigest() != digest:
raise ValueError(f"Checkpoint checksum mismatch: {name}")
tok = cast(PreTrainedTokenizerFast, AutoTokenizer.from_pretrained(path)) # pyright: ignore[reportUnknownMemberType] — third-party typing boundary
config = ModernBertConfig.from_pretrained(path) # pyright: ignore[reportUnknownMemberType] — third-party typing boundary
config._attn_implementation = "sdpa" # pyright: ignore[reportPrivateUsage] — third-party typing boundary
mlm = ModernBertForMaskedLM(config)
model = TabFix(mlm)
saved_categories = cast(object, json.loads((path / "categories.json").read_text()))
if saved_categories != list(CATEGORIES):
raise ValueError("Checkpoint uses another label map; use its bundled training script")
# Store tied embedding once; the duplicate decoder weight shares storage.
weights = load_file(str(path / "model.safetensors"))
weights["mlm.decoder.weight"] = weights["mlm.model.embeddings.tok_embeddings.weight"]
_ = model.load_state_dict(weights, strict=True)
model.to(device)
return model, tok, manifest
def save_checkpoint(model: TabFix, tok: PreTrainedTokenizerFast, optimizer: torch.optim.Optimizer,
cfg: RunConfig, step: int, rng: random.Random, report: dict[str, object],
api: HfApi | None) -> Path:
root = Path(cfg.output)
root.mkdir(parents=True, exist_ok=True)
stage = root / "staging"
stage.mkdir(exist_ok=True)
weights = {k: v.detach().cpu().contiguous() for k, v in cast(dict[str, Tensor], model.state_dict()).items() if k != "mlm.decoder.weight"}
save_file(weights, str(stage / "model.safetensors"))
del weights
model.mlm.config.save_pretrained(stage) # pyright: ignore[reportUnknownMemberType] — third-party typing boundary
_ = tok.save_pretrained(stage) # pyright: ignore[reportUnknownMemberType] — third-party typing boundary
torch.save({"optimizer": optimizer.state_dict(), "step": step, "python_rng": rng.getstate(),
"torch_rng": torch.get_rng_state(), "cuda_rng": torch.cuda.get_rng_state_all() if torch.cuda.is_available() else []}, stage / "training_state.pt")
write_json(stage / "run_config.json", asdict(cfg))
write_json(stage / "metrics.json", report)
write_json(stage / "categories.json", list(CATEGORIES))
shutil.copy2(__file__, stage / "train_tabfix.py")
(stage / "README.md").write_text(model_card(cfg, step, report))
(stage / "LICENSE").write_text("TabFix fine-tuning contributions: CC BY-NC 4.0. Commercial use requires the owner's permission.\nhttps://creativecommons.org/licenses/by-nc/4.0/legalcode\n\nThe underlying mmBERT weights retain their MIT license and attribution; see BASE_MODEL_CARD.md.\n")
if not cfg.tiny:
base_license = hf_hub_download(BASE, "README.md", revision=BASE_REVISION)
shutil.copy2(base_license, stage / "BASE_MODEL_CARD.md")
files = {p.name: hashlib.sha256(p.read_bytes()).hexdigest() for p in stage.iterdir() if p.is_file() and p.name != "checkpoint.json"}
write_json(stage / "checkpoint.json", {"version": 1, "step": step, "dataset_revision": REVISION,
"base_revision": BASE_REVISION, "files": files, "complete": True})
current = root / "checkpoint"
previous = root / "previous"
if previous.exists():
shutil.rmtree(previous)
if current.exists():
current.rename(previous)
stage.rename(current)
# One Hub commit atomically replaces all files; an interrupted upload retains old commit.
if api and cfg.repo_id:
commit = api.upload_folder(repo_id=cfg.repo_id, folder_path=current, commit_message=f"Joint training checkpoint step {step}")
emit("checkpoint_uploaded", step=step, commit=str(commit.oid))
write_json(root / "last_upload.json", {"commit": commit.oid, "step": step})
emit("checkpoint_saved", step=step, path=str(current))
return current
def model_card(cfg: RunConfig, step: int, report: dict[str, object]) -> str:
return f'''---
license: cc-by-nc-4.0
language: [en, fr, nl, de, pl, el, ja]
base_model: {BASE}
datasets: [{DATASET}]
tags: [tabular-data, masked-lm, token-classification, experimental]
---
# TabFix shared-encoder detector and masked corrector
{'Temporary random tiny smoke-test model, not a pretrained fine-tune.' if cfg.tiny else 'Experimental budget-limited mmBERT-base joint fine-tuning checkpoint.'}
One shared encoder, {K} residual category BIO/insertion outputs (`text.encoding`, `text.spelling`), and the pretrained MLM vocabulary head. Declared schema rules detect deterministic violations before neural classification. Correction generates a complete cell inside `` beside the visible `` in the same cell. Valid copies are training targets too. The MLM head emits replacement tokens followed by `END_EDIT` and visible padding tokens.
Completed optimizer updates: **{step}**. This checkpoint is a research baseline, not a validated automatic cleaning system.
## Use and resume
The custom model must be loaded through the accompanying self-contained `train_tabfix.py`, not `AutoModelForMaskedLM` directly. Its weights include detection heads and have a custom key layout. `load_checkpoint(path, device)` returns the model and tokenizer; `Consumer`, `predict_edits` and `generate_repair` implement input alignment and suggestions. The command `uv run train_tabfix.py --resume {cfg.repo_id or 'LOCAL_CHECKPOINT'} --verify-only` verifies checksums and loads the model. Resume training with `--resume ... --repo-id ... --train-seconds 1200`; hardware timeouts must be set separately.
The recoverable state includes model, tokenizer, optimizer and random generators. `training_state.pt` uses PyTorch serialization: load only from a trusted checkpoint you control. No API keys are saved. The original training script and resolved package pins are embedded in the repository.
## Scope
Canonical labels refer to decoded Unicode cell coordinates. Prediction is token-granular. Disabled categories are filtered before edits; overlapping proposals and linked structural repairs require abstention/group handling. Suggestions are not automatically applied. Missing values need not have recoverable intended contents. Real-source unlabelled cells are excluded from negative supervision.
Dataset revision: `{REVISION}`. Train/validation formats are held out; test data was not consumed. This dataset is mostly synthetic. The detector/correction task weight is `ln({K})/ln(256000) = {cfg.correction_weight:.6f}`; valid-copy correction examples are sampled with probability 0.30. The run records periodic held-out validation losses for both tasks in `validation_history`, then a final validation report. These are development diagnostics, not deployment precision or end-to-end repair accuracy. Conditional pseudo-perplexity is scored with each response token masked in turn, original visible; an acceptance threshold is fitted on the reported validation sample. Its sample size and observed errors limit its reliability. No latency guarantee is claimed. See `metrics.json` for sample sizes, loss, operating-point diagnostics and exact-repair outcomes.
## Validation results
Selected step: {report.get('selected_step', step)}. Stop reason: {report.get('stop_reason', 'training in progress')}.
```json
{json.dumps({key: value for key, value in mapping(report.get('validation', {})).items() if key != 'prediction_sample'}, ensure_ascii=False, indent=2)}
```
## Licensing
TabFix fine-tuning contributions are CC BY-NC 4.0; commercial use requires the owner's permission. The underlying mmBERT material remains under its MIT terms and attribution, attributed in BASE_MODEL_CARD.md for pretrained checkpoints. Dataset source terms and USDA attribution remain documented in the dataset card.
## Training controls
The learning rate warms up for {cfg.warmup_steps} updates, then follows cosine decay to 10% of its peak over a fixed {cfg.schedule_steps}-update horizon. A fixed, stratified validation panel of up to {cfg.validation_examples} records is evaluated before training and every {cfg.validation_steps} updates or {cfg.validation_seconds} seconds. Selection minimizes detection loss plus {cfg.correction_weight:.6f} times correction loss (equal averages of fully and partially masked views). Training stops after {cfg.patience} checks without an improvement of at least {cfg.min_delta}. Final publication restores the best complete checkpoint; recoverable latest and best revisions are retained as tags in this repository. These controls reduce risk; they do not guarantee generalization. The test split is not used for checkpoint selection.
## Recorded run results
```json
{json.dumps(report, ensure_ascii=False, indent=2)}
```
'''
@torch.no_grad()
def generate_repair(model: TabFix, consumer: Consumer, p: Prepared, edit: Edit, device: torch.device) -> str | None:
model.eval()
for capacity in (None, 256):
ids, positions, _, _, _ = consumer.correction(p, edit, random.Random(0), False, capacity)
end = integer(consumer.tokenizer.convert_tokens_to_ids(CONTROL[1]))
pad = integer(consumer.tokenizer.convert_tokens_to_ids(CONTROL[2]))
mask = integer(consumer.tokenizer.mask_token_id)
limit = len(positions)
finish: int | None = None
# PAD hypotheses before the selected END may need one additional pass.
for _ in range(2 * limit):
active = [j for j, pos in enumerate(positions) if ids[pos] == mask and (finish is None or j < finish)]
if not active:
break
with autocast(device):
logits = model.vocab(model.hidden(ids, device)[[positions[j] for j in active]]).float()
for special in consumer.tokenizer.all_special_ids:
if special not in (end, pad):
logits[:, special] = -torch.inf
if finish is not None:
logits[:, end] = logits[:, pad] = -torch.inf
probs = logits.softmax(-1)
# Keep original probabilities when restricting boundary candidates.
candidates = probs.clone()
for i, j in enumerate(active):
if finish is None and j == 0:
candidates[i, pad] = -1
if j == limit - 1:
candidates[i] = -1
candidates[i, end] = probs[i, end]
candidates[i, pad] = probs[i, pad]
confidence, prediction = candidates.max(-1)
confidence = torch.nan_to_num(confidence, nan=-1.0)
pick = int(confidence.argmax())
if float(confidence[pick]) < 0:
break
j, token = active[pick], int(prediction[pick])
ids[positions[j]] = token
if token == end:
finish = j
# END defines the replacement length, including immediate deletion.
# Earlier hypotheses to its right must not prevent termination.
for k, pos in enumerate(positions):
if k > j:
ids[pos] = pad
elif k < j and ids[pos] == pad:
ids[pos] = mask
if finish is not None and all(ids[positions[j]] not in (mask, pad, end) for j in range(finish)):
return consumer.literal_tokenizer.decode([ids[positions[j]] for j in range(finish)], skip_special_tokens=False)
return None
@torch.no_grad()
def conditional_perplexity(model: TabFix, consumer: Consumer, p: Prepared, edit: Edit,
answer: str, device: torch.device) -> float:
tokens = consumer.literal(answer)
end = integer(consumer.tokenizer.convert_tokens_to_ids(CONTROL[1]))
pad = integer(consumer.tokenizer.convert_tokens_to_ids(CONTROL[2]))
mask = integer(consumer.tokenizer.mask_token_id)
ids, positions, _, _, _ = consumer.correction(p, edit, random.Random(0), False,
max(8, len(tokens)))
target = tokens + [end]
for j, pos in enumerate(positions):
ids[pos] = target[j] if j < len(target) else pad
losses: list[float] = []
model.eval()
for start in range(0, len(target), 2):
indices = list(range(start, min(start + 2, len(target))))
views = [ids.copy() for _ in indices]
for view, j in zip(views, indices):
view[positions[j]] = mask
with autocast(device):
hidden = model.hidden_batch(views, device)
selected = torch.stack([hidden[k, positions[j]] for k, j in enumerate(indices)])
logits = model.vocab(selected).float()
losses.extend(cast(list[float], F.cross_entropy(logits, torch.tensor([target[j] for j in indices], device=device), reduction="none").cpu().tolist())) # pyright: ignore[reportUnknownMemberType] — torch typing boundary
return math.exp(min(80., sum(losses) / len(losses)))
def calibrate_perplexity(repairs: list[dict[str, object]]) -> dict[str, object]:
scored = sorted((number(r["perplexity"]), bool(r["exact"])) for r in repairs if r.get("perplexity") is not None)
threshold = 0.0
accepted = 0
# Strict empirical policy: accept no observed incorrect validation repair.
# Tied scores are considered together; tiny samples are explicitly reported.
for score in sorted({score for score, _ in scored}):
group = [(s, ok) for s, ok in scored if s <= score]
if not all(ok for _, ok in group):
break
threshold, accepted = score, len(group)
return {"threshold": threshold, "scored": len(scored), "accepted": accepted,
"observed_errors_at_threshold": 0, "policy": "largest validation prefix with zero observed exact-match errors; not a statistical guarantee"}
@torch.no_grad()
def predict_edits(model: TabFix, consumer: Consumer, p: Prepared, device: torch.device,
enabled: set[str], threshold: float = .9) -> list[Edit]:
"""Inference scans every cell; does not use metadata scope or gold errors."""
model.eval()
with autocast(device):
h = model.hidden(p.ids, device)
probs = model.bio.forward(h).reshape(-1, K, 3).float().softmax(-1)
insert = model.insert.forward(h).float().sigmoid()
proposed: list[Edit] = []
blocked = invalid_cells(p.record.xml)
columns = ET.fromstring(p.record.xml).findall("./schema/column")
for c in p.cells:
if (c.row, c.column) in blocked:
continue
if c.column < len(columns):
rules = mapping(xml_json(columns[c.column], "constraints", {}))
if rules.get("enum") or rules.get("lexicon"):
continue
tok = tokens_for(c, p.offsets)
for category, name in enumerate(CATEGORIES):
if name not in enabled:
continue
active: list[int] = []
groups: list[list[int]] = []
for t in tok:
error = float(probs[t, category, 1:].sum()) >= threshold
begin = int(probs[t, category].argmax()) == 1
if active and (not error or begin):
groups.append(active)
active = []
if error:
active.append(t)
if active:
groups.append(active)
for group in groups:
chars = [j for j, (a, b) in enumerate(c.bounds) if b > p.offsets[group[0]][0] and a < p.offsets[group[-1]][1]]
proposed.append(Edit(c.row, c.column, min(chars), max(chars) + 1, "", name))
boundaries = {0, len(c.value)}
boundaries.update(j for j, (a, _) in enumerate(c.bounds) if any(p.offsets[t][0] == a for t in tok))
for at in boundaries:
if float(insert[anchor_token(c, p.offsets, at), category]) >= threshold:
proposed.append(Edit(c.row, c.column, at, at, "", name))
return proposed
def validation_sample(records: list[Record], size: int, seed: int) -> list[Record]:
rng = random.Random(seed)
shuffled = records.copy()
rng.shuffle(shuffled)
buckets: dict[tuple[str, bool], list[Record]] = {}
for r in shuffled:
lang = string(sequence(r.metadata["languages"])[0])
buckets.setdefault((lang, bool(r.edits)), []).append(r)
selected: list[Record] = []
while len(selected) < size and any(buckets.values()):
for key in sorted(buckets):
if buckets[key] and len(selected) < size:
selected.append(buckets[key].pop())
return selected
def monitoring_sample(records: list[Record], size: int, seed: int) -> list[Record]:
shuffled = records.copy()
random.Random(seed).shuffle(shuffled)
buckets: dict[tuple[str, str], list[Record]] = {}
for r in shuffled:
language = string(sequence(r.metadata["languages"])[0])
category = "+".join(sorted({e.category for e in r.edits})) or "clean:" + str(mapping(r.metadata.get("negative", {})).get("subtype", "unspecified"))
buckets.setdefault((language, category), []).append(r)
result: list[Record] = []
while len(result) < size and any(buckets.values()):
for key in sorted(buckets):
if buckets[key] and len(result) < size:
result.append(buckets[key].pop())
return result
@torch.no_grad()
def monitor(model: TabFix, consumer: Consumer, records: list[Record], cfg: RunConfig,
device: torch.device) -> dict[str, object]:
"""Fixed held-out examples and masks; no generation or training RNG consumption."""
model.eval()
det: list[float] = []
full: list[float] = []
partial: list[float] = []
category_counts: Counter[str] = Counter()
confusion = [0, 0, 0]
residual = [r for r in records if r.metadata.get("neural_view")]
general = [r for r in records if not r.metadata.get("neural_view")]
monitored = (monitoring_sample(residual, max(1, cfg.validation_examples // 2), cfg.seed)
+ monitoring_sample(general, max(1, cfg.validation_examples // 2), cfg.seed)) if residual and general else monitoring_sample(records, cfg.validation_examples, cfg.seed)
for r in monitored:
p = consumer.prepare(r)
with autocast(device):
h = model.hidden(p.ids, device)
det.append(float(detection_from_hidden(model, p, h)))
bio = model.bio.forward(h[p.positions]).reshape(-1, K, 3).float().softmax(-1)[:, :, 1:].sum(-1)
gap = model.insert.forward(h[p.gaps]).float().sigmoid()
for probs, raw_labels in ((bio, p.labels), (gap, p.gap_labels)):
labels = torch.tensor(raw_labels, device=device).reshape(-1, K)
valid, actual, predicted = labels != -100, labels > 0, probs >= .9
confusion[0] += int((predicted & actual & valid).sum())
confusion[1] += int((predicted & ~actual & valid).sum())
confusion[2] += int((~predicted & actual & valid).sum())
for e in eligible(r)[:2]:
category_counts[e.category] += 1
for all_masked, output in ((True, full), (False, partial)):
view = consumer.correction(p, e, random.Random(r.id + str(e)), True, mask_all=all_masked)
with autocast(device):
output.append(float(correction_loss(model, view, device)))
if not det or not full:
raise ValueError("Monitoring requires nonempty detection and correction validation sets")
d, c = sum(det) / len(det), .5 * (sum(full) / len(full) + sum(partial) / len(partial))
return {"detection_loss": d, "correction_loss": c,
"correction_full_mask_loss": sum(full) / len(full),
"correction_partial_mask_loss": sum(partial) / len(partial),
"selection_score": d + cfg.correction_weight * c,
"examples": len(det), "correction_edits": len(full),
"correction_category_counts": dict(category_counts),
"token_and_gap_tp_fp_fn_at_0.9": confusion,
"sample_ids": [r.id for r in monitored]}
@torch.no_grad()
def evaluate(model: TabFix, consumer: Consumer, records: list[Record], cfg: RunConfig, device: torch.device) -> dict[str, object]:
model.eval()
losses: list[float] = []
correction_losses: list[float] = []
counts = {str(t): [0, 0, 0, 0] for t in (.5, .8, .9, .95, .99)}
repairs: list[dict[str, object]] = []
predictions: list[dict[str, object]] = []
for r in validation_sample(records, cfg.validation_examples, cfg.seed):
p = consumer.prepare(r)
with autocast(device):
losses.append(float(detection_loss(model, p, device)))
h = model.hidden(p.ids, device)
prob = model.bio.forward(h[p.positions]).reshape(-1, K, 3).float().softmax(-1)[:, :, 1:].sum(-1)
labels = torch.tensor(p.labels, device=device).reshape(-1, K)
valid = labels != -100
actual = labels > 0
for threshold, v in counts.items():
predicted = prob >= float(threshold)
v[0] += int((predicted & actual & valid).sum())
v[1] += int((predicted & ~actual & valid).sum())
v[2] += int((~predicted & actual & valid).sum())
v[3] += int((~predicted & ~actual & valid).sum())
predictions.append({"id": r.id, "predictions_at_0.9": [asdict(e) for e in predict_edits(model, consumer, p, device, set(CATEGORIES))][:50]})
eligible_edits = eligible(r)
if eligible_edits:
try:
correction_view = consumer.correction(p, eligible_edits[0], random.Random(0), True)
correction_losses.append(float(correction_loss(model, correction_view, device)))
if len(repairs) < cfg.validation_repairs:
e = eligible_edits[0]
_, _, _, _, target = consumer.correction(p, e, random.Random(0), False)
answer = generate_repair(model, consumer, p, e, device)
score = conditional_perplexity(model, consumer, p, e, answer, device) if answer is not None else math.inf
repairs.append({"id": r.id, "category": e.category, "target": target, "prediction": answer, "exact": answer == target, "perplexity": score if math.isfinite(score) else None})
except ValueError as exc:
consumer.stats[str(exc)] += 1
return {"sample_policy": "seeded language x positive/clean round-robin validation sample",
"validation_examples": len(losses), "detection_loss": sum(losses) / max(1, len(losses)),
"correction_examples": len(correction_losses), "correction_loss": sum(correction_losses) / max(1, len(correction_losses)),
"token_category_confusion_tp_fp_fn_tn": counts,
"gold_cell_repairs": repairs, "perplexity_calibration": calibrate_perplexity(repairs), "prediction_sample": predictions,
"limitations": "Token diagnostics, not cell/span calibration; limited gold-span repair sample; test untouched."}
def main(cfg: RunConfig) -> None:
torch.set_num_threads(min(4, os.cpu_count() or 1))
torch.manual_seed(cfg.seed) # pyright: ignore[reportUnknownMemberType] — third-party typing boundary
rng = random.Random(cfg.seed)
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
if device.type == "cuda" and not torch.cuda.is_bf16_supported():
raise RuntimeError("Use a BF16-capable GPU (L4/L40S/A100); do not silently change precision")
resume_path: Path | None = None
if cfg.resume:
resume_path = Path(cfg.resume) if Path(cfg.resume).is_dir() else Path(snapshot_download(cfg.resume, revision=cfg.resume_revision))
model, tok, manifest = load_checkpoint(resume_path, device)
if cfg.verify_only:
emit("verified", step=manifest["step"], parameters=sum(p.numel() for p in model.parameters()))
return
old = read_json(resume_path / "run_config.json")
for name in ("seed", "lr", "max_length", "correction_weight", "batch_size", "tiny", "warmup_steps", "schedule_steps", "validation_examples", "patience", "min_delta"):
if old.get(name) != asdict(cfg)[name]:
raise ValueError(f"Resume config mismatch: {name}")
else:
if cfg.verify_only:
raise ValueError("--verify-only requires --resume")
model, tok = make_model(cfg)
model.to(device)
if device.type == "cuda":
model.mlm.model.gradient_checkpointing_enable() # pyright: ignore[reportUnknownMemberType] — third-party typing boundary
consumer = Consumer(tok, cfg.max_length)
optimizer = torch.optim.AdamW(model.parameters(), lr=cfg.lr, weight_decay=.01)
step = 0
if resume_path:
state = mapping(cast(object, torch.load(resume_path / "training_state.pt", map_location="cpu", weights_only=False)))
optimizer.load_state_dict(cast(dict[str, object], state["optimizer"]))
step = integer(state["step"])
rng.setstate(cast(tuple[object, ...], state["python_rng"]))
torch.set_rng_state(cast(Tensor, state["torch_rng"]))
if device.type == "cuda":
torch.cuda.set_rng_state_all(cast(list[Tensor], state["cuda_rng"]))
train = read_records("train", cfg)
validation = read_records("validation", cfg)
weights = [number(r.metadata["sampling_weight"]) for r in train]
repair_indices = [i for i, r in enumerate(train) if eligible(r)]
repair_only_indices = [i for i in repair_indices if any(e.category != "copy" for e in eligible(train[i]))]
copy_indices = [i for i in repair_indices if any(e.category == "copy" for e in eligible(train[i]))]
neural_indices = [i for i, r in enumerate(train) if r.metadata.get("neural_view")]
if not neural_indices:
neural_indices = list(range(len(train))) # Tiny local smoke fixtures.
neural_weights = [weights[i] for i in neural_indices]
if cfg.batch_size < 1:
raise ValueError("Batch size must be positive")
api = HfApi() if cfg.repo_id else None
if api:
# A fresh objective intentionally replaces the existing preview repo;
# the upload commits remain atomic and the local checkpoint is recoverable.
_ = api.repo_info(cfg.repo_id)
@lru_cache(maxsize=64)
def prepared(i: int) -> Prepared:
return consumer.prepare(train[i])
validation_history: list[dict[str, object]] = []
if resume_path and (resume_path / "metrics.json").exists():
prior_report = read_json(resume_path / "metrics.json")
validation_history.extend(entry for item in sequence(prior_report.get("validation_history", [])) if integer((entry := mapping(item))["step"]) <= step)
report: dict[str, object] = {"tiny_smoke": cfg.tiny, "dataset_revision": REVISION,
"train_records": len(train), "correction_eligible_records": len(repair_indices),
"correction_request_version": 4, "batch_size": cfg.batch_size,
"start_step": step, "device": str(device), "parameters": sum(p.numel() for p in model.parameters()),
"validation_history": validation_history}
emit("ready", **report)
start = time.monotonic()
saved_at = start
last_validation_at = start
deadline = start + cfg.train_seconds
selection = Selection()
best_path = Path(cfg.output) / "best"
if resume_path:
state_selection = mapping(read_json(resume_path / "metrics.json")["selection"])
selection = Selection(number(state_selection["best_score"]), integer(state_selection["best_step"]), integer(state_selection["bad_checks"]))
if not best_path.exists():
if cfg.repo_id:
source_best = Path(snapshot_download(cfg.repo_id, revision=f"v2-best-step-{selection.best_step}"))
elif selection.best_step == step:
source_best = resume_path
else:
raise ValueError("Resume needs the saved best checkpoint as well as latest state")
shutil.copytree(source_best, best_path)
if cfg.validation_steps < 1 or cfg.patience < 1:
raise ValueError("Validation interval and patience must be positive")
def check_validation() -> None:
nonlocal last_validation_at, saved_at
optimizer.zero_grad(set_to_none=True)
result = monitor(model, consumer, validation, cfg, device)
improved = selection.observe(number(result["selection_score"]), step, cfg.min_delta)
entry: dict[str, object] = {"step": step, **result}
validation_history.append(entry)
report["selection"] = asdict(selection)
emit("validation", **{k: v for k, v in entry.items() if k != "sample_ids"}, improved=improved, bad_checks=selection.bad_checks)
current = save_checkpoint(model, tok, optimizer, cfg, step, rng, report, api)
if improved:
shutil.copytree(current, best_path, dirs_exist_ok=True)
if api:
uploaded = read_json(Path(cfg.output) / "last_upload.json")
api.create_tag(repo_id=cfg.repo_id, tag=f"v2-best-step-{step}", revision=string(uploaded["commit"]), exist_ok=True)
last_validation_at = saved_at = time.monotonic()
if not resume_path:
check_validation()
iterations = 0
interval_detection: list[float] = []
interval_correction: list[float] = []
while step < cfg.steps and time.monotonic() < deadline and not stop_requested:
model.train()
optimizer.zero_grad(set_to_none=True)
for group in optimizer.param_groups:
group["lr"] = learning_rate(cfg, step)
batch = [prepared(rng.choices(neural_indices, weights=neural_weights, k=1)[0]) for _ in range(cfg.batch_size)]
with autocast(device):
hidden = model.hidden_batch([p.ids for p in batch], device)
det = torch.stack([detection_from_hidden(model, p, hidden[i]) for i, p in enumerate(batch)]).mean()
if not bool(torch.isfinite(det)):
raise FloatingPointError("Nonfinite detection loss")
det.backward() # pyright: ignore[reportUnknownMemberType] — third-party typing boundary
views: list[tuple[list[int], list[int], list[int], list[bool], str]] = []
for _ in range(cfg.batch_size * 20):
copy_task = bool(copy_indices) and (not repair_only_indices or rng.random() < .3)
pool = copy_indices if copy_task else repair_only_indices
rp = prepared(rng.choice(pool))
choices = [e for e in eligible(rp.record) if (e.category == "copy") == copy_task]
try:
views.append(consumer.correction(rp, rng.choice(choices), rng, True))
except ValueError as exc:
consumer.stats[str(exc)] += 1
if len(views) == cfg.batch_size:
break
if len(views) != cfg.batch_size:
raise RuntimeError("Cannot obtain eligible correction batch")
with autocast(device):
hidden = model.hidden_batch([v[0] for v in views], device)
unweighted_cor = torch.stack([correction_from_hidden(model, v, hidden[i]) for i, v in enumerate(views)]).mean()
cor = unweighted_cor * cfg.correction_weight
if not bool(torch.isfinite(cor)):
raise FloatingPointError("Nonfinite correction loss")
cor.backward() # pyright: ignore[reportUnknownMemberType] — third-party typing boundary
cor_value = float(unweighted_cor.detach())
_ = torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0, error_if_nonfinite=True)
optimizer.step() # pyright: ignore[reportUnknownMemberType] — third-party typing boundary
step += 1
iterations += 1
if device.type == "cuda":
report["peak_gpu_bytes"] = torch.cuda.max_memory_allocated()
report.update(step=step, learning_rate=learning_rate(cfg, step - 1), last_detection_loss=float(det.detach()), last_correction_loss=cor_value,
training_elapsed_seconds=time.monotonic() - start, exclusions=dict(consumer.stats))
interval_detection.append(float(det.detach()))
interval_correction.append(cor_value)
if step % 10 == 0 or iterations <= 2:
report["mean_detection_loss"] = sum(interval_detection) / len(interval_detection)
report["mean_correction_loss"] = sum(interval_correction) / len(interval_correction)
report["loss_interval_updates"] = len(interval_detection)
emit("train", **{k: v for k, v in report.items() if k != "validation_history"})
interval_detection.clear()
interval_correction.clear()
if iterations == 10 or time.monotonic() - saved_at >= cfg.save_seconds:
_ = save_checkpoint(model, tok, optimizer, cfg, step, rng, report, api)
saved_at = time.monotonic()
if step % cfg.validation_steps == 0 or (cfg.validation_seconds > 0 and time.monotonic() - last_validation_at >= cfg.validation_seconds):
check_validation()
if selection.bad_checks >= cfg.patience:
break
report["stop_reason"] = "signal" if stop_requested else "early_stopping" if selection.bad_checks >= cfg.patience else "steps" if step >= cfg.steps else "time_budget"
# Save BEFORE evaluation so a timeout cannot erase trained progress.
current = save_checkpoint(model, tok, optimizer, cfg, step, rng, report, api)
if not stop_requested:
if integer(validation_history[-1]["step"]) != step:
check_validation()
if api:
last_upload = read_json(Path(cfg.output) / "last_upload.json")
api.create_tag(repo_id=cfg.repo_id, tag=f"v2-latest-step-{step}", revision=string(last_upload["commit"]), exist_ok=True)
report["last_training_step"] = step
report["selected_step"] = selection.best_step
report["selection_at_stop"] = asdict(selection)
report["selection"] = read_json(best_path / "metrics.json")["selection"]
del model, optimizer
model, tok, _ = load_checkpoint(best_path, device)
# Restore the entire matching checkpoint, including optimizer and RNG state.
shutil.copytree(best_path, current, dirs_exist_ok=True)
step = selection.best_step
report["step"] = step
# Evaluation is deterministic and does not advance the training RNG.
report["validation"] = evaluate(model, consumer, validation, cfg, device)
write_json(current / "metrics.json", report)
(current / "README.md").write_text(model_card(cfg, step, report))
manifest = read_json(current / "checkpoint.json")
files = mapping(manifest["files"])
for name in ("metrics.json", "README.md"):
files[name] = hashlib.sha256((current / name).read_bytes()).hexdigest()
write_json(current / "checkpoint.json", manifest)
if api:
_ = api.upload_folder(repo_id=cfg.repo_id, folder_path=current, commit_message="Validation diagnostics and finalized model card")
emit("complete", step=step, stop_reason=report["stop_reason"], metrics_path=str(current / "metrics.json"))
def parse_args() -> RunConfig:
parser = argparse.ArgumentParser(description=__doc__)
defaults = RunConfig()
for name, value in mapping(asdict(defaults)).items():
flag = "--" + name.replace("_", "-")
if isinstance(value, bool):
_ = parser.add_argument(flag, action="store_true", default=value)
else:
_ = parser.add_argument(flag, type=int if isinstance(value, int) else float if isinstance(value, float) else str, default=value)
args = mapping(vars(parser.parse_args()))
return RunConfig(
repo_id=string(args["repo_id"]), output=string(args["output"]), local_data=string(args["local_data"]),
tokenizer=string(args["tokenizer"]), resume=string(args["resume"]), resume_revision=string(args["resume_revision"]),
tiny=bool(args["tiny"]), steps=integer(args["steps"]), train_seconds=integer(args["train_seconds"]),
save_seconds=integer(args["save_seconds"]), lr=number(args["lr"]), seed=integer(args["seed"]),
validation_examples=integer(args["validation_examples"]), validation_repairs=integer(args["validation_repairs"]), validation_seconds=integer(args["validation_seconds"]),
validation_steps=integer(args["validation_steps"]), patience=integer(args["patience"]), min_delta=number(args["min_delta"]), warmup_steps=integer(args["warmup_steps"]), schedule_steps=integer(args["schedule_steps"]),
max_length=integer(args["max_length"]), correction_weight=number(args["correction_weight"]), batch_size=integer(args["batch_size"]), verify_only=bool(args["verify_only"]))
def handle_signal(_signum: int, _frame: object) -> None:
global stop_requested
stop_requested = True
if __name__ == "__main__":
signal.signal(signal.SIGTERM, handle_signal)
signal.signal(signal.SIGINT, handle_signal)
main(parse_args())