Spaces:
Running
Running
| """Run trained GNN inference on unlabeled spreadsheet grids. | |
| Loads versioned models from data/models/ and predicts cell labels | |
| for a raw grid (as returned by xlsx_client.fetch_xlsx_sheet). | |
| Supports batch prediction of all unlabeled sheets after training, | |
| storing results per model version in data/predictions/v{NNN}/. | |
| A connected-component heuristic assigns table IDs based on spatial gaps. | |
| """ | |
| from __future__ import annotations | |
| import json | |
| import logging | |
| from collections import deque | |
| from datetime import datetime, timezone | |
| from pathlib import Path | |
| from typing import Optional | |
| import torch | |
| import torch.nn.functional as F | |
| import numpy as np | |
| import config | |
| from features import _content_features, _formula_features, FTYPE_TO_IDX | |
| from models import ( | |
| GATModel, GCNModel, GraphTransformerModel, MLPModel, | |
| SequenceTransformerModel, AdjTransformerModel, | |
| SpatialEdgeTransformerModel, DualModalityGNNModel, | |
| LAYOUTLM_ARCHS, | |
| build_model, | |
| ) | |
| from label_utils import CELL_LABELS, COARSE_LABELS | |
| from train_gnn import ( | |
| EMBED_DIM, EDGE_TYPE_DIM, GraphConfig, | |
| _build_spatial_edges_cfg, _build_row_col_edges_cfg, | |
| _build_virtual_col_nodes, _build_virtual_row_nodes, | |
| _build_content_edge_features, | |
| EDGE_TYPE_SPATIAL, | |
| ) | |
| import math | |
| ARCH_REGISTRY: dict[str, type] = { | |
| "gat": GATModel, | |
| "gcn": GCNModel, | |
| "graph_transformer": GraphTransformerModel, | |
| "mlp": MLPModel, | |
| "seq_transformer": SequenceTransformerModel, | |
| "adj_transformer": AdjTransformerModel, | |
| "spatial_edge_transformer": SpatialEdgeTransformerModel, | |
| "dual_modality_gnn": DualModalityGNNModel, | |
| } | |
| SpreadsheetGNN = GATModel | |
| logger = logging.getLogger("labeler.predict") | |
| _cached_model: Optional[SpreadsheetGNN] = None | |
| _cached_config: Optional[dict] = None | |
| _cached_version: Optional[str] = None | |
| _cached_output_labels: Optional[list[str]] = None | |
| _cached_embed_model = None | |
| def get_latest_model_version() -> Optional[str]: | |
| """Scan data/models/ for the highest version number. Returns e.g. '003' or None.""" | |
| existing = sorted(config.MODELS_DIR.glob("model_v*.pt")) | |
| if not existing: | |
| return None | |
| return existing[-1].stem.split("_v")[1] | |
| def get_best_model_path() -> Optional[Path]: | |
| """Return best_model.pt if it exists, else the latest versioned model.""" | |
| best = config.MODELS_DIR / "best_model.pt" | |
| if best.exists(): | |
| return best | |
| existing = sorted(config.MODELS_DIR.glob("model_v*.pt")) | |
| return existing[-1] if existing else None | |
| def _checkpoint_kwargs_from_config(cfg: dict) -> dict: | |
| """Extra constructor kwargs from experiment checkpoints (e.g. CNN grid).""" | |
| from train_gnn import CNN_GRID_CHANNELS, CNN_GRID_OUT_DIM | |
| extra: dict = {} | |
| if cfg.get("use_cnn_grid"): | |
| extra["use_cnn_grid"] = True | |
| extra["cnn_grid_channels"] = int(cfg.get("cnn_grid_channels", CNN_GRID_CHANNELS)) | |
| extra["cnn_grid_out_dim"] = int(cfg.get("cnn_grid_out_dim", CNN_GRID_OUT_DIM)) | |
| return extra | |
| def _align_structural_embed_to_dim( | |
| x_structural: torch.Tensor, | |
| x_embed: torch.Tensor, | |
| expected_dim: int, | |
| ) -> torch.Tensor: | |
| """Resize structural features so ``[structural | embed]`` has width ``expected_dim``. | |
| Truncates or zero-pads the **structural** block (prefix) when the live feature | |
| pipeline drifts from the checkpoint (e.g. 89 vs 71 structural dims). | |
| """ | |
| from train_gnn import EMBED_DIM | |
| emb_d = x_embed.size(1) | |
| need_struct = expected_dim - emb_d | |
| if need_struct < 0: | |
| logger.warning( | |
| "Checkpoint expects in_dim=%d smaller than embedding width=%d; " | |
| "truncating concatenated features.", | |
| expected_dim, emb_d, | |
| ) | |
| x_full = torch.cat([x_structural, x_embed], dim=1) | |
| return x_full[:, :expected_dim] | |
| s = x_structural.size(1) | |
| if s == need_struct: | |
| return torch.cat([x_structural, x_embed], dim=1) | |
| if s > need_struct: | |
| logger.warning( | |
| "Aligned structural features %d -> %d to match checkpoint in_dim=%d " | |
| "(older / fewer structural dims in trained weights).", | |
| s, need_struct, expected_dim, | |
| ) | |
| x_structural = x_structural[:, :need_struct].contiguous() | |
| else: | |
| pad = torch.zeros( | |
| x_structural.size(0), need_struct - s, | |
| dtype=x_structural.dtype, device=x_structural.device, | |
| ) | |
| logger.warning( | |
| "Aligned structural features %d -> %d to match checkpoint in_dim=%d " | |
| "(padding missing dims with zeros).", | |
| s, need_struct, expected_dim, | |
| ) | |
| x_structural = torch.cat([x_structural, pad], dim=1) | |
| return torch.cat([x_structural, x_embed], dim=1) | |
| def load_model_from_path(model_path: Path) -> Optional[SpreadsheetGNN]: | |
| """Load any architecture from a checkpoint path. Not cached.""" | |
| if not model_path.exists(): | |
| return None | |
| try: | |
| checkpoint = torch.load(model_path, map_location="cpu", weights_only=False) | |
| cfg = checkpoint["config"] | |
| arch = checkpoint.get("arch", "gat") | |
| num_classes = cfg.get("num_classes", 9) | |
| saved_etd = cfg.get("edge_type_dim", 3) | |
| ck_kw = _checkpoint_kwargs_from_config(cfg) | |
| model = build_model( | |
| arch, | |
| in_dim=cfg["in_dim"], | |
| hidden_dim=cfg["hidden_dim"], | |
| num_classes=num_classes, | |
| edge_dim=saved_etd, | |
| **ck_kw, | |
| ) | |
| # train_compare enables coarse+level heads for non-LayoutLM; legacy train_gnn GAT does not. | |
| state = checkpoint["model_state_dict"] | |
| use_two_stage = cfg.get("use_two_stage") | |
| if use_two_stage is None: | |
| use_two_stage = any(k.startswith("coarse_head") for k in state.keys()) | |
| if use_two_stage and arch not in LAYOUTLM_ARCHS: | |
| model.enable_two_stage_head(cfg["hidden_dim"]) | |
| model.load_state_dict(state, strict=False) | |
| model.eval() | |
| model._edge_type_dim = saved_etd | |
| model._output_labels = checkpoint.get("cell_labels", COARSE_LABELS) | |
| model._legacy = (num_classes <= 9) | |
| model._arch = arch | |
| model._train_in_dim = cfg.get("in_dim") | |
| gc_dict = cfg.get("graph_config") | |
| model._graph_config = GraphConfig.from_dict(gc_dict) if gc_dict else None | |
| return model | |
| except Exception: | |
| logger.exception("Failed to load model from %s", model_path) | |
| return None | |
| def load_model(version: Optional[str] = None) -> Optional[SpreadsheetGNN]: | |
| """Load a trained model, caching it in memory. Returns None if no model file. | |
| Args: | |
| version: Specific version string (e.g. '001'). If None, loads the best | |
| available model (best_model.pt or latest versioned). | |
| """ | |
| global _cached_model, _cached_config, _cached_version, _cached_output_labels | |
| if version is not None: | |
| model_path = config.MODELS_DIR / f"model_v{version}.pt" | |
| else: | |
| model_path = get_best_model_path() | |
| if model_path is None: | |
| model_path = config.DATA_DIR / "spreadsheet_gnn.pt" | |
| if not model_path.exists(): | |
| return None | |
| version_key = version or model_path.stem | |
| if _cached_model is not None and _cached_version == version_key: | |
| return _cached_model | |
| model = load_model_from_path(model_path) | |
| if model is None: | |
| return None | |
| _cached_model = model | |
| _cached_config = torch.load(model_path, map_location="cpu", weights_only=False)["config"] | |
| _cached_version = version_key | |
| _cached_output_labels = model._output_labels | |
| logger.info( | |
| "Loaded model from %s (arch=%s, version=%s)", | |
| model_path, getattr(model, '_arch', 'unknown'), version_key, | |
| ) | |
| return model | |
| def _cell_features_inference(cell: dict, max_row: int, max_col: int, | |
| median_font_size: float = 11.0, | |
| col_row_stats: dict | None = None, | |
| cohesion_stats: dict | None = None) -> list[float]: | |
| """Delegate to the shared _cell_features in train_gnn (keeps train/inference in sync).""" | |
| from train_gnn import _cell_features | |
| return _cell_features(cell, max_row, max_col, median_font_size, col_row_stats, cohesion_stats) | |
| def _build_flat_spatial_edges(cells: list[dict], idx_map: dict) -> list[list[int]]: | |
| """Build spatial edges connecting all adjacent cells (no table boundaries).""" | |
| coord_to_idx = {} | |
| for c in cells: | |
| key = (c["row"], c["col"]) | |
| if key in idx_map: | |
| coord_to_idx[key] = idx_map[key] | |
| src, dst = [], [] | |
| for (r, c), node_idx in coord_to_idx.items(): | |
| for dr, dc in [(-1, 0), (1, 0), (0, -1), (0, 1)]: | |
| neighbor = (r + dr, c + dc) | |
| if neighbor in coord_to_idx: | |
| src.append(node_idx) | |
| dst.append(coord_to_idx[neighbor]) | |
| return [src, dst] | |
| def _build_inference_edges( | |
| all_cells: list[dict], idx_map: dict, graph_cfg: GraphConfig, | |
| ) -> tuple[torch.Tensor, torch.Tensor, int]: | |
| """Build inference-time edges matching training GraphConfig. | |
| Returns (edge_index, edge_attr, num_virtual_nodes). | |
| Inference never has table labels, so table_groups is empty. | |
| """ | |
| table_groups = {} # no table scoping at inference | |
| num_real = len(all_cells) | |
| spatial_edges, spatial_types = _build_spatial_edges_cfg( | |
| all_cells, idx_map, table_groups, graph_cfg) | |
| rc_edges, rc_types = _build_row_col_edges_cfg( | |
| all_cells, idx_map, table_groups, graph_cfg) | |
| all_src = spatial_edges[0] + rc_edges[0] | |
| all_dst = spatial_edges[1] + rc_edges[1] | |
| all_types = spatial_types + rc_types | |
| num_virtual = 0 | |
| if graph_cfg.virtual_col_nodes: | |
| vc_edges, vc_types, n_vc = _build_virtual_col_nodes( | |
| all_cells, idx_map, num_real + num_virtual) | |
| all_src += vc_edges[0]; all_dst += vc_edges[1]; all_types += vc_types | |
| num_virtual += n_vc | |
| if graph_cfg.virtual_row_nodes: | |
| vr_edges, vr_types, n_vr = _build_virtual_row_nodes( | |
| all_cells, idx_map, num_real + num_virtual) | |
| all_src += vr_edges[0]; all_dst += vr_edges[1]; all_types += vr_types | |
| num_virtual += n_vr | |
| total_nodes = num_real + num_virtual | |
| if not all_src: | |
| return (torch.zeros(2, 0, dtype=torch.long), | |
| torch.zeros(0, graph_cfg.edge_type_dim), | |
| num_virtual) | |
| edge_index = torch.tensor([all_src, all_dst], dtype=torch.long) | |
| net = graph_cfg.n_edge_types | |
| num_edges = len(all_types) | |
| edge_type_onehot = torch.zeros(num_edges, net) | |
| et_arr = torch.tensor(all_types, dtype=torch.long).clamp(max=net - 1) | |
| edge_type_onehot.scatter_(1, et_arr.unsqueeze(1), 1.0) | |
| parts = [edge_type_onehot] | |
| if graph_cfg.use_rel_pe: | |
| node_rows_f = torch.zeros(total_nodes, dtype=torch.float) | |
| node_cols_f = torch.zeros(total_nodes, dtype=torch.float) | |
| for i, c in enumerate(all_cells): | |
| node_rows_f[i] = c["row"] | |
| node_cols_f[i] = c["col"] | |
| src_t = torch.tensor(all_src, dtype=torch.long) | |
| dst_t = torch.tensor(all_dst, dtype=torch.long) | |
| dr = node_rows_f[dst_t] - node_rows_f[src_t] | |
| dc = node_cols_f[dst_t] - node_cols_f[src_t] | |
| REL_PE_DIM = 16 | |
| half = REL_PE_DIM // 4 | |
| div_term = torch.exp( | |
| torch.arange(0, half, dtype=torch.float) * -(math.log(10000.0) / max(half, 1)) | |
| ) | |
| rel_pe = torch.cat([ | |
| torch.sin(dr.unsqueeze(-1) * div_term), | |
| torch.cos(dr.unsqueeze(-1) * div_term), | |
| torch.sin(dc.unsqueeze(-1) * div_term), | |
| torch.cos(dc.unsqueeze(-1) * div_term), | |
| ], dim=1) | |
| parts.append(rel_pe) | |
| if graph_cfg.content_edge_features: | |
| content_feat = _build_content_edge_features(all_cells, all_src, all_dst) | |
| parts.append(content_feat) | |
| if graph_cfg.distance_edge_features: | |
| if not graph_cfg.use_rel_pe: | |
| node_rows_f = torch.zeros(total_nodes, dtype=torch.float) | |
| node_cols_f = torch.zeros(total_nodes, dtype=torch.float) | |
| for i, c in enumerate(all_cells): | |
| node_rows_f[i] = c["row"] | |
| node_cols_f[i] = c["col"] | |
| src_t = torch.tensor(all_src, dtype=torch.long) | |
| dst_t = torch.tensor(all_dst, dtype=torch.long) | |
| dr = node_rows_f[dst_t] - node_rows_f[src_t] | |
| dc = node_cols_f[dst_t] - node_cols_f[src_t] | |
| max_row = max((c["row"] for c in all_cells), default=0) + 1 | |
| max_col = max((c["col"] for c in all_cells), default=0) + 1 | |
| manhattan = dr.abs() + dc.abs() | |
| norm_dist = manhattan / max(max_row + max_col, 1) | |
| dist_decay = torch.exp(-0.3 * manhattan) | |
| dist_feat = torch.stack([norm_dist, dist_decay], dim=1) | |
| parts.append(dist_feat) | |
| edge_attr = torch.cat(parts, dim=1) | |
| return edge_index, edge_attr, num_virtual | |
| def _assign_table_ids(cells: list[dict], predictions: list[dict]) -> None: | |
| """Heuristic: assign table IDs via connected components of non-empty cells. | |
| Cells predicted as metadata that are spatially isolated from data cells | |
| get table 0 (sheet-level). Connected components of data cells each get | |
| a distinct table ID starting from 1. | |
| """ | |
| occupied = set() | |
| coord_to_pred = {} | |
| for p in predictions: | |
| key = (p["row"], p["col"]) | |
| occupied.add(key) | |
| coord_to_pred[key] = p | |
| data_labels = {"value", "row_header_1", "col_header_1", "aggregation"} | |
| data_coords = set() | |
| for p in predictions: | |
| if p["label"] in data_labels or p["label"].startswith("row_header_") or p["label"].startswith("col_header_"): | |
| data_coords.add((p["row"], p["col"])) | |
| visited = set() | |
| components = [] | |
| for coord in sorted(data_coords): | |
| if coord in visited: | |
| continue | |
| component = set() | |
| queue = deque([coord]) | |
| while queue: | |
| cur = queue.popleft() | |
| if cur in visited: | |
| continue | |
| visited.add(cur) | |
| if cur not in occupied: | |
| continue | |
| component.add(cur) | |
| r, c = cur | |
| for dr, dc in [(-1, 0), (1, 0), (0, -1), (0, 1)]: | |
| nb = (r + dr, c + dc) | |
| if nb in occupied and nb not in visited: | |
| queue.append(nb) | |
| if component: | |
| components.append(component) | |
| # Merge vertically adjacent components whose column spans overlap | |
| # significantly. This prevents header rows from getting a different | |
| # table ID than their data rows when separated by an empty row. | |
| def _col_span(comp): | |
| return {c for _, c in comp} | |
| def _row_range(comp): | |
| rows = [r for r, _ in comp] | |
| return min(rows), max(rows) | |
| merged = True | |
| while merged: | |
| merged = False | |
| for i in range(len(components)): | |
| if components[i] is None: | |
| continue | |
| cols_i = _col_span(components[i]) | |
| rmin_i, rmax_i = _row_range(components[i]) | |
| for j in range(i + 1, len(components)): | |
| if components[j] is None: | |
| continue | |
| cols_j = _col_span(components[j]) | |
| rmin_j, rmax_j = _row_range(components[j]) | |
| gap = min(abs(rmin_j - rmax_i), abs(rmin_i - rmax_j)) | |
| if gap > 3: | |
| continue | |
| overlap = len(cols_i & cols_j) | |
| union = len(cols_i | cols_j) | |
| if union > 0 and overlap / union > 0.4: | |
| components[i] = components[i] | components[j] | |
| components[j] = None | |
| cols_i = _col_span(components[i]) | |
| rmin_i, rmax_i = _row_range(components[i]) | |
| merged = True | |
| components = [c for c in components if c is not None] | |
| for p in predictions: | |
| key = (p["row"], p["col"]) | |
| if p["label"] == "metadata": | |
| is_near_data = False | |
| r, c = key | |
| for dr in range(-2, 3): | |
| for dc in range(-2, 3): | |
| if (r + dr, c + dc) in data_coords: | |
| is_near_data = True | |
| break | |
| if is_near_data: | |
| break | |
| if not is_near_data: | |
| p["table"] = 0 | |
| continue | |
| for tid, comp in enumerate(components, start=1): | |
| if key in comp: | |
| p["table"] = tid | |
| break | |
| else: | |
| if p["label"] == "metadata": | |
| p["table"] = 0 | |
| elif p["label"] == "empty": | |
| p["table"] = 0 | |
| else: | |
| p["table"] = 1 | |
| def _assign_header_levels(predictions: list[dict]) -> None: | |
| """Assign header levels based on position within each table. | |
| Convention: L1 = most specific (closest to data), higher = broader. | |
| Column headers: bottom-to-top numbering (closest to data = L1). | |
| Row headers: right-to-left numbering (closest to data = L1). | |
| """ | |
| from collections import defaultdict | |
| tables = defaultdict(list) | |
| for p in predictions: | |
| tables[p.get("table", 0)].append(p) | |
| for tid, preds in tables.items(): | |
| col_header_rows = sorted({p["row"] for p in preds if p["label"] == "col_header"}) | |
| row_to_level = {row: lvl for lvl, row in enumerate(reversed(col_header_rows), start=1)} | |
| row_header_cols = sorted({p["col"] for p in preds if p["label"] == "row_header"}) | |
| col_to_level = {col: lvl for lvl, col in enumerate(reversed(row_header_cols), start=1)} | |
| for p in preds: | |
| if p["label"] == "col_header": | |
| level = row_to_level.get(p["row"], 1) | |
| p["label"] = f"col_header_{level}" | |
| elif p["label"] == "row_header": | |
| level = col_to_level.get(p["col"], 1) | |
| p["label"] = f"row_header_{level}" | |
| def _get_embed_model(): | |
| """Lazy-load the sentence-transformer model for text embeddings.""" | |
| global _cached_embed_model | |
| if _cached_embed_model is None: | |
| try: | |
| from embed_text import _get_model | |
| _cached_embed_model = _get_model() | |
| except Exception: | |
| logger.warning("Could not load embedding model; using zero embeddings") | |
| return _cached_embed_model | |
| def _is_numeric_value(val: str) -> bool: | |
| try: | |
| float(val.replace(",", "").replace(" ", "").replace("%", "")) | |
| return True | |
| except ValueError: | |
| return False | |
| def _generate_embeddings(cells: list[dict]) -> torch.Tensor: | |
| """Generate text embeddings for cells, matching training behavior. | |
| Text cells get real embeddings; numeric-only cells get zeros | |
| (same logic as embed_text.py). | |
| """ | |
| model = _get_embed_model() | |
| if model is None: | |
| return torch.zeros(len(cells), EMBED_DIM, dtype=torch.float) | |
| text_indices = [] | |
| texts = [] | |
| for i, c in enumerate(cells): | |
| val = c.get("value", "").strip() | |
| if val and not _is_numeric_value(val): | |
| text_indices.append(i) | |
| texts.append(f"passage: {val}") | |
| embeddings = np.zeros((len(cells), EMBED_DIM), dtype=np.float32) | |
| if texts: | |
| vecs = model.encode(texts, batch_size=64, show_progress_bar=False, normalize_embeddings=True) | |
| for idx, vec in zip(text_indices, vecs): | |
| embeddings[idx] = vec | |
| return torch.tensor(embeddings, dtype=torch.float) | |
| def _cluster_by_table_edges( | |
| edge_index: torch.Tensor, same_table_probs: torch.Tensor, | |
| num_nodes: int, threshold: float = 0.5, | |
| ) -> list[int]: | |
| """Connected components on predicted same-table edges.""" | |
| import networkx as nx | |
| G = nx.Graph() | |
| G.add_nodes_from(range(num_nodes)) | |
| mask = same_table_probs > threshold | |
| for i in range(edge_index.size(1)): | |
| if mask[i]: | |
| G.add_edge(edge_index[0, i].item(), edge_index[1, i].item()) | |
| components = sorted(nx.connected_components(G), key=len, reverse=True) | |
| table_ids = [0] * num_nodes | |
| for tid, comp in enumerate(components, start=1): | |
| for node_idx in comp: | |
| table_ids[node_idx] = tid | |
| return table_ids | |
| import re as _re | |
| _NUMERIC_RE = _re.compile( | |
| r"^[\s]*[\$β¬Β£Β₯]?\s*[\-\+]?\s*[\d,]+(?:\.\d+)?\s*[%]?\s*$" | |
| ) | |
| def _postprocess_value_attribute( | |
| cells: list[dict], predictions: list[dict] | |
| ) -> None: | |
| """Reclassify attribute cells as value using three complementary heuristics. | |
| 1. **Zero-value fallback** (original): when the model predicts zero value | |
| cells, flip numeric attribute cells that have text above in the same column. | |
| 2. **Column-level voting**: if >60% of non-empty predictions in a column are | |
| already ``value``, flip remaining numeric ``attribute`` cells in that column. | |
| 3. **Low-confidence flip**: if an attribute prediction has confidence < 0.7 and | |
| the cell is numeric with a text cell above, flip to ``value``. | |
| """ | |
| from collections import defaultdict | |
| pred_by_coord = {(p["row"], p["col"]): p for p in predictions} | |
| cell_by_coord = {(c["row"], c["col"]): c for c in cells} | |
| has_value = any(p["label"] == "value" for p in predictions) | |
| reclassified_rows = set() | |
| def _is_numeric_cell(p: dict) -> bool: | |
| cell = cell_by_coord.get((p["row"], p["col"])) | |
| if cell is None: | |
| return False | |
| return bool(_NUMERIC_RE.match(cell.get("value", ""))) | |
| def _has_text_above(p: dict) -> bool: | |
| for scan_r in range(p["row"] - 1, -1, -1): | |
| above_cell = cell_by_coord.get((scan_r, p["col"])) | |
| if above_cell is None: | |
| continue | |
| above_val = above_cell.get("value", "").strip() | |
| if above_val and not _NUMERIC_RE.match(above_val): | |
| return True | |
| return False | |
| # --- Heuristic 1: zero-value fallback (original behavior) --- | |
| if not has_value: | |
| for p in predictions: | |
| if p["label"] != "attribute": | |
| continue | |
| if _is_numeric_cell(p) and _has_text_above(p): | |
| p["label"] = "value" | |
| reclassified_rows.add(p["row"]) | |
| # --- Heuristic 2: column-level voting --- | |
| preds_by_col: dict[int, list[dict]] = defaultdict(list) | |
| for p in predictions: | |
| if p["label"] in ("value", "attribute", "aggregation"): | |
| preds_by_col[p["col"]].append(p) | |
| for col_idx, col_preds in preds_by_col.items(): | |
| value_count = sum(1 for p in col_preds if p["label"] == "value") | |
| total = len(col_preds) | |
| if total > 0 and value_count / total > 0.6: | |
| for p in col_preds: | |
| if p["label"] == "attribute" and _is_numeric_cell(p): | |
| p["label"] = "value" | |
| reclassified_rows.add(p["row"]) | |
| # --- Heuristic 3: low-confidence flip --- | |
| for p in predictions: | |
| if p["label"] != "attribute": | |
| continue | |
| if p.get("confidence", 1.0) < 0.7: | |
| if _is_numeric_cell(p) and _has_text_above(p): | |
| p["label"] = "value" | |
| reclassified_rows.add(p["row"]) | |
| # For rows that gained reclassified value cells, promote the | |
| # leftmost remaining text attribute cell to row_header_1 so that | |
| # value chunks get proper row context. | |
| if reclassified_rows: | |
| rows_attrs: dict[int, list[dict]] = defaultdict(list) | |
| for p in predictions: | |
| if p["row"] in reclassified_rows and p["label"] == "attribute": | |
| cell = cell_by_coord.get((p["row"], p["col"])) | |
| if cell: | |
| val = cell.get("value", "").strip() | |
| if val and not _NUMERIC_RE.match(val): | |
| rows_attrs[p["row"]].append(p) | |
| for row, attrs in rows_attrs.items(): | |
| attrs.sort(key=lambda a: a["col"]) | |
| attrs[0]["label"] = "row_header_1" | |
| def predict_sheet( | |
| grid: dict, | |
| use_heuristics: bool = False, | |
| model_override=None, | |
| ) -> Optional[list[dict]]: | |
| """Run GNN inference on a raw grid and return predicted labels. | |
| Args: | |
| grid: Sheet data dict with ``cells``, ``grid_size``, etc. | |
| use_heuristics: Force heuristic post-processing even for new models | |
| model_override: Use this model instead of the cached/default one. | |
| (useful for ablation comparison). | |
| Returns a list of {"row", "col", "label", "confidence", "table"} dicts, | |
| or None if no model is available. | |
| """ | |
| model = model_override or load_model() | |
| if model is None: | |
| return None | |
| output_labels = getattr(model, "_output_labels", COARSE_LABELS) | |
| legacy = getattr(model, "_legacy", True) | |
| cells = grid.get("cells", []) | |
| all_cells = [c for c in cells if c.get("value", "") != ""] | |
| if len(all_cells) < 3: | |
| return None | |
| max_row = max(c["row"] for c in all_cells) + 1 | |
| max_col = max(c["col"] for c in all_cells) + 1 | |
| idx_map = {} | |
| for i, c in enumerate(all_cells): | |
| idx_map[(c["row"], c["col"])] = i | |
| font_sizes = [c.get("font_size") for c in all_cells if c.get("font_size")] | |
| median_fs = float(sorted(font_sizes)[len(font_sizes) // 2]) if font_sizes else 11.0 | |
| from train_gnn import _compute_col_row_stats, _compute_cohesion_stats | |
| crs = _compute_col_row_stats(all_cells) | |
| coh = _compute_cohesion_stats(all_cells) | |
| x_structural = torch.tensor( | |
| [_cell_features_inference(c, max_row, max_col, median_fs, crs, coh) for c in all_cells], | |
| dtype=torch.float, | |
| ) | |
| x_embed = _generate_embeddings(all_cells) | |
| expected_dim = getattr(model, "_train_in_dim", None) | |
| if expected_dim is None and _cached_config is not None: | |
| expected_dim = _cached_config.get("in_dim") | |
| if expected_dim is None: | |
| expected_dim = x_structural.size(1) + x_embed.size(1) | |
| x = _align_structural_embed_to_dim(x_structural, x_embed, int(expected_dim)) | |
| graph_cfg = getattr(model, "_graph_config", None) | |
| num_virtual = 0 | |
| if graph_cfg is not None: | |
| edge_index, edge_attr, num_virtual = _build_inference_edges( | |
| all_cells, idx_map, graph_cfg) | |
| if edge_index.size(1) == 0: | |
| return None | |
| if num_virtual > 0: | |
| x = torch.cat([x, torch.zeros(num_virtual, x.size(1))], dim=0) | |
| else: | |
| edges = _build_flat_spatial_edges(all_cells, idx_map) | |
| if not edges[0]: | |
| return None | |
| edge_index = torch.tensor(edges, dtype=torch.long) | |
| model_etd = getattr(model, "_edge_type_dim", EDGE_TYPE_DIM) | |
| edge_attr = torch.zeros(edge_index.size(1), model_etd) | |
| edge_attr[:, 0] = 1.0 | |
| grid_kwargs = {} | |
| if getattr(model, "use_cnn_grid", False): | |
| from train_gnn import _build_cnn_grid, CNN_GRID_MAX_ROWS, CNN_GRID_MAX_COLS | |
| grid_features = _build_cnn_grid(all_cells, max_row, max_col) | |
| grid_H = min(max_row, CNN_GRID_MAX_ROWS) | |
| grid_W = min(max_col, CNN_GRID_MAX_COLS) | |
| num_real = len(all_cells) | |
| grid_row_col = torch.zeros(num_real + num_virtual, 2, dtype=torch.long) | |
| grid_valid = torch.zeros(num_real + num_virtual, dtype=torch.bool) | |
| for i, c in enumerate(all_cells): | |
| r, col = c["row"], c["col"] | |
| if r < grid_H and col < grid_W: | |
| grid_row_col[i, 0] = r | |
| grid_row_col[i, 1] = col | |
| grid_valid[i] = True | |
| grid_kwargs = { | |
| "grid_features": grid_features, | |
| "grid_row_col": grid_row_col, | |
| "grid_valid": grid_valid, | |
| } | |
| is_graph_learner = getattr(model, "_is_graph_learner", False) | |
| if is_graph_learner: | |
| node_rows = torch.tensor([c["row"] for c in all_cells], dtype=torch.long) | |
| node_cols = torch.tensor([c["col"] for c in all_cells], dtype=torch.long) | |
| if num_virtual > 0: | |
| node_rows = torch.cat([node_rows, torch.zeros(num_virtual, dtype=torch.long)]) | |
| node_cols = torch.cat([node_cols, torch.zeros(num_virtual, dtype=torch.long)]) | |
| if hasattr(model, "_struct_refine_active"): | |
| model._struct_refine_active = True | |
| h = model(x, edge_index, edge_attr, node_rows=node_rows, node_cols=node_cols) | |
| else: | |
| h = model(x, edge_index, edge_attr, **grid_kwargs) | |
| logits = model.classify_nodes(h) | |
| probs = F.softmax(logits, dim=-1) | |
| pred_classes = probs.argmax(dim=-1).cpu().numpy() | |
| confidences = probs.max(dim=-1).values.cpu().numpy() | |
| # Build label predictions using the correct label list for this model | |
| predictions = [] | |
| for i, c in enumerate(all_cells): | |
| cls_idx = int(pred_classes[i]) | |
| label_name = output_labels[cls_idx] if cls_idx < len(output_labels) else "value" | |
| predictions.append({ | |
| "row": c["row"], | |
| "col": c["col"], | |
| "label": label_name, | |
| "confidence": round(float(confidences[i]), 3), | |
| "table": 1, | |
| }) | |
| pred_adjs = {} | |
| if legacy or use_heuristics: | |
| _assign_table_ids(all_cells, predictions) | |
| _assign_header_levels(predictions) | |
| elif is_graph_learner and hasattr(model, "predict_adjacency"): | |
| import networkx as _nx | |
| from models import SpatialEdgeTransformerModel, SpatialBilinearTransformerModel | |
| if isinstance(model, SpatialEdgeTransformerModel): | |
| adj_pred = model.predict_adjacency(h, node_rows, node_cols) | |
| elif isinstance(model, SpatialBilinearTransformerModel): | |
| adj_pred = model.predict_adjacency(h, node_rows, node_cols) | |
| else: | |
| adj_pred = model.predict_adjacency(h) | |
| for adj_type, logit_mat in adj_pred.items(): | |
| prob_mat = torch.sigmoid(logit_mat) | |
| sym = (prob_mat + prob_mat.T) / 2 if adj_type != "hier" else prob_mat | |
| edge_mask = sym > 0.5 | |
| pairs = edge_mask.nonzero(as_tuple=False) | |
| pred_adjs[adj_type] = {(r[0].item(), r[1].item()) for r in pairs} | |
| if "table" in pred_adjs and pred_adjs["table"]: | |
| G = _nx.Graph() | |
| G.add_nodes_from(range(len(all_cells))) | |
| G.add_edges_from(pred_adjs["table"]) | |
| for comp_idx, comp in enumerate( | |
| sorted(_nx.connected_components(G), key=len, reverse=True), start=1 | |
| ): | |
| for node in comp: | |
| if node < len(predictions): | |
| predictions[node]["table"] = comp_idx | |
| else: | |
| _assign_table_ids(all_cells, predictions) | |
| else: | |
| if hasattr(model, "predict_table_edges"): | |
| table_logits = model.predict_table_edges(h, edge_index) | |
| table_probs = torch.sigmoid(table_logits) | |
| table_ids = _cluster_by_table_edges( | |
| edge_index, table_probs, len(all_cells), | |
| ) | |
| for i, p in enumerate(predictions): | |
| p["table"] = table_ids[i] | |
| else: | |
| _assign_table_ids(all_cells, predictions) | |
| _postprocess_value_attribute(all_cells, predictions) | |
| return predictions | |
| # βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| # Batch prediction & pre-computed prediction storage | |
| # βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| def _safe_filename(drive_file_id: str, sheet_name: str) -> str: | |
| """Build a safe filename from drive_file_id and sheet_name.""" | |
| return f"{drive_file_id}_{sheet_name}".replace("/", "_").replace(" ", "_") + ".json" | |
| def batch_predict_all(version: str) -> None: | |
| """Predict all unlabeled sheets with the given model version and save to disk. | |
| Stores results in data/predictions/v{version}/{drive_file_id}_{sheet_name}.json. | |
| Skips sheets that already have a prediction file for this version (idempotent). | |
| """ | |
| import drive_client | |
| import metadata_client | |
| import xlsx_client | |
| model = load_model(version) | |
| if model is None: | |
| print(f"No model found for version v{version}. Skipping batch prediction.") | |
| return | |
| version_dir = config.PREDICTIONS_DIR / f"v{version}" | |
| version_dir.mkdir(parents=True, exist_ok=True) | |
| rows = metadata_client.get_all_rows() | |
| pending = [r for r in rows if r["status"] in {"unlabelled", "in_progress"}] | |
| if not pending: | |
| print("No unlabeled sheets to predict.") | |
| return | |
| print(f"Predicting {len(pending)} sheets with model v{version}...") | |
| predicted = 0 | |
| skipped = 0 | |
| errors = 0 | |
| for i, row in enumerate(pending, 1): | |
| fname = _safe_filename(row["drive_file_id"], row["sheet_name"]) | |
| pred_path = version_dir / fname | |
| if pred_path.exists(): | |
| skipped += 1 | |
| continue | |
| try: | |
| local_path = str(config.TEMP_DIR / f"{row['drive_file_id']}.xlsx") | |
| if not Path(local_path).exists(): | |
| drive_client.download_file(row["drive_file_id"], local_path) | |
| grid = xlsx_client.fetch_xlsx_sheet(local_path, row["sheet_name"]) | |
| preds = predict_sheet(grid) | |
| if preds: | |
| pred_data = { | |
| "model_version": f"v{version}", | |
| "timestamp": datetime.now(timezone.utc).isoformat(), | |
| "drive_file_id": row["drive_file_id"], | |
| "sheet_name": row["sheet_name"], | |
| "predictions": preds, | |
| } | |
| pred_path.write_text( | |
| json.dumps(pred_data, ensure_ascii=False, indent=1) | |
| ) | |
| predicted += 1 | |
| print(f" [{i}/{len(pending)}] Predicted {row['sheet_name']} from {row['file_name']}") | |
| else: | |
| skipped += 1 | |
| print(f" [{i}/{len(pending)}] Skipped {row['sheet_name']} (too few cells)") | |
| except Exception as e: | |
| errors += 1 | |
| logger.warning("Error predicting %s/%s: %s", row["file_name"], row["sheet_name"], e) | |
| print(f" [{i}/{len(pending)}] Error: {row['sheet_name']} - {e}") | |
| print(f"\nBatch prediction complete: {predicted} predicted, {skipped} skipped, {errors} errors") | |
| def load_predictions( | |
| drive_file_id: str, sheet_name: str | |
| ) -> tuple[Optional[list[dict]], Optional[str]]: | |
| """Load pre-computed predictions from the latest model version directory. | |
| Returns (predictions_list, model_version) or (None, None) if not found. | |
| """ | |
| latest = get_latest_model_version() | |
| if latest is None: | |
| return None, None | |
| version_dir = config.PREDICTIONS_DIR / f"v{latest}" | |
| if not version_dir.exists(): | |
| return None, None | |
| fname = _safe_filename(drive_file_id, sheet_name) | |
| pred_path = version_dir / fname | |
| if not pred_path.exists(): | |
| return None, None | |
| try: | |
| data = json.loads(pred_path.read_text()) | |
| return data.get("predictions"), data.get("model_version", f"v{latest}") | |
| except (json.JSONDecodeError, OSError) as e: | |
| logger.warning("Error reading prediction file %s: %s", pred_path, e) | |
| return None, None | |