spreadsheetlabeller / label_utils.py
zofiasmolenasana
Merge value+attribute labels, add header label, RAG eval improvements
d051588 unverified
Raw
History Blame Contribute Delete
4.3 kB
"""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