# /// 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())