"""Shared label constants and helpers used across the labeling pipeline.""" from __future__ import annotations # --------------------------------------------------------------------------- # Fine-grained labels (13 classes) — used by new models # --------------------------------------------------------------------------- CELL_LABELS = [ "value", "aggregation", "header", "metadata", "comment", "empty", "junk", "col_header_1", "col_header_2", "col_header_3", "row_header_1", "row_header_2", "row_header_3", ] LABEL2IDX = {label: i for i, label in enumerate(CELL_LABELS)} NUM_CLASSES = len(CELL_LABELS) MAX_HEADER_LEVEL = 3 # --------------------------------------------------------------------------- # Coarse labels (9 classes) — legacy models and backward-compatible evaluation # --------------------------------------------------------------------------- COARSE_LABELS = [ "value", "col_header", "row_header", "aggregation", "header", "metadata", "comment", "empty", "junk", ] COARSE_LABEL2IDX = {label: i for i, label in enumerate(COARSE_LABELS)} NUM_COARSE_CLASSES = len(COARSE_LABELS) # Map each fine-grained index to the corresponding coarse index fine_to_coarse_idx: dict[int, int] = {} for _fine_idx, _fine_label in enumerate(CELL_LABELS): _coarse_label = ( "col_header" if _fine_label.startswith("col_header_") else "row_header" if _fine_label.startswith("row_header_") else _fine_label ) fine_to_coarse_idx[_fine_idx] = COARSE_LABEL2IDX[_coarse_label] # --------------------------------------------------------------------------- # Convenience sets # --------------------------------------------------------------------------- VALUE_LABELS = {"value", "aggregation"} DATA_LABELS = {"value", "aggregation"} KNOWN_FLAT_LABELS = {"value", "aggregation", "header", "metadata", "comment", "empty", "junk"} # --------------------------------------------------------------------------- # Predicates # --------------------------------------------------------------------------- def is_row_header(label: str | None) -> bool: return label is not None and label.startswith("row_header_") def is_col_header(label: str | None) -> bool: return label is not None and label.startswith("col_header_") def header_level(label: str) -> int: return int(label.rsplit("_", 1)[1]) def clamp_header_level(label: str, max_level: int = MAX_HEADER_LEVEL) -> str: """Clamp header level to *max_level*. e.g. col_header_8 -> col_header_3.""" if is_col_header(label): lvl = min(header_level(label), max_level) return f"col_header_{lvl}" if is_row_header(label): lvl = min(header_level(label), max_level) return f"row_header_{lvl}" return label def is_header(label: str | None) -> bool: return label == "header" def normalize_label(label: str | None) -> str: """Collapse row_header_N -> row_header, col_header_N -> col_header. Also maps legacy/orphan labels to their current equivalents. """ if label is None: return "unlabeled" if label == "column_group": return "metadata" if label == "attribute": return "value" if is_row_header(label): return "row_header" if is_col_header(label): return "col_header" return label def is_known_label(label: str | None) -> bool: """Check if label is a recognized cell label (handles any header level).""" if label is None: return False clamped = clamp_header_level(label) if clamped in LABEL2IDX: return True normalized = normalize_label(label) return normalized in LABEL2IDX or normalized in COARSE_LABEL2IDX def is_fine_grained_label(label: str | None) -> bool: """Return True if *label* is one of the 13 fine-grained class names.""" return label is not None and label in LABEL2IDX def is_comment(label: str | None) -> bool: return label == "comment" def is_junk(label: str | None) -> bool: return label == "junk" def is_data_label(label: str | None) -> bool: """Return True for any label that represents labeled content (not unlabeled).""" if label in KNOWN_FLAT_LABELS: return True if is_col_header(label) or is_row_header(label): return True return False