parse-bench 1.0.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (227) hide show
  1. parse_bench/__init__.py +3 -0
  2. parse_bench/analysis/__init__.py +6 -0
  3. parse_bench/analysis/aggregation_report.py +582 -0
  4. parse_bench/analysis/cli.py +472 -0
  5. parse_bench/analysis/comparison.py +382 -0
  6. parse_bench/analysis/comparison_core.py +357 -0
  7. parse_bench/analysis/comparison_report.py +2066 -0
  8. parse_bench/analysis/detailed_report.py +2254 -0
  9. parse_bench/analysis/leaderboard_report.py +852 -0
  10. parse_bench/analysis/metric_definitions.py +771 -0
  11. parse_bench/cli.py +267 -0
  12. parse_bench/data/__init__.py +1 -0
  13. parse_bench/data/cli.py +118 -0
  14. parse_bench/data/download.py +127 -0
  15. parse_bench/evaluation/__init__.py +11 -0
  16. parse_bench/evaluation/cli.py +435 -0
  17. parse_bench/evaluation/evaluators/__init__.py +17 -0
  18. parse_bench/evaluation/evaluators/base.py +34 -0
  19. parse_bench/evaluation/evaluators/extract.py +429 -0
  20. parse_bench/evaluation/evaluators/layoutdet.py +1682 -0
  21. parse_bench/evaluation/evaluators/parse.py +1353 -0
  22. parse_bench/evaluation/evaluators/qa.py +199 -0
  23. parse_bench/evaluation/layout_adapters/__init__.py +21 -0
  24. parse_bench/evaluation/layout_adapters/adapters.py +3180 -0
  25. parse_bench/evaluation/layout_adapters/base.py +105 -0
  26. parse_bench/evaluation/layout_adapters/registry.py +109 -0
  27. parse_bench/evaluation/layout_label_mappers/__init__.py +22 -0
  28. parse_bench/evaluation/layout_label_mappers/base.py +66 -0
  29. parse_bench/evaluation/layout_label_mappers/mappers.py +332 -0
  30. parse_bench/evaluation/layout_label_mappers/projection.py +74 -0
  31. parse_bench/evaluation/layout_label_mappers/registry.py +119 -0
  32. parse_bench/evaluation/metric_aggregation.py +56 -0
  33. parse_bench/evaluation/metrics/__init__.py +5 -0
  34. parse_bench/evaluation/metrics/attribution/__init__.py +35 -0
  35. parse_bench/evaluation/metrics/attribution/constants.py +12 -0
  36. parse_bench/evaluation/metrics/attribution/core.py +1108 -0
  37. parse_bench/evaluation/metrics/attribution/evaluate.py +446 -0
  38. parse_bench/evaluation/metrics/attribution/geometry.py +161 -0
  39. parse_bench/evaluation/metrics/attribution/text_utils.py +233 -0
  40. parse_bench/evaluation/metrics/base.py +33 -0
  41. parse_bench/evaluation/metrics/downstream/__init__.py +0 -0
  42. parse_bench/evaluation/metrics/extract/__init__.py +29 -0
  43. parse_bench/evaluation/metrics/extract/json_subset_match.py +473 -0
  44. parse_bench/evaluation/metrics/extract/json_subset_match_metric.py +81 -0
  45. parse_bench/evaluation/metrics/extract/list_unwrap.py +340 -0
  46. parse_bench/evaluation/metrics/extract/rule_based_metric.py +90 -0
  47. parse_bench/evaluation/metrics/extract/test_rules.py +409 -0
  48. parse_bench/evaluation/metrics/extract/test_types.py +11 -0
  49. parse_bench/evaluation/metrics/field_grounding/__init__.py +21 -0
  50. parse_bench/evaluation/metrics/field_grounding/core.py +437 -0
  51. parse_bench/evaluation/metrics/field_grounding/extract_adapter.py +1224 -0
  52. parse_bench/evaluation/metrics/field_grounding/parse_adapter.py +697 -0
  53. parse_bench/evaluation/metrics/field_grounding/rule_filters.py +19 -0
  54. parse_bench/evaluation/metrics/field_grounding/value_compare.py +190 -0
  55. parse_bench/evaluation/metrics/layoutdet/__init__.py +17 -0
  56. parse_bench/evaluation/metrics/layoutdet/classification_utils.py +300 -0
  57. parse_bench/evaluation/metrics/layoutdet/iou.py +76 -0
  58. parse_bench/evaluation/metrics/parse/__init__.py +5 -0
  59. parse_bench/evaluation/metrics/parse/_vendor_grits_reference.py +531 -0
  60. parse_bench/evaluation/metrics/parse/cross_page_table_consistency.py +165 -0
  61. parse_bench/evaluation/metrics/parse/emphasis_spans.py +242 -0
  62. parse_bench/evaluation/metrics/parse/fast_tree_edit.py +282 -0
  63. parse_bench/evaluation/metrics/parse/grits_metric.py +1125 -0
  64. parse_bench/evaluation/metrics/parse/grits_reference_metric.py +142 -0
  65. parse_bench/evaluation/metrics/parse/header_accuracy_metric.py +1662 -0
  66. parse_bench/evaluation/metrics/parse/llm_normalization/__init__.py +51 -0
  67. parse_bench/evaluation/metrics/parse/llm_normalization/base.py +125 -0
  68. parse_bench/evaluation/metrics/parse/llm_normalization/config.py +44 -0
  69. parse_bench/evaluation/metrics/parse/llm_normalization/postprocess.py +322 -0
  70. parse_bench/evaluation/metrics/parse/llm_normalization/strategy_judge.py +541 -0
  71. parse_bench/evaluation/metrics/parse/mermaid_graph.py +682 -0
  72. parse_bench/evaluation/metrics/parse/rule_based_judge_metric.py +56 -0
  73. parse_bench/evaluation/metrics/parse/rule_based_metric.py +434 -0
  74. parse_bench/evaluation/metrics/parse/rules_bag.py +1161 -0
  75. parse_bench/evaluation/metrics/parse/rules_base.py +751 -0
  76. parse_bench/evaluation/metrics/parse/rules_chart.py +1556 -0
  77. parse_bench/evaluation/metrics/parse/rules_diagram.py +591 -0
  78. parse_bench/evaluation/metrics/parse/rules_form.py +2274 -0
  79. parse_bench/evaluation/metrics/parse/rules_formatting.py +1500 -0
  80. parse_bench/evaluation/metrics/parse/rules_heading.py +228 -0
  81. parse_bench/evaluation/metrics/parse/rules_list.py +226 -0
  82. parse_bench/evaluation/metrics/parse/rules_page_decoration.py +276 -0
  83. parse_bench/evaluation/metrics/parse/rules_table.py +1666 -0
  84. parse_bench/evaluation/metrics/parse/rules_text.py +340 -0
  85. parse_bench/evaluation/metrics/parse/rules_watermark.py +105 -0
  86. parse_bench/evaluation/metrics/parse/structural_consistency_metric.py +251 -0
  87. parse_bench/evaluation/metrics/parse/table_extraction.py +152 -0
  88. parse_bench/evaluation/metrics/parse/table_merging.py +195 -0
  89. parse_bench/evaluation/metrics/parse/table_pairing.py +87 -0
  90. parse_bench/evaluation/metrics/parse/table_parsing.py +955 -0
  91. parse_bench/evaluation/metrics/parse/table_record_match_metric.py +1453 -0
  92. parse_bench/evaluation/metrics/parse/table_splitting.py +301 -0
  93. parse_bench/evaluation/metrics/parse/table_title_stripping.py +530 -0
  94. parse_bench/evaluation/metrics/parse/teds_metric.py +600 -0
  95. parse_bench/evaluation/metrics/parse/test_rules.py +120 -0
  96. parse_bench/evaluation/metrics/parse/test_types.py +103 -0
  97. parse_bench/evaluation/metrics/parse/text_content_projection.py +175 -0
  98. parse_bench/evaluation/metrics/parse/text_similarity_metric.py +61 -0
  99. parse_bench/evaluation/metrics/parse/utils.py +885 -0
  100. parse_bench/evaluation/metrics/qa/__init__.py +5 -0
  101. parse_bench/evaluation/metrics/qa/answer_comparison.py +380 -0
  102. parse_bench/evaluation/qa/__init__.py +5 -0
  103. parse_bench/evaluation/qa/llm_service.py +335 -0
  104. parse_bench/evaluation/reports/__init__.py +8 -0
  105. parse_bench/evaluation/reports/csv.py +64 -0
  106. parse_bench/evaluation/reports/html.py +338 -0
  107. parse_bench/evaluation/reports/markdown.py +98 -0
  108. parse_bench/evaluation/reports/rule_csv.py +22 -0
  109. parse_bench/evaluation/runner.py +1864 -0
  110. parse_bench/evaluation/stats.py +104 -0
  111. parse_bench/extensions.py +72 -0
  112. parse_bench/inference/__init__.py +33 -0
  113. parse_bench/inference/chunkr_layout_extraction.py +160 -0
  114. parse_bench/inference/cli.py +484 -0
  115. parse_bench/inference/layout_extraction.py +422 -0
  116. parse_bench/inference/pipelines/__init__.py +59 -0
  117. parse_bench/inference/pipelines/extract.py +39 -0
  118. parse_bench/inference/pipelines/layout.py +142 -0
  119. parse_bench/inference/pipelines/parse.py +2603 -0
  120. parse_bench/inference/pipelines.py +0 -0
  121. parse_bench/inference/providers/__init__.py +28 -0
  122. parse_bench/inference/providers/base.py +196 -0
  123. parse_bench/inference/providers/cancellation.py +137 -0
  124. parse_bench/inference/providers/extract/__init__.py +22 -0
  125. parse_bench/inference/providers/extract/citations.py +549 -0
  126. parse_bench/inference/providers/extract/extend.py +851 -0
  127. parse_bench/inference/providers/extract/llamaextract_v2_api.py +583 -0
  128. parse_bench/inference/providers/layoutdet/__init__.py +25 -0
  129. parse_bench/inference/providers/layoutdet/adapters.py +946 -0
  130. parse_bench/inference/providers/layoutdet/base.py +203 -0
  131. parse_bench/inference/providers/layoutdet/chandra.py +449 -0
  132. parse_bench/inference/providers/layoutdet/docling.py +125 -0
  133. parse_bench/inference/providers/layoutdet/dots_ocr.py +606 -0
  134. parse_bench/inference/providers/layoutdet/layout_v3.py +137 -0
  135. parse_bench/inference/providers/layoutdet/layout_v3_byoc.py +204 -0
  136. parse_bench/inference/providers/layoutdet/paddle.py +117 -0
  137. parse_bench/inference/providers/layoutdet/qwen3vl.py +360 -0
  138. parse_bench/inference/providers/layoutdet/surya.py +250 -0
  139. parse_bench/inference/providers/layoutdet/yolo.py +109 -0
  140. parse_bench/inference/providers/parse/__init__.py +64 -0
  141. parse_bench/inference/providers/parse/_docling_common.py +233 -0
  142. parse_bench/inference/providers/parse/_layout_utils.py +611 -0
  143. parse_bench/inference/providers/parse/amazon_nova.py +515 -0
  144. parse_bench/inference/providers/parse/anthropic.py +882 -0
  145. parse_bench/inference/providers/parse/azure_document_intelligence.py +700 -0
  146. parse_bench/inference/providers/parse/chandra2.py +633 -0
  147. parse_bench/inference/providers/parse/chunkr.py +268 -0
  148. parse_bench/inference/providers/parse/databricks_ai_parse.py +724 -0
  149. parse_bench/inference/providers/parse/datalab.py +370 -0
  150. parse_bench/inference/providers/parse/deepseekocr2.py +382 -0
  151. parse_bench/inference/providers/parse/docling.py +281 -0
  152. parse_bench/inference/providers/parse/docling_serve.py +289 -0
  153. parse_bench/inference/providers/parse/dots_ocr.py +574 -0
  154. parse_bench/inference/providers/parse/extend_parse.py +710 -0
  155. parse_bench/inference/providers/parse/falconocr.py +436 -0
  156. parse_bench/inference/providers/parse/florin_parser_nano.py +559 -0
  157. parse_bench/inference/providers/parse/gemma4.py +472 -0
  158. parse_bench/inference/providers/parse/glm_zai.py +229 -0
  159. parse_bench/inference/providers/parse/google.py +1125 -0
  160. parse_bench/inference/providers/parse/google_agentic_vision.py +819 -0
  161. parse_bench/inference/providers/parse/google_docai.py +776 -0
  162. parse_bench/inference/providers/parse/google_docai_layout_normalization.py +573 -0
  163. parse_bench/inference/providers/parse/granite_vision.py +515 -0
  164. parse_bench/inference/providers/parse/infinity_parser2.py +704 -0
  165. parse_bench/inference/providers/parse/kdl_frontier_nano.py +3327 -0
  166. parse_bench/inference/providers/parse/landingai.py +452 -0
  167. parse_bench/inference/providers/parse/liteparse.py +350 -0
  168. parse_bench/inference/providers/parse/llamaparse.py +677 -0
  169. parse_bench/inference/providers/parse/llamaparse_v2_normalization.py +1013 -0
  170. parse_bench/inference/providers/parse/markitdown.py +138 -0
  171. parse_bench/inference/providers/parse/mineru25.py +405 -0
  172. parse_bench/inference/providers/parse/mineru2605pro.py +432 -0
  173. parse_bench/inference/providers/parse/mineru_diffusion.py +371 -0
  174. parse_bench/inference/providers/parse/mistral_ocr.py +546 -0
  175. parse_bench/inference/providers/parse/nemotron_omni.py +473 -0
  176. parse_bench/inference/providers/parse/oi_parser.py +222 -0
  177. parse_bench/inference/providers/parse/openai.py +740 -0
  178. parse_bench/inference/providers/parse/opendataloader.py +152 -0
  179. parse_bench/inference/providers/parse/paddleocr.py +624 -0
  180. parse_bench/inference/providers/parse/pdf_inspector.py +142 -0
  181. parse_bench/inference/providers/parse/pulse.py +785 -0
  182. parse_bench/inference/providers/parse/pymupdf.py +207 -0
  183. parse_bench/inference/providers/parse/pymupdf4llm.py +356 -0
  184. parse_bench/inference/providers/parse/pypdf.py +179 -0
  185. parse_bench/inference/providers/parse/qwen.py +678 -0
  186. parse_bench/inference/providers/parse/rakedoc_nano.py +70 -0
  187. parse_bench/inference/providers/parse/reducto.py +546 -0
  188. parse_bench/inference/providers/parse/surya2.py +372 -0
  189. parse_bench/inference/providers/parse/tesseract.py +301 -0
  190. parse_bench/inference/providers/parse/textract.py +694 -0
  191. parse_bench/inference/providers/parse/unlimitedocr.py +346 -0
  192. parse_bench/inference/providers/parse/unstructured.py +485 -0
  193. parse_bench/inference/providers/parse/warp_ingest.py +199 -0
  194. parse_bench/inference/providers/registry.py +49 -0
  195. parse_bench/inference/renormalize.py +170 -0
  196. parse_bench/inference/runner.py +2023 -0
  197. parse_bench/layout_label_mapping.py +424 -0
  198. parse_bench/layout_projection.py +179 -0
  199. parse_bench/pipeline/__init__.py +1 -0
  200. parse_bench/pipeline/cli.py +549 -0
  201. parse_bench/schemas/__init__.py +33 -0
  202. parse_bench/schemas/evaluation.py +93 -0
  203. parse_bench/schemas/extract_output.py +36 -0
  204. parse_bench/schemas/layout_detection_output.py +545 -0
  205. parse_bench/schemas/layout_ontology.py +315 -0
  206. parse_bench/schemas/metrics.py +69 -0
  207. parse_bench/schemas/parse_output.py +152 -0
  208. parse_bench/schemas/pipeline.py +22 -0
  209. parse_bench/schemas/pipeline_io.py +106 -0
  210. parse_bench/schemas/product.py +97 -0
  211. parse_bench/test_cases/__init__.py +25 -0
  212. parse_bench/test_cases/bbox_value_strict_comparator.py +880 -0
  213. parse_bench/test_cases/extract_field_paths.py +164 -0
  214. parse_bench/test_cases/layout_attribution_generation.py +287 -0
  215. parse_bench/test_cases/loader.py +652 -0
  216. parse_bench/test_cases/parse_rule_schemas.py +1071 -0
  217. parse_bench/test_cases/rule_filters.py +32 -0
  218. parse_bench/test_cases/rule_ids.py +107 -0
  219. parse_bench/test_cases/schema.py +427 -0
  220. parse_bench/utils/__init__.py +15 -0
  221. parse_bench/utils/gemini_layout_utils.py +670 -0
  222. parse_bench/utils/text_aggregation.py +100 -0
  223. parse_bench-1.0.0.dist-info/METADATA +476 -0
  224. parse_bench-1.0.0.dist-info/RECORD +227 -0
  225. parse_bench-1.0.0.dist-info/WHEEL +4 -0
  226. parse_bench-1.0.0.dist-info/entry_points.txt +2 -0
  227. parse_bench-1.0.0.dist-info/licenses/LICENSE +201 -0
@@ -0,0 +1,1662 @@
1
+ """Header accuracy metric for HTML table comparison.
2
+
3
+ Evaluates how accurately table headers are reproduced by comparing
4
+ predicted HTML tables against ground-truth HTML tables. Produces a
5
+ composite score from eight submetrics (``header_perfect`` is still
6
+ emitted but excluded from the composite):
7
+
8
+ 1. **header_cell_count**: Ratio of predicted header cell count to expected.
9
+ Penalises both missing and extra header cells symmetrically.
10
+ 2. **header_grits**: GriTS-Con applied to contiguous header blocks,
11
+ measuring how well the mapping between header regions is preserved.
12
+ Extra predicted blocks are penalised.
13
+ 3. **header_content_bag**: Bag-of-cells exact content overlap — counts how
14
+ many expected header texts appear (exact match after formatting
15
+ normalization) in the prediction.
16
+ 4. **header_perfect**: Binary — 1.0 iff the header structure
17
+ (cell texts, positions, colspan/rowspan) matches exactly.
18
+ *Emitted but not included in the composite.*
19
+ 5. **header_block_extent**: Measures how well each header block's location
20
+ and size within the full table matches what is expected.
21
+ 6. **header_block_proximity**: Averaged nearest-edge distance similarity
22
+ between matched header block pairs. Extra predicted blocks are
23
+ penalised via ``max(gt_pairs, pred_pairs)`` denominator.
24
+ 7. **header_block_relative_direction**: Averaged cosine similarity of
25
+ signed edge vectors between matched header block pairs, mapped to
26
+ [0, 1]. Extra predicted blocks are penalised via
27
+ ``max(gt_pairs, pred_pairs)`` denominator.
28
+ 8. **multilevel_header_depth**: Compares the depth of the header hierarchy
29
+ tree (number of nesting levels) between GT and prediction using
30
+ ``min(depth_gt, depth_pred) / max(depth_gt, depth_pred)``.
31
+ 9. **header_data_alignment**: Uses GriTS row/col alignment maps to check
32
+ whether each GT header cell's text appears at the corresponding
33
+ position in the prediction grid. When GriTS alignment is not
34
+ available, computes its own alignment as a fallback.
35
+
36
+ The overall ``header_composite_v3`` score is the mean of *applicable*
37
+ submetrics from 1–3 and 5–9 (i.e. excluding ``header_perfect``).
38
+ Submetrics that are trivially 1.0 because both GT and prediction have
39
+ the same degenerate block count (e.g. both have ≤1 block, so proximity
40
+ and direction metrics are vacuous) are excluded from the composite
41
+ denominator. When the counts *differ* (e.g. GT has 0 blocks but
42
+ prediction has 2), the submetric is kept so the mismatch is still
43
+ penalised.
44
+
45
+ Table-level matching is delegated to the caller (typically from GriTS
46
+ matching) via the ``table_pairs`` parameter. Block-level matching is
47
+ computed once (via GriTS-Con Hungarian) and shared across all block
48
+ submetrics for consistency.
49
+ """
50
+
51
+ from __future__ import annotations
52
+
53
+ from dataclasses import dataclass, field
54
+ from typing import Any
55
+
56
+ import numpy as np
57
+ from bs4 import BeautifulSoup
58
+ from scipy.optimize import linear_sum_assignment
59
+
60
+ from parse_bench.evaluation.metrics.base import Metric
61
+ from parse_bench.evaluation.metrics.parse.grits_metric import (
62
+ _lcs_similarity,
63
+ factored_2dmss,
64
+ )
65
+ from parse_bench.evaluation.metrics.parse.table_extraction import extract_html_tables
66
+ from parse_bench.evaluation.metrics.parse.table_parsing import _sup_sub_to_unicode
67
+ from parse_bench.evaluation.metrics.parse.utils import (
68
+ cells_match_leader_insensitive,
69
+ normalize_cell_text,
70
+ normalize_text,
71
+ )
72
+ from parse_bench.schemas.evaluation import MetricValue
73
+
74
+
75
+ def _normalize_header_text(raw: str) -> str:
76
+ """Normalize cell text for header comparison.
77
+
78
+ Applies cell-level normalization (formatting stripping, dash/dot
79
+ normalization) via ``normalize_cell_text``, then the full
80
+ ``normalize_text`` pipeline (lowercasing, accent removal, etc.).
81
+ """
82
+ text = normalize_cell_text(raw)
83
+ return normalize_text(text)
84
+
85
+
86
+ # ---------------------------------------------------------------------------
87
+ # Header block extraction
88
+ # ---------------------------------------------------------------------------
89
+
90
+
91
+ @dataclass
92
+ class HeaderCell:
93
+ """A single header cell with its grid position and span info."""
94
+
95
+ text: str
96
+ row: int
97
+ col: int
98
+ rowspan: int
99
+ colspan: int
100
+
101
+
102
+ @dataclass
103
+ class HeaderBlock:
104
+ """A contiguous rectangular region of header cells."""
105
+
106
+ cells: list[HeaderCell] = field(default_factory=list)
107
+ min_row: int = 0
108
+ max_row: int = 0 # exclusive
109
+ min_col: int = 0
110
+ max_col: int = 0 # exclusive
111
+
112
+ def extent(self, table_rows: int, table_cols: int) -> tuple[float, float, float, float]:
113
+ """Normalised (row_start, col_start, row_end, col_end) in [0, 1]."""
114
+ if table_rows == 0 or table_cols == 0:
115
+ return (0.0, 0.0, 0.0, 0.0)
116
+ return (
117
+ self.min_row / table_rows,
118
+ self.min_col / table_cols,
119
+ self.max_row / table_rows,
120
+ self.max_col / table_cols,
121
+ )
122
+
123
+ def center(self, table_rows: int, table_cols: int) -> tuple[float, float]:
124
+ """Normalised center (row, col) of this block."""
125
+ if table_rows == 0 or table_cols == 0:
126
+ return (0.0, 0.0)
127
+ return (
128
+ (self.min_row + self.max_row) / 2.0 / table_rows,
129
+ (self.min_col + self.max_col) / 2.0 / table_cols,
130
+ )
131
+
132
+
133
+ def _parse_header_cells(table_html: str) -> tuple[list[HeaderCell], int, int]:
134
+ """Extract header cells from an HTML table string.
135
+
136
+ Returns (header_cells, num_rows, num_cols).
137
+ """
138
+ # Use lxml for robustness with malformed HTML (e.g. <th>...</td> mismatches)
139
+ soup = BeautifulSoup(table_html, "lxml")
140
+ table = soup.find("table")
141
+ if not table:
142
+ return [], 0, 0
143
+
144
+ rows = table.find_all("tr")
145
+ if not rows:
146
+ return [], 0, 0
147
+
148
+ # Build occupied grid to resolve spans
149
+ occupied: dict[tuple[int, int], bool] = {}
150
+ header_cells: list[HeaderCell] = []
151
+ max_col = 0
152
+
153
+ # Identify thead rows
154
+ thead = table.find("thead")
155
+ thead_row_indices: set[int] = set()
156
+ if thead:
157
+ for tr in thead.find_all("tr"):
158
+ if tr in rows:
159
+ thead_row_indices.add(rows.index(tr))
160
+
161
+ for row_idx, row in enumerate(rows):
162
+ col_idx = 0
163
+ for cell in row.find_all(["td", "th"]):
164
+ while (row_idx, col_idx) in occupied:
165
+ col_idx += 1
166
+
167
+ rowspan = int(str(cell.get("rowspan", "1")))
168
+ colspan = int(str(cell.get("colspan", "1")))
169
+ is_header = cell.name == "th" or row_idx in thead_row_indices
170
+
171
+ # Mark occupied
172
+ for r in range(row_idx, row_idx + rowspan):
173
+ for c in range(col_idx, col_idx + colspan):
174
+ occupied[(r, c)] = True
175
+
176
+ if is_header:
177
+ # Convert <sup>/<sub> digit content to Unicode equivalents
178
+ # so "Name<sup>1</sup>" becomes "Name¹", matching sources
179
+ # that already use Unicode superscripts.
180
+ _sup_sub_to_unicode(cell)
181
+ text = _normalize_header_text(cell.get_text(strip=True))
182
+ header_cells.append(
183
+ HeaderCell(
184
+ text=text,
185
+ row=row_idx,
186
+ col=col_idx,
187
+ rowspan=rowspan,
188
+ colspan=colspan,
189
+ )
190
+ )
191
+
192
+ max_col = max(max_col, col_idx + colspan)
193
+ col_idx += colspan
194
+
195
+ num_rows = len(rows)
196
+ num_cols = max_col if max_col > 0 else 0
197
+ return header_cells, num_rows, num_cols
198
+
199
+
200
+ def _find_header_blocks(cells: list[HeaderCell]) -> list[HeaderBlock]:
201
+ """Group header cells into contiguous rectangular blocks.
202
+
203
+ Two header cells belong to the same block if they are adjacent
204
+ (horizontally or vertically, using 4-connected adjacency).
205
+ """
206
+ if not cells:
207
+ return []
208
+
209
+ # Build adjacency via grid positions occupied by each cell
210
+ cell_positions: dict[int, set[tuple[int, int]]] = {}
211
+ for idx, cell in enumerate(cells):
212
+ positions = set()
213
+ for r in range(cell.row, cell.row + cell.rowspan):
214
+ for c in range(cell.col, cell.col + cell.colspan):
215
+ positions.add((r, c))
216
+ cell_positions[idx] = positions
217
+
218
+ # Union-Find
219
+ parent = list(range(len(cells)))
220
+
221
+ def find(x: int) -> int:
222
+ while parent[x] != x:
223
+ parent[x] = parent[parent[x]]
224
+ x = parent[x]
225
+ return x
226
+
227
+ def union(a: int, b: int) -> None:
228
+ ra, rb = find(a), find(b)
229
+ if ra != rb:
230
+ parent[ra] = rb
231
+
232
+ # Two cells are in the same block if any of their occupied positions
233
+ # are adjacent (4-connected: up/down/left/right only).
234
+ all_pos_to_idx: dict[tuple[int, int], int] = {}
235
+ for idx, positions in cell_positions.items():
236
+ for pos in positions:
237
+ all_pos_to_idx[pos] = idx
238
+
239
+ for idx, positions in cell_positions.items():
240
+ for r, c in positions:
241
+ for dr, dc in ((-1, 0), (1, 0), (0, -1), (0, 1)):
242
+ neighbor = (r + dr, c + dc)
243
+ if neighbor in all_pos_to_idx:
244
+ union(idx, all_pos_to_idx[neighbor])
245
+
246
+ # Group by root
247
+ groups: dict[int, list[int]] = {}
248
+ for idx in range(len(cells)):
249
+ root = find(idx)
250
+ groups.setdefault(root, []).append(idx)
251
+
252
+ blocks: list[HeaderBlock] = []
253
+ for indices in groups.values():
254
+ block_cells = [cells[i] for i in indices]
255
+ min_row = min(c.row for c in block_cells)
256
+ max_row = max(c.row + c.rowspan for c in block_cells)
257
+ min_col = min(c.col for c in block_cells)
258
+ max_col = max(c.col + c.colspan for c in block_cells)
259
+ blocks.append(
260
+ HeaderBlock(
261
+ cells=block_cells,
262
+ min_row=min_row,
263
+ max_row=max_row,
264
+ min_col=min_col,
265
+ max_col=max_col,
266
+ )
267
+ )
268
+
269
+ # Sort blocks by (min_row, min_col) for stable ordering
270
+ blocks.sort(key=lambda b: (b.min_row, b.min_col))
271
+ return blocks
272
+
273
+
274
+ # ---------------------------------------------------------------------------
275
+ # Block matching (shared across all block submetrics)
276
+ # ---------------------------------------------------------------------------
277
+
278
+
279
+ def _block_to_text_grid(block: HeaderBlock) -> np.ndarray:
280
+ """Build a text grid for a header block (for GriTS-Con)."""
281
+ rows = block.max_row - block.min_row
282
+ cols = block.max_col - block.min_col
283
+ if rows <= 0 or cols <= 0:
284
+ return np.array([[""]], dtype=object)
285
+ grid = np.full((rows, cols), "", dtype=object)
286
+ for cell in block.cells:
287
+ for r in range(cell.row, cell.row + cell.rowspan):
288
+ for c in range(cell.col, cell.col + cell.colspan):
289
+ grid[r - block.min_row, c - block.min_col] = cell.text
290
+ return grid
291
+
292
+
293
+ def _block_extent_iou(
294
+ gt_b: HeaderBlock,
295
+ pred_b: HeaderBlock,
296
+ gt_rows: int,
297
+ gt_cols: int,
298
+ pred_rows: int,
299
+ pred_cols: int,
300
+ ) -> float:
301
+ """IoU between normalised extents of two blocks in their respective tables."""
302
+ e1 = gt_b.extent(gt_rows, gt_cols)
303
+ e2 = pred_b.extent(pred_rows, pred_cols)
304
+ inter_r1, inter_c1 = max(e1[0], e2[0]), max(e1[1], e2[1])
305
+ inter_r2, inter_c2 = min(e1[2], e2[2]), min(e1[3], e2[3])
306
+ inter = max(0.0, inter_r2 - inter_r1) * max(0.0, inter_c2 - inter_c1)
307
+ area1 = (e1[2] - e1[0]) * (e1[3] - e1[1])
308
+ area2 = (e2[2] - e2[0]) * (e2[3] - e2[1])
309
+ union = area1 + area2 - inter
310
+ return inter / union if union > 0 else 0.0
311
+
312
+
313
+ # Small weight for extent IoU tiebreaker — enough to break ties in GriTS-Con
314
+ # but not enough to override a genuine content difference.
315
+ _EXTENT_TIEBREAK_WEIGHT = 1e-4
316
+
317
+
318
+ def _match_blocks(
319
+ gt_blocks: list[HeaderBlock],
320
+ pred_blocks: list[HeaderBlock],
321
+ gt_rows: int = 0,
322
+ gt_cols: int = 0,
323
+ pred_rows: int = 0,
324
+ pred_cols: int = 0,
325
+ ) -> tuple[dict[int, int], dict[tuple[int, int], float]]:
326
+ """Match GT blocks to pred blocks via GriTS-Con Hungarian matching.
327
+
328
+ When GriTS-Con scores tie, extent IoU (positional overlap) is used
329
+ as a tiebreaker so that blocks in similar positions are preferred.
330
+
331
+ Returns:
332
+ gt_to_pred: mapping from GT block index to pred block index
333
+ grits_scores: dict of (gt_idx, pred_idx) -> GriTS-Con f-score
334
+ for all pairs considered during matching
335
+ """
336
+ if not gt_blocks or not pred_blocks:
337
+ return {}, {}
338
+
339
+ n_gt = len(gt_blocks)
340
+ n_pred = len(pred_blocks)
341
+
342
+ # Infer table dimensions from block extents if not provided
343
+ if gt_rows == 0:
344
+ gt_rows = max(b.max_row for b in gt_blocks)
345
+ if gt_cols == 0:
346
+ gt_cols = max(b.max_col for b in gt_blocks)
347
+ if pred_rows == 0:
348
+ pred_rows = max(b.max_row for b in pred_blocks)
349
+ if pred_cols == 0:
350
+ pred_cols = max(b.max_col for b in pred_blocks)
351
+
352
+ cost = np.zeros((n_gt, n_pred))
353
+ grits_scores: dict[tuple[int, int], float] = {}
354
+ for i, gt_b in enumerate(gt_blocks):
355
+ gt_grid = _block_to_text_grid(gt_b)
356
+ for j, pred_b in enumerate(pred_blocks):
357
+ pred_grid = _block_to_text_grid(pred_b)
358
+ fscore, _, _, _ = factored_2dmss(gt_grid, pred_grid, _lcs_similarity)
359
+ grits_scores[(i, j)] = fscore
360
+ # Extent IoU tiebreaker: when GriTS-Con ties, prefer positional match
361
+ iou = _block_extent_iou(
362
+ gt_b,
363
+ pred_b,
364
+ gt_rows,
365
+ gt_cols,
366
+ pred_rows,
367
+ pred_cols,
368
+ )
369
+ cost[i, j] = -(fscore + _EXTENT_TIEBREAK_WEIGHT * iou)
370
+
371
+ row_ind, col_ind = linear_sum_assignment(cost)
372
+
373
+ gt_to_pred: dict[int, int] = {}
374
+ for r, c in zip(row_ind, col_ind, strict=True):
375
+ gt_to_pred[int(r)] = int(c)
376
+
377
+ return gt_to_pred, grits_scores
378
+
379
+
380
+ # ---------------------------------------------------------------------------
381
+ # Submetric implementations
382
+ # ---------------------------------------------------------------------------
383
+
384
+
385
+ _COMPOSITE_KEYS = [
386
+ "header_cell_count",
387
+ "header_grits",
388
+ "header_content_bag",
389
+ "header_block_extent",
390
+ "header_block_proximity",
391
+ "header_block_relative_direction",
392
+ "multilevel_header_depth",
393
+ "header_data_alignment",
394
+ ]
395
+
396
+
397
+ def _build_text_lookup(table_html: str) -> dict[tuple[int, int], str]:
398
+ """Build a (row, col) -> normalized_text lookup for ALL cells in a table.
399
+
400
+ Unlike _parse_header_cells which only returns <th>/thead cells, this
401
+ returns text for every cell (th and td) so we can check what text
402
+ appears at any grid position.
403
+ """
404
+ soup = BeautifulSoup(table_html, "lxml")
405
+ table = soup.find("table")
406
+ if not table:
407
+ return {}
408
+
409
+ rows = table.find_all("tr")
410
+ if not rows:
411
+ return {}
412
+
413
+ occupied: dict[tuple[int, int], bool] = {}
414
+ lookup: dict[tuple[int, int], str] = {}
415
+
416
+ for row_idx, row in enumerate(rows):
417
+ col_idx = 0
418
+ for cell in row.find_all(["td", "th"]):
419
+ while (row_idx, col_idx) in occupied:
420
+ col_idx += 1
421
+ rowspan = int(str(cell.get("rowspan", "1")))
422
+ colspan = int(str(cell.get("colspan", "1")))
423
+ text = _normalize_header_text(cell.get_text(strip=True))
424
+ for r in range(row_idx, row_idx + rowspan):
425
+ for c in range(col_idx, col_idx + colspan):
426
+ occupied[(r, c)] = True
427
+ lookup[(r, c)] = text
428
+ col_idx += colspan
429
+
430
+ return lookup
431
+
432
+
433
+ _GENEROUS_LCS_THRESHOLD = 0.8
434
+
435
+
436
+ def _is_bottom_left_block(block: HeaderBlock, num_rows: int) -> bool:
437
+ """Check if a header block is a rectangle in the bottom-left of the table.
438
+
439
+ A block is bottom-left if:
440
+ - It touches the leftmost column (min_col == 0)
441
+ - It has at least one cell in the bottom row (max_row == num_rows)
442
+ - It does NOT extend to the top row (min_row > 0)
443
+ - It is a perfect rectangle (no holes or irregular shape)
444
+ """
445
+ if not (block.min_col == 0 and block.max_row == num_rows and block.min_row > 0):
446
+ return False
447
+ # Check rectangularity: count occupied grid positions across all cells
448
+ occupied: set[tuple[int, int]] = set()
449
+ for cell in block.cells:
450
+ for r in range(cell.row, cell.row + cell.rowspan):
451
+ for c in range(cell.col, cell.col + cell.colspan):
452
+ occupied.add((r, c))
453
+ expected_area = (block.max_row - block.min_row) * (block.max_col - block.min_col)
454
+ return len(occupied) == expected_area
455
+
456
+
457
+ def _find_contiguous_groups(
458
+ positions: set[tuple[int, int]],
459
+ ) -> list[set[tuple[int, int]]]:
460
+ """Group positions into contiguous sets using 4-connected adjacency."""
461
+ if not positions:
462
+ return []
463
+
464
+ parent: dict[tuple[int, int], tuple[int, int]] = {p: p for p in positions}
465
+
466
+ def find(x: tuple[int, int]) -> tuple[int, int]:
467
+ while parent[x] != x:
468
+ parent[x] = parent[parent[x]]
469
+ x = parent[x]
470
+ return x
471
+
472
+ def union(a: tuple[int, int], b: tuple[int, int]) -> None:
473
+ ra, rb = find(a), find(b)
474
+ if ra != rb:
475
+ parent[ra] = rb
476
+
477
+ for r, c in positions:
478
+ for dr, dc in ((-1, 0), (1, 0), (0, -1), (0, 1)):
479
+ neighbor = (r + dr, c + dc)
480
+ if neighbor in positions:
481
+ union((r, c), neighbor)
482
+
483
+ groups: dict[tuple[int, int], set[tuple[int, int]]] = {}
484
+ for p in positions:
485
+ root = find(p)
486
+ groups.setdefault(root, set()).add(p)
487
+
488
+ return list(groups.values())
489
+
490
+
491
+ def _promote_cells_at_positions(pred_html: str, positions: set[tuple[int, int]]) -> str:
492
+ """Promote <td> cells at specific grid positions to <th>."""
493
+ soup = BeautifulSoup(pred_html, "lxml")
494
+ table = soup.find("table")
495
+ if not table:
496
+ return pred_html
497
+
498
+ rows = table.find_all("tr")
499
+ occupied: dict[tuple[int, int], bool] = {}
500
+
501
+ for row_idx, row in enumerate(rows):
502
+ col_idx = 0
503
+ for cell in row.find_all(["td", "th"]):
504
+ while (row_idx, col_idx) in occupied:
505
+ col_idx += 1
506
+ rowspan = int(str(cell.get("rowspan", "1")))
507
+ colspan = int(str(cell.get("colspan", "1")))
508
+ # Check if any of this cell's positions should be promoted
509
+ should_promote = False
510
+ for r in range(row_idx, row_idx + rowspan):
511
+ for c in range(col_idx, col_idx + colspan):
512
+ occupied[(r, c)] = True
513
+ if (r, c) in positions:
514
+ should_promote = True
515
+ if should_promote and cell.name == "td":
516
+ cell.name = "th"
517
+ col_idx += colspan
518
+
519
+ return str(table)
520
+
521
+
522
+ def _promote_bottom_left_to_header(gt_html: str, pred_html: str, threshold: float = _GENEROUS_LCS_THRESHOLD) -> str:
523
+ """Promote pred cells to <th> where GT has a bottom-left header block.
524
+
525
+ For each bottom-left header block in the GT, check if the corresponding
526
+ cells in the pred table have similar text. Promote matching pred cells
527
+ that form a contiguous block to <th>.
528
+ """
529
+ gt_cells, gt_num_rows, gt_num_cols = _parse_header_cells(gt_html)
530
+ if not gt_cells or gt_num_rows == 0:
531
+ return pred_html
532
+
533
+ gt_blocks = _find_header_blocks(gt_cells)
534
+ bl_blocks = [b for b in gt_blocks if _is_bottom_left_block(b, gt_num_rows)]
535
+ if not bl_blocks:
536
+ return pred_html
537
+
538
+ pred_text_lookup = _build_text_lookup(pred_html)
539
+ if not pred_text_lookup:
540
+ return pred_html
541
+
542
+ # Check if pred already has headers in the bottom-left region
543
+ pred_cells, pred_num_rows, _ = _parse_header_cells(pred_html)
544
+ pred_blocks = _find_header_blocks(pred_cells)
545
+ pred_bl_blocks = [b for b in pred_blocks if _is_bottom_left_block(b, pred_num_rows)]
546
+ if pred_bl_blocks:
547
+ return pred_html # pred already has bottom-left headers
548
+
549
+ # Collect positions to promote
550
+ positions_to_promote: set[tuple[int, int]] = set()
551
+ for block in bl_blocks:
552
+ matching_positions: set[tuple[int, int]] = set()
553
+ for cell in block.cells:
554
+ gt_text = cell.text
555
+ for r in range(cell.row, cell.row + cell.rowspan):
556
+ for c in range(cell.col, cell.col + cell.colspan):
557
+ pred_text = pred_text_lookup.get((r, c), "")
558
+ if _lcs_similarity(gt_text, pred_text) >= threshold:
559
+ matching_positions.add((r, c))
560
+
561
+ # Check contiguity: find the contiguous group that contains
562
+ # the bottom-left cell of the *pred* table (not the GT table,
563
+ # since the pred may be truncated).
564
+ if matching_positions:
565
+ contiguous_groups = _find_contiguous_groups(matching_positions)
566
+ pred_bottom_left = (pred_num_rows - 1, 0)
567
+ for group in contiguous_groups:
568
+ if pred_bottom_left in group:
569
+ positions_to_promote.update(group)
570
+ break
571
+
572
+ if not positions_to_promote:
573
+ return pred_html
574
+
575
+ # Promote the cells at those positions
576
+ return _promote_cells_at_positions(pred_html, positions_to_promote)
577
+
578
+
579
+ def _header_data_alignment_score(
580
+ gt_cells: list[HeaderCell],
581
+ pred_text_lookup: dict[tuple[int, int], str],
582
+ row_map: dict[int, int],
583
+ col_map: dict[int, int],
584
+ ) -> float:
585
+ """Submetric 10: header-data alignment via GriTS grid mapping.
586
+
587
+ For each GT header cell at anchor (row, col) with normalized text T,
588
+ maps to the prediction grid via (row_map[row], col_map[col]) and
589
+ checks whether the text at that position matches T.
590
+
591
+ Returns fraction of GT header cells whose text matches at the
592
+ aligned position. Returns 1.0 when GT has no headers.
593
+
594
+ The text check is leader-insensitive — a trailing dot/period run is
595
+ decoration, not content. See ``utils.cells_match_leader_insensitive``.
596
+ """
597
+ if not gt_cells:
598
+ return 1.0
599
+ if not row_map or not col_map:
600
+ return 0.0
601
+
602
+ hits = 0
603
+ for gc in gt_cells:
604
+ mapped_r = row_map.get(gc.row)
605
+ mapped_c = col_map.get(gc.col)
606
+ if mapped_r is not None and mapped_c is not None:
607
+ pred_text = pred_text_lookup.get((mapped_r, mapped_c), "")
608
+ if cells_match_leader_insensitive(gc.text, pred_text):
609
+ hits += 1
610
+
611
+ return hits / len(gt_cells)
612
+
613
+
614
+ def _header_data_alignment_score_fallback(
615
+ gt_html: str,
616
+ pred_html: str,
617
+ gt_cells: list[HeaderCell],
618
+ pred_text_lookup: dict[tuple[int, int], str],
619
+ ) -> float:
620
+ """Compute header_data_alignment without pre-computed GriTS alignment.
621
+
622
+ Builds text grids from the HTML and runs _align_2d_outer to get
623
+ row/col mappings, then delegates to _header_data_alignment_score.
624
+ """
625
+ from parse_bench.evaluation.metrics.parse.grits_metric import (
626
+ _align_2d_outer,
627
+ _lcs_similarity,
628
+ cells_to_grid,
629
+ html_to_cells,
630
+ )
631
+
632
+ if not gt_cells:
633
+ return 1.0
634
+
635
+ true_cells = html_to_cells(gt_html)
636
+ pred_cells_parsed = html_to_cells(pred_html)
637
+ if not true_cells or not pred_cells_parsed:
638
+ return 0.0
639
+
640
+ true_text = np.array(cells_to_grid(true_cells, key="cell_text"), dtype=object)
641
+ pred_text = np.array(cells_to_grid(pred_cells_parsed, key="cell_text"), dtype=object)
642
+
643
+ # Compute reward lookup (same as factored_2dmss)
644
+ pre_computed: dict[tuple[int, int, int, int], float] = {}
645
+ transpose_rewards: dict[tuple[int, int, int, int], float] = {}
646
+ for trow in range(true_text.shape[0]):
647
+ for tcol in range(true_text.shape[1]):
648
+ for prow in range(pred_text.shape[0]):
649
+ for pcol in range(pred_text.shape[1]):
650
+ reward = _lcs_similarity(true_text[trow, tcol], pred_text[prow, pcol])
651
+ pre_computed[(trow, tcol, prow, pcol)] = reward
652
+ transpose_rewards[(tcol, trow, pcol, prow)] = reward
653
+
654
+ true_row_nums, pred_row_nums, _ = _align_2d_outer(true_text.shape[:2], pred_text.shape[:2], pre_computed)
655
+ true_col_nums, pred_col_nums, _ = _align_2d_outer(
656
+ true_text.shape[:2][::-1], pred_text.shape[:2][::-1], transpose_rewards
657
+ )
658
+
659
+ row_map = dict(zip(true_row_nums, pred_row_nums, strict=True))
660
+ col_map = dict(zip(true_col_nums, pred_col_nums, strict=True))
661
+
662
+ return _header_data_alignment_score(gt_cells, pred_text_lookup, row_map, col_map)
663
+
664
+
665
+ def _header_cell_count_score(
666
+ gt_cells: list[HeaderCell],
667
+ pred_cells: list[HeaderCell],
668
+ ) -> float:
669
+ """Submetric 1: ratio-based cell count similarity.
670
+
671
+ Returns min(gt_count, pred_count) / max(gt_count, pred_count).
672
+ If both are 0, returns 1.0 (both agree there are no headers).
673
+ """
674
+ gt_n = len(gt_cells)
675
+ pred_n = len(pred_cells)
676
+ if gt_n == 0 and pred_n == 0:
677
+ return 1.0
678
+ if gt_n == 0 or pred_n == 0:
679
+ return 0.0
680
+ return min(gt_n, pred_n) / max(gt_n, pred_n)
681
+
682
+
683
+ def _header_grits_score(
684
+ gt_blocks: list[HeaderBlock],
685
+ pred_blocks: list[HeaderBlock],
686
+ gt_to_pred: dict[int, int],
687
+ grits_scores: dict[tuple[int, int], float],
688
+ ) -> float:
689
+ """Submetric 2: GriTS-Con on matched header blocks.
690
+
691
+ Uses the shared block matching. Unmatched GT blocks score 0. Extra
692
+ pred blocks are penalised by averaging over max(n_gt, n_pred).
693
+ """
694
+ if not gt_blocks and not pred_blocks:
695
+ return 1.0
696
+ if not gt_blocks or not pred_blocks:
697
+ return 0.0
698
+
699
+ n_gt = len(gt_blocks)
700
+ n_pred = len(pred_blocks)
701
+
702
+ matched_scores = [grits_scores[(gi, pi)] for gi, pi in gt_to_pred.items()]
703
+
704
+ denom = max(n_gt, n_pred)
705
+ return sum(matched_scores) / denom
706
+
707
+
708
+ def _header_content_bag_score(
709
+ gt_cells: list[HeaderCell],
710
+ pred_cells: list[HeaderCell],
711
+ ) -> float:
712
+ """Submetric 3: bag-of-cells exact content overlap.
713
+
714
+ For each GT header cell text, check if any pred header cell text
715
+ matches exactly (after formatting normalization).
716
+ Score = matched / total GT.
717
+
718
+ "Exactly" is leader-insensitive: a header that differs only in a trailing
719
+ dot/period run ("no." vs "no") is the same header. See
720
+ ``utils.cells_match_leader_insensitive``.
721
+ """
722
+ if not gt_cells and not pred_cells:
723
+ return 1.0
724
+ if not gt_cells or not pred_cells:
725
+ return 0.0
726
+
727
+ pred_texts = [c.text for c in pred_cells]
728
+ used = [False] * len(pred_texts)
729
+ matched = 0
730
+
731
+ for gt_cell in gt_cells:
732
+ for j, pt in enumerate(pred_texts):
733
+ if used[j]:
734
+ continue
735
+ if cells_match_leader_insensitive(gt_cell.text, pt):
736
+ used[j] = True
737
+ matched += 1
738
+ break
739
+
740
+ return matched / len(gt_cells)
741
+
742
+
743
+ def _header_perfect_score(
744
+ gt_cells: list[HeaderCell],
745
+ pred_cells: list[HeaderCell],
746
+ ) -> float:
747
+ """Submetric 4: binary exact structure match.
748
+
749
+ Returns 1.0 iff the header cells have the same count and each
750
+ (text, row, col, rowspan, colspan) matches exactly (in sorted order).
751
+
752
+ Geometry must match exactly; the text is compared leader-insensitively,
753
+ the same way ``header_content_bag`` and ``header_data_alignment`` compare
754
+ it, so the three submetrics cannot disagree about whether a trailing dot
755
+ run makes two headers different. Sorting still uses the raw text, which
756
+ is inert here: (row, col, rowspan, colspan) is already unique per cell.
757
+ """
758
+ if not gt_cells and not pred_cells:
759
+ return 1.0
760
+ if len(gt_cells) != len(pred_cells):
761
+ return 0.0
762
+
763
+ def _key(c: HeaderCell) -> tuple[int, int, int, int, str]:
764
+ return (c.row, c.col, c.rowspan, c.colspan, c.text)
765
+
766
+ gt_sorted = sorted(gt_cells, key=_key)
767
+ pred_sorted = sorted(pred_cells, key=_key)
768
+
769
+ for g, p in zip(gt_sorted, pred_sorted, strict=True):
770
+ if _key(g)[:4] != _key(p)[:4] or not cells_match_leader_insensitive(g.text, p.text):
771
+ return 0.0
772
+ return 1.0
773
+
774
+
775
+ def _header_block_extent_score(
776
+ gt_blocks: list[HeaderBlock],
777
+ pred_blocks: list[HeaderBlock],
778
+ gt_to_pred: dict[int, int],
779
+ table_rows_gt: int,
780
+ table_cols_gt: int,
781
+ table_rows_pred: int,
782
+ table_cols_pred: int,
783
+ ) -> float:
784
+ """Submetric 7: header block location/extent similarity.
785
+
786
+ For each matched GT-pred block pair, computes IoU of their normalised
787
+ extents within their respective tables. Averages over max(n_gt, n_pred).
788
+ """
789
+ if not gt_blocks and not pred_blocks:
790
+ return 1.0
791
+ if not gt_blocks or not pred_blocks:
792
+ return 0.0
793
+
794
+ n_gt = len(gt_blocks)
795
+ n_pred = len(pred_blocks)
796
+
797
+ def _extent_iou(
798
+ gt_b: HeaderBlock,
799
+ pred_b: HeaderBlock,
800
+ ) -> float:
801
+ e1 = gt_b.extent(table_rows_gt, table_cols_gt)
802
+ e2 = pred_b.extent(table_rows_pred, table_cols_pred)
803
+
804
+ # IoU on normalised rectangles
805
+ r1, c1, r2, c2 = e1
806
+ r3, c3, r4, c4 = e2
807
+
808
+ inter_r1 = max(r1, r3)
809
+ inter_c1 = max(c1, c3)
810
+ inter_r2 = min(r2, r4)
811
+ inter_c2 = min(c2, c4)
812
+
813
+ inter_area = max(0.0, inter_r2 - inter_r1) * max(0.0, inter_c2 - inter_c1)
814
+ area1 = (r2 - r1) * (c2 - c1)
815
+ area2 = (r4 - r3) * (c4 - c3)
816
+ union_area = area1 + area2 - inter_area
817
+
818
+ if union_area <= 0:
819
+ return 0.0
820
+ return inter_area / union_area
821
+
822
+ total = 0.0
823
+ for gi, pi in gt_to_pred.items():
824
+ total += _extent_iou(gt_blocks[gi], pred_blocks[pi])
825
+
826
+ denom = max(n_gt, n_pred)
827
+ return total / denom
828
+
829
+
830
+ def _block_edge_vector(a: HeaderBlock, b: HeaderBlock) -> tuple[float, float]:
831
+ """Signed vector from nearest edge of *a* to nearest edge of *b*, in cells.
832
+
833
+ Returns (dr, dc) where positive dr means *b* is below *a* and
834
+ positive dc means *b* is to the right of *a*. Components are zero
835
+ when the blocks overlap along that axis.
836
+ """
837
+ # Row component (signed gap)
838
+ if b.min_row >= a.max_row:
839
+ dr = float(b.min_row - a.max_row)
840
+ elif a.min_row >= b.max_row:
841
+ dr = -float(a.min_row - b.max_row)
842
+ else:
843
+ dr = 0.0
844
+
845
+ # Col component (signed gap)
846
+ if b.min_col >= a.max_col:
847
+ dc = float(b.min_col - a.max_col)
848
+ elif a.min_col >= b.max_col:
849
+ dc = -float(a.min_col - b.max_col)
850
+ else:
851
+ dc = 0.0
852
+
853
+ return (dr, dc)
854
+
855
+
856
+ def _block_edge_distance(a: HeaderBlock, b: HeaderBlock) -> float:
857
+ """Shortest Euclidean distance between edges/corners of two blocks, in cells.
858
+
859
+ If blocks overlap or are adjacent, returns 0.
860
+ """
861
+ dr, dc = _block_edge_vector(a, b)
862
+ return float((dr**2 + dc**2) ** 0.5)
863
+
864
+
865
+ def _direction_similarity(
866
+ gt_dr: float,
867
+ gt_dc: float,
868
+ pred_dr: float,
869
+ pred_dc: float,
870
+ ) -> float:
871
+ """Cosine similarity between two edge vectors, mapped to [0, 1].
872
+
873
+ Returns (cos_sim + 1) / 2 so that:
874
+ - parallel vectors → 1.0
875
+ - perpendicular → 0.5
876
+ - opposite → 0.0
877
+
878
+ If either vector is zero-length, returns 1.0 (co-located blocks,
879
+ direction is irrelevant).
880
+ """
881
+ gt_mag = (gt_dr**2 + gt_dc**2) ** 0.5
882
+ pred_mag = (pred_dr**2 + pred_dc**2) ** 0.5
883
+ if gt_mag < 1e-9 or pred_mag < 1e-9:
884
+ return 1.0
885
+ cos_sim = (gt_dr * pred_dr + gt_dc * pred_dc) / (gt_mag * pred_mag)
886
+ # Clamp for floating-point safety
887
+ cos_sim = max(-1.0, min(1.0, cos_sim))
888
+ return float((cos_sim + 1.0) / 2.0)
889
+
890
+
891
+ def _header_block_relative_position_score(
892
+ gt_blocks: list[HeaderBlock],
893
+ pred_blocks: list[HeaderBlock],
894
+ gt_to_pred: dict[int, int],
895
+ table_rows_gt: int,
896
+ table_cols_gt: int,
897
+ table_rows_pred: int,
898
+ table_cols_pred: int,
899
+ ) -> tuple[float, float]:
900
+ """Proximity and direction scores for header block pairs.
901
+
902
+ For every pair of matched GT blocks, computes:
903
+ - **proximity**: similarity of nearest-edge distances in cell units
904
+ ``1 - |dist_gt - dist_pred| / max(dist_gt, dist_pred)``
905
+ - **direction**: cosine similarity of the signed edge vectors, mapped
906
+ to [0, 1] via ``(cos + 1) / 2``
907
+
908
+ The denominator is ``max(gt_pairs, pred_pairs)`` so that extra
909
+ predicted blocks are penalised (unmatched pairs contribute 0).
910
+
911
+ Returns (proximity, direction) averages.
912
+ If <=1 block on both sides with the same count, returns (1.0, 1.0).
913
+ If counts differ (0 vs 1), returns (0.0, 0.0).
914
+ """
915
+ n_gt = len(gt_blocks)
916
+ n_pred = len(pred_blocks)
917
+
918
+ if n_gt <= 1 and n_pred <= 1:
919
+ # Both sides have ≤1 block — but if counts differ (0 vs 1)
920
+ # that is a mismatch, not vacuous agreement.
921
+ if n_gt != n_pred:
922
+ return 0.0, 0.0
923
+ return 1.0, 1.0
924
+ if n_pred == 0:
925
+ return 0.0, 0.0
926
+
927
+ matched_gt_indices = sorted(gt_to_pred.keys())
928
+ prox_scores: list[float] = []
929
+ dir_scores: list[float] = []
930
+
931
+ for idx_a in range(len(matched_gt_indices)):
932
+ for idx_b in range(idx_a + 1, len(matched_gt_indices)):
933
+ gi_a = matched_gt_indices[idx_a]
934
+ gi_b = matched_gt_indices[idx_b]
935
+ pi_a = gt_to_pred[gi_a]
936
+ pi_b = gt_to_pred[gi_b]
937
+
938
+ gt_dr, gt_dc = _block_edge_vector(gt_blocks[gi_a], gt_blocks[gi_b])
939
+ pred_dr, pred_dc = _block_edge_vector(pred_blocks[pi_a], pred_blocks[pi_b])
940
+
941
+ gt_dist = (gt_dr**2 + gt_dc**2) ** 0.5
942
+ pred_dist = (pred_dr**2 + pred_dc**2) ** 0.5
943
+
944
+ max_dist = max(gt_dist, pred_dist)
945
+ if max_dist < 1e-9:
946
+ prox_scores.append(1.0)
947
+ else:
948
+ prox_scores.append(1.0 - abs(gt_dist - pred_dist) / max_dist)
949
+
950
+ dir_scores.append(_direction_similarity(gt_dr, gt_dc, pred_dr, pred_dc))
951
+
952
+ gt_pairs = n_gt * (n_gt - 1) // 2
953
+ pred_pairs = n_pred * (n_pred - 1) // 2
954
+ total_pairs = max(gt_pairs, pred_pairs)
955
+
956
+ if total_pairs == 0:
957
+ return 1.0, 1.0
958
+
959
+ # Sum matched pair scores and divide by total (unmatched pairs contribute 0)
960
+ sum_prox = sum(prox_scores)
961
+ sum_dir = sum(dir_scores)
962
+ avg_prox = sum_prox / total_pairs
963
+ avg_dir = sum_dir / total_pairs
964
+ return avg_prox, avg_dir
965
+
966
+
967
+ # ---------------------------------------------------------------------------
968
+ # Header hierarchy depth
969
+ # ---------------------------------------------------------------------------
970
+
971
+
972
+ def _header_hierarchy_depth(cells: list[HeaderCell]) -> int:
973
+ """Compute the depth of the header hierarchy.
974
+
975
+ The depth is the number of distinct levels in the header tree.
976
+ A cell with ``colspan > 1`` (or ``rowspan > 1``) is a parent that
977
+ groups child cells underneath (or beside) it.
978
+
979
+ The algorithm assigns each header cell to a *level* by tracing how
980
+ many ancestor cells span over it from rows above. A cell at
981
+ ``(row, col)`` is a child of the innermost cell in a prior row whose
982
+ column span covers ``col``. The depth is the maximum nesting depth
983
+ across all cells.
984
+
985
+ Returns 0 when there are no header cells.
986
+ """
987
+ if not cells:
988
+ return 0
989
+
990
+ # Build occupancy: for each grid position, record the cell that owns it
991
+ # sorted by row so we process top-down
992
+ # For each cell, compute its level in the hierarchy.
993
+ # A cell's level = 1 + level of its nearest ancestor (a cell in a
994
+ # prior row whose column span covers this cell's columns).
995
+ # Cells in the first header row (or with no ancestor) are at level 1.
996
+
997
+ # Sort cells by (row, col) for top-down processing
998
+ sorted_cells = sorted(cells, key=lambda c: (c.row, c.col))
999
+
1000
+ # Map grid positions to the cell that "owns" them and its level
1001
+ # For each column, track the stack of spanning cells
1002
+ # Simple approach: for each cell, find the deepest ancestor
1003
+
1004
+ # Build a grid mapping (row, col) -> cell index
1005
+ cell_index: dict[tuple[int, int], int] = {}
1006
+ for idx, cell in enumerate(sorted_cells):
1007
+ for r in range(cell.row, cell.row + cell.rowspan):
1008
+ for c in range(cell.col, cell.col + cell.colspan):
1009
+ cell_index[(r, c)] = idx
1010
+
1011
+ cell_level: dict[int, int] = {}
1012
+
1013
+ for idx, cell in enumerate(sorted_cells):
1014
+ # Look for an ancestor: a cell in a prior row whose column span
1015
+ # covers at least one column of this cell, and which is a
1016
+ # *different* cell (not the same cell spanning multiple rows).
1017
+ best_ancestor_level = 0
1018
+ # Check the row just above this cell's start row
1019
+ if cell.row > 0:
1020
+ # Look at all columns this cell spans
1021
+ ancestor_candidates: set[int] = set()
1022
+ for c in range(cell.col, cell.col + cell.colspan):
1023
+ for r in range(cell.row - 1, -1, -1):
1024
+ if (r, c) in cell_index:
1025
+ anc_idx = cell_index[(r, c)]
1026
+ if anc_idx != idx:
1027
+ ancestor_candidates.add(anc_idx)
1028
+ break # found the nearest cell above in this column
1029
+
1030
+ for anc_idx in ancestor_candidates:
1031
+ # The ancestor must have colspan > 1 OR be a spanning cell
1032
+ # that groups this cell — i.e. its column span must be
1033
+ # strictly wider than this cell's, OR it must be in a
1034
+ # different row. A same-width cell in a row above still
1035
+ # forms a parent if it spans across.
1036
+ # Actually: any cell in a row above that covers our columns
1037
+ # is a potential parent in the hierarchy.
1038
+ if anc_idx in cell_level:
1039
+ best_ancestor_level = max(best_ancestor_level, cell_level[anc_idx])
1040
+
1041
+ cell_level[idx] = best_ancestor_level + 1
1042
+
1043
+ return max(cell_level.values()) if cell_level else 0
1044
+
1045
+
1046
+ def _header_hierarchy_depth_score(
1047
+ gt_cells: list[HeaderCell],
1048
+ pred_cells: list[HeaderCell],
1049
+ ) -> float:
1050
+ """Submetric 9: header hierarchy depth similarity.
1051
+
1052
+ Compares the depth of the header hierarchy tree between GT and
1053
+ prediction using ``min(d_gt, d_pred) / max(d_gt, d_pred)``.
1054
+ Returns 1.0 when both have the same depth (including both 0).
1055
+ """
1056
+ gt_depth = _header_hierarchy_depth(gt_cells)
1057
+ pred_depth = _header_hierarchy_depth(pred_cells)
1058
+ if gt_depth == 0 and pred_depth == 0:
1059
+ return 1.0
1060
+ if gt_depth == 0 or pred_depth == 0:
1061
+ return 0.0
1062
+ return min(gt_depth, pred_depth) / max(gt_depth, pred_depth)
1063
+
1064
+
1065
+ # ---------------------------------------------------------------------------
1066
+ # Per-table composite computation
1067
+ # ---------------------------------------------------------------------------
1068
+
1069
+
1070
+ def compute_header_composite_for_table_pair(
1071
+ gt_html: str,
1072
+ pred_html: str,
1073
+ ) -> dict[str, float]:
1074
+ """Compute all header accuracy submetrics for a single table pair.
1075
+
1076
+ Returns a dict with keys for each submetric plus the composite
1077
+ header_composite_v3 (mean of all submetrics).
1078
+ """
1079
+ result = _detailed_header_composite_for_table_pair(gt_html, pred_html)
1080
+ return result[0]
1081
+
1082
+
1083
+ def _detailed_header_composite_for_table_pair(
1084
+ gt_html: str,
1085
+ pred_html: str,
1086
+ row_map: dict[int, int] | None = None,
1087
+ col_map: dict[int, int] | None = None,
1088
+ ) -> tuple[dict[str, float], dict[str, list[str]]]:
1089
+ """Compute header accuracy scores and rich per-submetric diagnostics.
1090
+
1091
+ Returns:
1092
+ (scores, details) where scores is a dict of metric_name -> float
1093
+ and details is a dict of metric_name -> list of detail strings.
1094
+ """
1095
+ gt_cells, gt_rows, gt_cols = _parse_header_cells(gt_html)
1096
+ pred_cells, pred_rows, pred_cols = _parse_header_cells(pred_html)
1097
+
1098
+ gt_blocks = _find_header_blocks(gt_cells)
1099
+ pred_blocks = _find_header_blocks(pred_cells)
1100
+
1101
+ gt_to_pred, block_grits_scores = _match_blocks(
1102
+ gt_blocks,
1103
+ pred_blocks,
1104
+ gt_rows,
1105
+ gt_cols,
1106
+ pred_rows,
1107
+ pred_cols,
1108
+ )
1109
+
1110
+ scores: dict[str, float] = {}
1111
+ details: dict[str, list[str]] = {}
1112
+
1113
+ gt_n = len(gt_cells)
1114
+ pred_n = len(pred_cells)
1115
+ gt_b = len(gt_blocks)
1116
+ pred_b = len(pred_blocks)
1117
+
1118
+ # --- 1. header_cell_count ---
1119
+ s = _header_cell_count_score(gt_cells, pred_cells)
1120
+ scores["header_cell_count"] = s
1121
+ if gt_n == 0 and pred_n == 0:
1122
+ details["header_cell_count"] = [f"{s:.3f} — no header cells"]
1123
+ else:
1124
+ details["header_cell_count"] = [
1125
+ f"{s:.3f} — {pred_n}/{gt_n} cells predicted (min/max = {min(gt_n, pred_n)}/{max(gt_n, pred_n)})"
1126
+ ]
1127
+
1128
+ # --- 2. header_grits ---
1129
+ s = _header_grits_score(gt_blocks, pred_blocks, gt_to_pred, block_grits_scores)
1130
+ scores["header_grits"] = s
1131
+ grits_lines: list[str] = [f"{s:.3f} — {pred_b}/{gt_b} blocks predicted"]
1132
+ for gi, pi in sorted(gt_to_pred.items()):
1133
+ gs = block_grits_scores.get((gi, pi), 0.0)
1134
+ gt_texts = sorted({c.text for c in gt_blocks[gi].cells if c.text})
1135
+ pred_texts = sorted({c.text for c in pred_blocks[pi].cells if c.text})
1136
+ if gt_texts or pred_texts:
1137
+ grits_lines.append(
1138
+ f" block {gi + 1}↔{pi + 1}: grits={gs:.3f}"
1139
+ f" | expected [{', '.join(repr(t) for t in gt_texts[:5])}]"
1140
+ f" predicted [{', '.join(repr(t) for t in pred_texts[:5])}]"
1141
+ )
1142
+ else:
1143
+ gb, pb = gt_blocks[gi], pred_blocks[pi]
1144
+ grits_lines.append(
1145
+ f" block {gi + 1}↔{pi + 1}: grits={gs:.3f}"
1146
+ f" | GT rows [{gb.min_row},{gb.max_row})"
1147
+ f" cols [{gb.min_col},{gb.max_col})"
1148
+ f" pred rows [{pb.min_row},{pb.max_row})"
1149
+ f" cols [{pb.min_col},{pb.max_col})"
1150
+ )
1151
+ # Flag unmatched GT blocks
1152
+ for gi in range(gt_b):
1153
+ if gi not in gt_to_pred:
1154
+ gt_texts = sorted({c.text for c in gt_blocks[gi].cells if c.text})
1155
+ if gt_texts:
1156
+ grits_lines.append(
1157
+ f" block {gi + 1}: unmatched | expected [{', '.join(repr(t) for t in gt_texts[:5])}]"
1158
+ )
1159
+ else:
1160
+ gb = gt_blocks[gi]
1161
+ grits_lines.append(
1162
+ f" block {gi + 1}: unmatched | rows [{gb.min_row},{gb.max_row}) cols [{gb.min_col},{gb.max_col})"
1163
+ )
1164
+ details["header_grits"] = grits_lines
1165
+
1166
+ # --- 3. header_content_bag (with matched/missing/unexpected) ---
1167
+ s = _header_content_bag_score(gt_cells, pred_cells)
1168
+ scores["header_content_bag"] = s
1169
+ # Recompute matching to get per-cell info
1170
+ pred_texts_list = [c.text for c in pred_cells]
1171
+ used = [False] * len(pred_texts_list)
1172
+ matched_texts: list[str] = []
1173
+ missing_texts: list[str] = []
1174
+ for gc in gt_cells:
1175
+ found = False
1176
+ for j, pt in enumerate(pred_texts_list):
1177
+ if used[j]:
1178
+ continue
1179
+ if gc.text == pt:
1180
+ used[j] = True
1181
+ matched_texts.append(gc.text)
1182
+ found = True
1183
+ break
1184
+ if not found:
1185
+ missing_texts.append(gc.text)
1186
+ unexpected_texts = [pred_texts_list[j] for j in range(len(pred_texts_list)) if not used[j]]
1187
+
1188
+ bag_lines: list[str] = [f"{s:.3f} — {len(matched_texts)}/{gt_n} expected cells found"]
1189
+ if missing_texts:
1190
+ bag_lines.append(f" missing: {list(missing_texts)}")
1191
+ if unexpected_texts:
1192
+ bag_lines.append(f" unexpected: {list(unexpected_texts)}")
1193
+ details["header_content_bag"] = bag_lines
1194
+
1195
+ # --- 4. perfect_header ---
1196
+ s = _header_perfect_score(gt_cells, pred_cells)
1197
+ scores["header_perfect"] = s
1198
+ if s == 1.0:
1199
+ details["header_perfect"] = [f"{s:.3f} — exact match ({gt_n} cells)"]
1200
+ elif gt_n != pred_n:
1201
+ details["header_perfect"] = [f"{s:.3f} — cell count differs ({gt_n} expected, {pred_n} predicted)"]
1202
+ else:
1203
+ # Same count but position/span mismatch — show first difference
1204
+ def _key(c: HeaderCell) -> tuple[int, int, int, int, str]:
1205
+ return (c.row, c.col, c.rowspan, c.colspan, c.text)
1206
+
1207
+ gt_sorted = sorted(gt_cells, key=_key)
1208
+ pred_sorted = sorted(pred_cells, key=_key)
1209
+ diffs: list[str] = []
1210
+ for g, p in zip(gt_sorted, pred_sorted, strict=True):
1211
+ gk, pk = _key(g), _key(p)
1212
+ if gk != pk:
1213
+ diffs.append(
1214
+ f" expected ({g.row},{g.col}) {g.rowspan}x{g.colspan} {g.text!r}"
1215
+ f" vs predicted ({p.row},{p.col}) {p.rowspan}x{p.colspan} {p.text!r}"
1216
+ )
1217
+ if len(diffs) >= 3:
1218
+ break
1219
+ struct_lines = [f"{s:.3f} — position/span mismatch ({gt_n} cells)"]
1220
+ struct_lines.extend(diffs)
1221
+ details["header_perfect"] = struct_lines
1222
+
1223
+ # --- 5. header_block_extent ---
1224
+ s = _header_block_extent_score(
1225
+ gt_blocks,
1226
+ pred_blocks,
1227
+ gt_to_pred,
1228
+ gt_rows,
1229
+ gt_cols,
1230
+ pred_rows,
1231
+ pred_cols,
1232
+ )
1233
+ scores["header_block_extent"] = s
1234
+ extent_lines: list[str] = [f"{s:.3f} — {len(gt_to_pred)}/{max(gt_b, pred_b)} blocks matched"]
1235
+ for gi, pi in sorted(gt_to_pred.items()):
1236
+ e1 = gt_blocks[gi].extent(gt_rows, gt_cols)
1237
+ e2 = pred_blocks[pi].extent(pred_rows, pred_cols)
1238
+ # Compute IoU inline
1239
+ inter_r1, inter_c1 = max(e1[0], e2[0]), max(e1[1], e2[1])
1240
+ inter_r2, inter_c2 = min(e1[2], e2[2]), min(e1[3], e2[3])
1241
+ inter = max(0.0, inter_r2 - inter_r1) * max(0.0, inter_c2 - inter_c1)
1242
+ a1 = (e1[2] - e1[0]) * (e1[3] - e1[1])
1243
+ a2 = (e2[2] - e2[0]) * (e2[3] - e2[1])
1244
+ union = a1 + a2 - inter
1245
+ iou = inter / union if union > 0 else 0.0
1246
+ extent_lines.append(
1247
+ f" block {gi + 1}↔{pi + 1}: IoU={iou:.3f}"
1248
+ f" | expected rows [{e1[0]:.2f},{e1[2]:.2f}] cols [{e1[1]:.2f},{e1[3]:.2f}]"
1249
+ f" predicted rows [{e2[0]:.2f},{e2[2]:.2f}] cols [{e2[1]:.2f},{e2[3]:.2f}]"
1250
+ )
1251
+ details["header_block_extent"] = extent_lines
1252
+
1253
+ # --- 7/8. header_block_proximity & header_block_relative_direction ---
1254
+ s_prox, s_dir = _header_block_relative_position_score(
1255
+ gt_blocks,
1256
+ pred_blocks,
1257
+ gt_to_pred,
1258
+ gt_rows,
1259
+ gt_cols,
1260
+ pred_rows,
1261
+ pred_cols,
1262
+ )
1263
+ scores["header_block_proximity"] = s_prox
1264
+ scores["header_block_relative_direction"] = s_dir
1265
+ if gt_b == pred_b and gt_b <= 1:
1266
+ details["header_block_proximity"] = [f"{s_prox:.3f} — ≤1 block, no pairwise distances"]
1267
+ details["header_block_relative_direction"] = [f"{s_dir:.3f} — ≤1 block, no pairwise distances"]
1268
+ elif gt_b <= 1 and pred_b <= 1:
1269
+ details["header_block_proximity"] = [
1270
+ f"{s_prox:.3f} — block count mismatch ({gt_b} expected, {pred_b} predicted)"
1271
+ ]
1272
+ details["header_block_relative_direction"] = [
1273
+ f"{s_dir:.3f} — block count mismatch ({gt_b} expected, {pred_b} predicted)"
1274
+ ]
1275
+ else:
1276
+ matched_indices = sorted(gt_to_pred.keys())
1277
+ prox_pair_details: list[str] = []
1278
+ dir_pair_details: list[str] = []
1279
+ for idx_a in range(len(matched_indices)):
1280
+ for idx_b in range(idx_a + 1, len(matched_indices)):
1281
+ gi_a, gi_b = matched_indices[idx_a], matched_indices[idx_b]
1282
+ pi_a, pi_b = gt_to_pred[gi_a], gt_to_pred[gi_b]
1283
+ gt_dr, gt_dc = _block_edge_vector(gt_blocks[gi_a], gt_blocks[gi_b])
1284
+ pred_dr, pred_dc = _block_edge_vector(pred_blocks[pi_a], pred_blocks[pi_b])
1285
+ gt_dist = (gt_dr**2 + gt_dc**2) ** 0.5
1286
+ pred_dist = (pred_dr**2 + pred_dc**2) ** 0.5
1287
+ max_dist = max(gt_dist, pred_dist)
1288
+ pair_prox = 1.0 if max_dist < 1e-9 else 1.0 - abs(gt_dist - pred_dist) / max_dist
1289
+ pair_dir = _direction_similarity(gt_dr, gt_dc, pred_dr, pred_dc)
1290
+ prox_pair_details.append(
1291
+ f" blocks {gi_a + 1}↔{gi_b + 1}: proximity={pair_prox:.3f}"
1292
+ f" | gt_dist={gt_dist:.1f} pred_dist={pred_dist:.1f}"
1293
+ )
1294
+ dir_pair_details.append(
1295
+ f" blocks {gi_a + 1}↔{gi_b + 1}: direction={pair_dir:.3f}"
1296
+ f" | gt_vec=({gt_dr:.1f},{gt_dc:.1f})"
1297
+ f" pred_vec=({pred_dr:.1f},{pred_dc:.1f})"
1298
+ )
1299
+ gt_pairs_count = gt_b * (gt_b - 1) // 2
1300
+ pred_pairs_count = pred_b * (pred_b - 1) // 2
1301
+ total_pairs_count = max(gt_pairs_count, pred_pairs_count)
1302
+ prox_lines = [f"{s_prox:.3f} — {len(prox_pair_details)} matched / {total_pairs_count} total pair(s)"]
1303
+ dir_lines = [f"{s_dir:.3f} — {len(dir_pair_details)} matched / {total_pairs_count} total pair(s)"]
1304
+ if pred_b > gt_b:
1305
+ extra_msg = f" extra pred blocks: {pred_b - gt_b} (penalised via denominator)"
1306
+ prox_lines.append(extra_msg)
1307
+ dir_lines.append(extra_msg)
1308
+ prox_lines.extend(prox_pair_details[:5])
1309
+ dir_lines.extend(dir_pair_details[:5])
1310
+ details["header_block_proximity"] = prox_lines
1311
+ details["header_block_relative_direction"] = dir_lines
1312
+
1313
+ # --- 9. multilevel_header_depth ---
1314
+ s = _header_hierarchy_depth_score(gt_cells, pred_cells)
1315
+ scores["multilevel_header_depth"] = s
1316
+ gt_depth = _header_hierarchy_depth(gt_cells)
1317
+ pred_depth = _header_hierarchy_depth(pred_cells)
1318
+ if gt_depth == 0 and pred_depth == 0:
1319
+ details["multilevel_header_depth"] = [f"{s:.3f} — no header hierarchy"]
1320
+ else:
1321
+ details["multilevel_header_depth"] = [f"{s:.3f} — expected depth {gt_depth}, predicted depth {pred_depth}"]
1322
+
1323
+ # --- 10. header_data_alignment ---
1324
+ pred_text_lookup = _build_text_lookup(pred_html)
1325
+ has_grits_alignment = bool(row_map) and bool(col_map)
1326
+ if has_grits_alignment:
1327
+ assert row_map is not None and col_map is not None
1328
+ s = _header_data_alignment_score(gt_cells, pred_text_lookup, row_map, col_map)
1329
+ scores["header_data_alignment"] = s
1330
+ aligned = int(s * len(gt_cells)) if gt_cells else 0
1331
+ details["header_data_alignment"] = [
1332
+ f"{s:.3f} — {aligned}/{len(gt_cells)} GT headers aligned (via GriTS row/col mapping)"
1333
+ ]
1334
+ else:
1335
+ s = _header_data_alignment_score_fallback(
1336
+ gt_html,
1337
+ pred_html,
1338
+ gt_cells,
1339
+ pred_text_lookup,
1340
+ )
1341
+ scores["header_data_alignment"] = s
1342
+ aligned = int(s * len(gt_cells)) if gt_cells else 0
1343
+ details["header_data_alignment"] = [
1344
+ f"{s:.3f} — {aligned}/{len(gt_cells)} GT headers aligned (computed via standalone alignment, no GriTS data)"
1345
+ ]
1346
+
1347
+ # --- composite (mean of applicable _COMPOSITE_KEYS, excludes header_perfect) ---
1348
+ # Exclude submetrics that are trivially 1.0 because both GT and prediction
1349
+ # fall in the degenerate case (e.g. ≤1 block on both sides). When counts
1350
+ # differ (e.g. GT has 0 blocks but prediction has 2) the metric is kept so
1351
+ # that the mismatch is penalised.
1352
+ trivial_keys: set[str] = set()
1353
+ if gt_b == pred_b and gt_b <= 1:
1354
+ trivial_keys.add("header_block_proximity")
1355
+ trivial_keys.add("header_block_relative_direction")
1356
+
1357
+ applicable_keys = [k for k in _COMPOSITE_KEYS if k not in trivial_keys]
1358
+ if applicable_keys:
1359
+ scores["header_composite_v3"] = sum(scores[k] for k in applicable_keys) / len(applicable_keys)
1360
+ else:
1361
+ # All submetrics trivial — perfect by default
1362
+ scores["header_composite_v3"] = 1.0
1363
+
1364
+ sub_strs = [f"{k}={scores[k]:.3f}" for k in _COMPOSITE_KEYS]
1365
+ skipped_strs = [f"{k} (trivial, excluded)" for k in _COMPOSITE_KEYS if k in trivial_keys]
1366
+ composite_lines = [
1367
+ f"{scores['header_composite_v3']:.3f} — " + ", ".join(sub_strs),
1368
+ f"{pred_n}/{gt_n} cells, {pred_b}/{gt_b} blocks",
1369
+ ]
1370
+ if skipped_strs:
1371
+ composite_lines.append(f"excluded from composite: {', '.join(skipped_strs)}")
1372
+ details["header_composite_v3"] = composite_lines
1373
+
1374
+ return scores, details
1375
+
1376
+
1377
+ # ---------------------------------------------------------------------------
1378
+ # Metric class (multi-table document level)
1379
+ # ---------------------------------------------------------------------------
1380
+
1381
+ # All submetric keys emitted by this metric
1382
+ SUBMETRIC_KEYS = [
1383
+ "header_cell_count",
1384
+ "header_grits",
1385
+ "header_content_bag",
1386
+ "header_perfect",
1387
+ "header_block_extent",
1388
+ "header_block_proximity",
1389
+ "header_block_relative_direction",
1390
+ "multilevel_header_depth",
1391
+ "header_data_alignment",
1392
+ "header_composite_v3",
1393
+ ]
1394
+
1395
+
1396
+ class HeaderAccuracyMetric(Metric):
1397
+ """Header accuracy metric for comparing HTML tables in markdown content.
1398
+
1399
+ Computes header accuracy between expected and actual HTML tables.
1400
+ Table-level matching can be provided externally (e.g. from GriTS) via
1401
+ the ``table_pairs`` parameter, or computed internally as a fallback.
1402
+ """
1403
+
1404
+ @property
1405
+ def name(self) -> str:
1406
+ return "header_composite_v3"
1407
+
1408
+ def compute( # type: ignore[override]
1409
+ self,
1410
+ expected: str,
1411
+ actual: str,
1412
+ table_pairs: list[tuple[str, str]] | None = None,
1413
+ table_alignments: list[tuple[dict[int, int], dict[int, int]]] | None = None,
1414
+ **kwargs: Any,
1415
+ ) -> list[MetricValue]:
1416
+ """Compute header accuracy scores between expected and actual content.
1417
+
1418
+ Args:
1419
+ expected: Full document markdown/HTML with ground-truth tables.
1420
+ actual: Full document markdown/HTML with predicted tables.
1421
+ table_pairs: Optional pre-matched list of (gt_html, pred_html)
1422
+ table pairs (e.g. from GriTS matching). If None, tables are
1423
+ extracted and matched internally via Hungarian on the
1424
+ overall header_composite_v3 score.
1425
+ table_alignments: Optional per-table GriTS row/col alignment
1426
+ maps as [(row_map, col_map), ...]. Used for the
1427
+ header_data_alignment submetric.
1428
+
1429
+ Returns a list of MetricValues: one for the overall header_composite_v3
1430
+ and one for each submetric.
1431
+ """
1432
+ if table_pairs is not None:
1433
+ return self._compute_from_pairs(table_pairs, table_alignments)
1434
+
1435
+ # Fallback: extract and match tables internally
1436
+ expected_tables = extract_html_tables(expected)
1437
+ actual_tables = extract_html_tables(actual)
1438
+
1439
+ if not expected_tables:
1440
+ meta: dict[str, Any] = {
1441
+ "tables_found_expected": 0,
1442
+ "tables_found_actual": len(actual_tables),
1443
+ }
1444
+ return [MetricValue(metric_name="header_composite_v3", value=0.0, metadata=meta)]
1445
+
1446
+ if not actual_tables:
1447
+ meta = {
1448
+ "tables_found_expected": len(expected_tables),
1449
+ "tables_found_actual": 0,
1450
+ }
1451
+ return [MetricValue(metric_name="header_composite_v3", value=0.0, metadata=meta)]
1452
+
1453
+ n_gt = len(expected_tables)
1454
+ n_pred = len(actual_tables)
1455
+
1456
+ # Compute all pairwise scores
1457
+ pair_scores: dict[tuple[int, int], dict[str, float]] = {}
1458
+ for i, gt_t in enumerate(expected_tables):
1459
+ for j, pred_t in enumerate(actual_tables):
1460
+ pair_scores[(i, j)] = compute_header_composite_for_table_pair(gt_t, pred_t)
1461
+
1462
+ # Hungarian matching on overall header_composite_v3
1463
+ cost = np.zeros((n_gt, n_pred))
1464
+ for i in range(n_gt):
1465
+ for j in range(n_pred):
1466
+ cost[i, j] = -pair_scores[(i, j)]["header_composite_v3"]
1467
+
1468
+ row_ind, col_ind = linear_sum_assignment(cost)
1469
+
1470
+ # Build paired list
1471
+ pairs: list[tuple[str, str]] = []
1472
+ matched_gt: set[int] = set()
1473
+ for gt_idx, pred_idx in zip(row_ind, col_ind, strict=True):
1474
+ pairs.append((expected_tables[int(gt_idx)], actual_tables[int(pred_idx)]))
1475
+ matched_gt.add(int(gt_idx))
1476
+
1477
+ # Unmatched GT tables get paired with empty string
1478
+ for i in range(n_gt):
1479
+ if i not in matched_gt:
1480
+ pairs.append((expected_tables[i], ""))
1481
+
1482
+ return self._compute_from_pairs(pairs)
1483
+
1484
+ def _compute_from_pairs(
1485
+ self,
1486
+ table_pairs: list[tuple[str, str]],
1487
+ table_alignments: list[tuple[dict[int, int], dict[int, int]]] | None = None,
1488
+ ) -> list[MetricValue]:
1489
+ """Compute header accuracy from pre-matched table pairs."""
1490
+ if not table_pairs:
1491
+ return [MetricValue(metric_name="header_composite_v3", value=0.0, metadata={})]
1492
+
1493
+ accumulators: dict[str, list[float]] = {k: [] for k in SUBMETRIC_KEYS}
1494
+ per_table_details: list[dict[str, Any]] = []
1495
+ # Per-table rich detail strings keyed by submetric
1496
+ per_table_rich_details: list[dict[str, list[str]]] = []
1497
+ per_table_gt_cells: list[int] = []
1498
+ per_table_pred_cells: list[int] = []
1499
+ per_table_gt_blocks: list[int] = []
1500
+ per_table_pred_blocks: list[int] = []
1501
+
1502
+ for idx, (gt_html, pred_html) in enumerate(table_pairs):
1503
+ if not pred_html:
1504
+ # Unmatched GT table
1505
+ for k in SUBMETRIC_KEYS:
1506
+ accumulators[k].append(0.0)
1507
+ per_table_details.append({"table_pair_index": idx, **dict.fromkeys(SUBMETRIC_KEYS, 0.0)})
1508
+ gt_cells, _, _ = _parse_header_cells(gt_html)
1509
+ gt_n = len(gt_cells)
1510
+ gt_b = len(_find_header_blocks(gt_cells))
1511
+ per_table_gt_cells.append(gt_n)
1512
+ per_table_pred_cells.append(0)
1513
+ per_table_gt_blocks.append(gt_b)
1514
+ per_table_pred_blocks.append(0)
1515
+ # All-zero details for unmatched table
1516
+ unmatched_details: dict[str, list[str]] = {}
1517
+ for k in SUBMETRIC_KEYS:
1518
+ if k == "header_composite_v3":
1519
+ unmatched_details[k] = [f"0.000 — unmatched table ({gt_n} expected cells, 0 predicted)"]
1520
+ else:
1521
+ unmatched_details[k] = ["0.000 — no predicted table to compare"]
1522
+ per_table_rich_details.append(unmatched_details)
1523
+ continue
1524
+
1525
+ gt_cells_parsed, _, _ = _parse_header_cells(gt_html)
1526
+ pred_cells_parsed, _, _ = _parse_header_cells(pred_html)
1527
+ per_table_gt_cells.append(len(gt_cells_parsed))
1528
+ per_table_pred_cells.append(len(pred_cells_parsed))
1529
+ per_table_gt_blocks.append(len(_find_header_blocks(gt_cells_parsed)))
1530
+ per_table_pred_blocks.append(len(_find_header_blocks(pred_cells_parsed)))
1531
+
1532
+ if table_alignments and idx < len(table_alignments):
1533
+ pair_row_map, pair_col_map = table_alignments[idx]
1534
+ else:
1535
+ pair_row_map, pair_col_map = None, None
1536
+
1537
+ scores, rich_details = _detailed_header_composite_for_table_pair(
1538
+ gt_html,
1539
+ pred_html,
1540
+ row_map=pair_row_map,
1541
+ col_map=pair_col_map,
1542
+ )
1543
+ for k in SUBMETRIC_KEYS:
1544
+ accumulators[k].append(scores[k])
1545
+ per_table_details.append({"table_pair_index": idx, **scores})
1546
+ per_table_rich_details.append(rich_details)
1547
+
1548
+ shared_meta: dict[str, Any] = {
1549
+ "table_pairs": len(table_pairs),
1550
+ "per_table_details": per_table_details,
1551
+ "alignment_source": "grits" if table_alignments else "fallback",
1552
+ }
1553
+
1554
+ # Build human-readable detail strings per submetric
1555
+ total_gt = sum(per_table_gt_cells)
1556
+ total_pred = sum(per_table_pred_cells)
1557
+ summary_line = f"{total_gt} header cell(s) expected, {total_pred} predicted across {len(table_pairs)} table(s)"
1558
+
1559
+ submetric_details: dict[str, list[str]] = {}
1560
+ for k in SUBMETRIC_KEYS:
1561
+ lines: list[str] = [summary_line]
1562
+ for idx in range(len(table_pairs)):
1563
+ table_detail_lines = per_table_rich_details[idx].get(k, [])
1564
+ if len(table_pairs) > 1:
1565
+ # Prefix first line with table number
1566
+ if table_detail_lines:
1567
+ lines.append(f"Table {idx + 1}: {table_detail_lines[0]}")
1568
+ lines.extend(table_detail_lines[1:])
1569
+ else:
1570
+ td = per_table_details[idx]
1571
+ lines.append(f"Table {idx + 1}: {k}={td.get(k, 0.0):.3f}")
1572
+ else:
1573
+ # Single table — just append detail lines directly
1574
+ lines.extend(table_detail_lines)
1575
+ submetric_details[k] = lines
1576
+
1577
+ results: list[MetricValue] = []
1578
+ for k in SUBMETRIC_KEYS:
1579
+ vals = accumulators[k]
1580
+ avg = sum(vals) / len(vals) if vals else 0.0
1581
+ results.append(
1582
+ MetricValue(
1583
+ metric_name=k,
1584
+ value=avg,
1585
+ metadata=shared_meta,
1586
+ details=submetric_details.get(k, []),
1587
+ )
1588
+ )
1589
+
1590
+ return results
1591
+
1592
+
1593
+ # ---------------------------------------------------------------------------
1594
+ # Generous header normalization
1595
+ # ---------------------------------------------------------------------------
1596
+
1597
+
1598
+ def _promote_top_row_to_header(pred_html: str) -> str:
1599
+ """Convert all <td> cells in the top row of *pred_html* to <th> cells."""
1600
+ soup = BeautifulSoup(pred_html, "lxml")
1601
+ table = soup.find("table")
1602
+ if not table:
1603
+ return pred_html
1604
+ rows = table.find_all("tr")
1605
+ if not rows:
1606
+ return pred_html
1607
+ for cell in rows[0].find_all("td"):
1608
+ cell.name = "th"
1609
+ return str(table)
1610
+
1611
+
1612
+ def _apply_generous_header_normalization(gt_html: str, pred_html: str) -> str:
1613
+ """Promote pred's top row to a header if GT has headers and pred has none.
1614
+ Also promote bottom-left cells if GT has a bottom-left header block."""
1615
+ gt_cells, _, _ = _parse_header_cells(gt_html)
1616
+ if not gt_cells:
1617
+ return pred_html
1618
+ pred_cells, _, _ = _parse_header_cells(pred_html)
1619
+ if not pred_cells:
1620
+ pred_html = _promote_top_row_to_header(pred_html)
1621
+ # Always try bottom-left promotion (pred may have top headers but no
1622
+ # bottom-left headers, or we just promoted the top row above)
1623
+ pred_html = _promote_bottom_left_to_header(gt_html, pred_html)
1624
+ return pred_html
1625
+
1626
+
1627
+ class HeaderAccuracyMetricGenerous(HeaderAccuracyMetric):
1628
+ """Variant of HeaderAccuracyMetric with generous header normalization.
1629
+
1630
+ When the GT table has header cells but the prediction has none,
1631
+ the prediction's top row is promoted to a header before scoring.
1632
+ Only the composite score is emitted (as ``header_composite_v3_generous``);
1633
+ sub-metrics are omitted to avoid name collisions with the base metric.
1634
+ """
1635
+
1636
+ @property
1637
+ def name(self) -> str:
1638
+ return "exp_header_composite_v3_generous"
1639
+
1640
+ def compute( # type: ignore[override]
1641
+ self,
1642
+ expected: str,
1643
+ actual: str,
1644
+ table_pairs: list[tuple[str, str]] | None = None,
1645
+ table_alignments: list[tuple[dict[int, int], dict[int, int]]] | None = None,
1646
+ **kwargs: Any,
1647
+ ) -> list[MetricValue]:
1648
+ if table_pairs is not None:
1649
+ table_pairs = [
1650
+ (gt, _apply_generous_header_normalization(gt, pred) if pred else pred) for gt, pred in table_pairs
1651
+ ]
1652
+ results = super().compute(expected, actual, table_pairs, table_alignments, **kwargs)
1653
+ return self._rename_composite(results)
1654
+
1655
+ @staticmethod
1656
+ def _rename_composite(results: list[MetricValue]) -> list[MetricValue]:
1657
+ """Keep only the composite MetricValue, renamed to exp_header_composite_v3_generous."""
1658
+ for mv in results:
1659
+ if mv.metric_name == "header_composite_v3":
1660
+ mv.metric_name = "exp_header_composite_v3_generous"
1661
+ return [mv]
1662
+ return []