Download train_tabfix.py from Antix5/tabfix-preview: direct link, hf CLI and curl.
- Browser
- Download file 68.9 kB
-
https://huggingface.co/Antix5/tabfix-preview/resolve/main/train_tabfix.py
- Command line
-
hf download hf://Antix5/tabfix-preview/train_tabfix.py
-
curl -L -o train_tabfix.py https://huggingface.co/Antix5/tabfix-preview/resolve/main/train_tabfix.py
68.9 kB
| # /// 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) | |
| 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 "<empty/>" | |
| 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} | |
| class Edit: | |
| row: int | |
| column: int | |
| start: int | |
| end: int | |
| replacement: str | |
| category: str | |
| 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"])) | |
| class Record: | |
| id: str | |
| xml: str | |
| edits: list[Edit] | |
| metadata: dict[str, object] | |
| 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"])))) ) | |
| 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'<row index="(\d+)">(.*?)</row>', xml, re.S): | |
| for col, cell in enumerate(re.finditer(r'<cell\b[^>]*>(.*?)</cell>', row[2], re.S)): | |
| raw = cell[1] | |
| origin = row.start(2) + cell.start(1) | |
| bounds: list[tuple[int, int]] = [] | |
| chars: list[str] = [] | |
| if raw != "<empty/>": | |
| 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 | |
| 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>" + original + "</original><replacement>" + string(self.tokenizer.mask_token) * (budget + 1) + "</replacement>" | |
| 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 | |
| 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) | |
| 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 `<replacement>` beside the visible `<original>` 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)} | |
| ``` | |
| ''' | |
| 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 | |
| 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"} | |
| 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 | |
| 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]} | |
| 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) | |
| 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()) | |