1pub use crate::assemble::TableGrid;
15use image::RgbImage;
16
17pub const SIDE: u32 = 448;
19#[allow(clippy::excessive_precision)]
22pub const MEAN: [f32; 3] = [0.94247851, 0.94254675, 0.94292611];
23#[allow(clippy::excessive_precision)]
24pub const STD: [f32; 3] = [0.17910956, 0.17940403, 0.17931663];
25pub const MAX_STEPS: usize = 1024;
27pub const MAX_ROW_TAGS: usize = 256;
34pub const EMBED_DIM: usize = 512;
36
37pub const START: i64 = 2;
39pub const END: i64 = 3;
40pub const ECEL: i64 = 4; pub const FCEL: i64 = 5; pub const LCEL: i64 = 6; pub const UCEL: i64 = 7; pub const XCEL: i64 = 8; pub const NL: i64 = 9; pub const CHED: i64 = 10; pub const RHED: i64 = 11; pub const SROW: i64 = 12; const CELL_TAGS: [i64; 6] = [FCEL, ECEL, XCEL, CHED, RHED, SROW];
51
52#[derive(Debug, Clone)]
56pub struct TableCell {
57 pub row: usize,
58 pub col: usize,
59 pub colspan: usize,
60 pub rowspan: usize,
61 pub tag: i64,
62 pub class: i64,
63 pub cx: f32,
64 pub cy: f32,
65 pub w: f32,
66 pub h: f32,
67}
68
69pub fn preprocess_input(img: &RgbImage) -> Vec<f32> {
74 let nn = (SIDE * SIDE) as usize;
75 let side = SIDE as usize;
76 let (sw, sh) = (img.width() as i32, img.height() as i32);
77 let sxr = sw as f32 / SIDE as f32;
78 let syr = sh as f32 / SIDE as f32;
79 let mut data = vec![0f32; 3 * nn];
80 for h in 0..side {
81 let fy = (h as f32 + 0.5) * syr - 0.5;
82 let wy = fy - fy.floor();
83 let y0c = (fy.floor() as i32).clamp(0, sh - 1) as u32;
84 let y1c = (fy.floor() as i32 + 1).clamp(0, sh - 1) as u32;
85 for w in 0..side {
86 let fx = (w as f32 + 0.5) * sxr - 0.5;
87 let wx = fx - fx.floor();
88 let x0c = (fx.floor() as i32).clamp(0, sw - 1) as u32;
89 let x1c = (fx.floor() as i32 + 1).clamp(0, sw - 1) as u32;
90 let p00 = img.get_pixel(x0c, y0c);
91 let p01 = img.get_pixel(x1c, y0c);
92 let p10 = img.get_pixel(x0c, y1c);
93 let p11 = img.get_pixel(x1c, y1c);
94 let idx = w * side + h; for c in 0..3 {
96 let top = p00[c] as f32 * (1.0 - wx) + p01[c] as f32 * wx;
97 let bot = p10[c] as f32 * (1.0 - wx) + p11[c] as f32 * wx;
98 let v = top * (1.0 - wy) + bot * wy;
99 data[c * nn + idx] = (v / 255.0 - MEAN[c]) / STD[c];
100 }
101 }
102 }
103 data
104}
105
106pub fn correct(raw: i64, prev_ucel: bool) -> i64 {
110 let mut tag = raw;
111 if tag == XCEL {
112 tag = LCEL;
113 }
114 if prev_ucel && tag == LCEL {
115 tag = FCEL;
116 }
117 tag
118}
119
120#[derive(Default)]
127pub struct BboxBook {
128 pub tags: Vec<i64>,
130 pub otsl: Vec<i64>,
132 pub hiddens: Vec<f32>,
134 pub n: usize,
136 pub merge: std::collections::HashMap<usize, i64>,
138 prev_ucel: bool,
139 skip: bool,
140 first_lcel: bool,
141 bbox_ind: usize,
142 cur_bbox_ind: usize,
143 row_len: usize,
145}
146
147impl BboxBook {
148 pub fn new() -> Self {
149 Self {
150 tags: vec![START],
151 skip: true, first_lcel: true,
153 ..Default::default()
154 }
155 }
156
157 pub fn step(&mut self, raw: i64, hidden: &[f32]) -> bool {
161 let tag = correct(raw, self.prev_ucel);
162 if tag == END {
163 return false;
164 }
165 if !self.skip && matches!(tag, FCEL | ECEL | CHED | RHED | SROW | NL | UCEL) {
167 self.hiddens.extend_from_slice(hidden);
168 self.n += 1;
169 if !self.first_lcel {
170 self.merge.insert(self.cur_bbox_ind, self.bbox_ind as i64);
171 }
172 self.bbox_ind += 1;
173 }
174 if tag != LCEL {
175 self.first_lcel = true;
176 } else if self.first_lcel {
177 self.hiddens.extend_from_slice(hidden);
178 self.n += 1;
179 self.first_lcel = false;
180 self.cur_bbox_ind = self.bbox_ind;
181 self.merge.insert(self.cur_bbox_ind, -1);
182 self.bbox_ind += 1;
183 }
184 self.skip = matches!(tag, NL | UCEL | XCEL);
185 self.prev_ucel = tag == UCEL;
186 self.otsl.push(tag);
187 self.tags.push(tag);
188 self.row_len = if tag == NL { 0 } else { self.row_len + 1 };
189 self.row_len < MAX_ROW_TAGS
190 }
191
192 pub fn runaway(&self) -> bool {
199 self.row_len >= MAX_ROW_TAGS
200 }
201}
202
203fn mergebboxes(b1: [f32; 4], b2: [f32; 4]) -> [f32; 4] {
206 let new_w = (b2[0] + b2[2] / 2.0) - (b1[0] - b1[2] / 2.0);
207 let new_h = (b2[1] + b2[3] / 2.0) - (b1[1] - b1[3] / 2.0);
208 let new_left = b1[0] - b1[2] / 2.0;
209 let new_top = (b2[1] - b2[3] / 2.0).min(b1[1] - b1[3] / 2.0);
210 [new_left + new_w / 2.0, new_top + new_h / 2.0, new_w, new_h]
211}
212
213pub fn merge_spans(
217 boxes: &[[f32; 4]],
218 classes: &[i64],
219 merge: &std::collections::HashMap<usize, i64>,
220) -> (Vec<[f32; 4]>, Vec<i64>) {
221 let skip: std::collections::HashSet<usize> = merge
222 .values()
223 .filter(|&&v| v >= 0)
224 .map(|&v| v as usize)
225 .collect();
226 let mut out = Vec::new();
227 let mut out_classes = Vec::new();
228 for (i, &b) in boxes.iter().enumerate() {
229 let class = classes.get(i).copied().unwrap_or(2);
230 if let Some(&j) = merge.get(&i) {
231 let partner = if j < 0 { boxes.len() - 1 } else { j as usize };
232 out.push(mergebboxes(b, boxes[partner.min(boxes.len() - 1)]));
233 out_classes.push(class);
234 } else if !skip.contains(&i) {
235 out.push(b);
236 out_classes.push(class);
237 }
238 }
239 (out, out_classes)
240}
241
242pub fn build_table_cells(otsl: &[i64], boxes: &[[f32; 4]], classes: &[i64]) -> Vec<TableCell> {
248 let mut grid: Vec<Vec<i64>> = vec![Vec::new()];
250 for &t in otsl {
251 if t == NL {
252 grid.push(Vec::new());
253 } else {
254 grid.last_mut().unwrap().push(t);
255 }
256 }
257 let mut cells = Vec::new();
258 let mut cell_id = 0usize;
259 for (r, row) in grid.iter().enumerate() {
260 for (c, &tag) in row.iter().enumerate() {
261 if !CELL_TAGS.contains(&tag) {
262 continue;
263 }
264 let mut colspan = 1;
265 while c + colspan < row.len() && matches!(row[c + colspan], LCEL | XCEL) {
266 colspan += 1;
267 }
268 let mut rowspan = 1;
269 while r + rowspan < grid.len()
270 && grid[r + rowspan]
271 .get(c)
272 .is_some_and(|&t| matches!(t, UCEL | XCEL))
273 {
274 rowspan += 1;
275 }
276 let b = boxes.get(cell_id).copied().unwrap_or([0.0; 4]);
277 let class = classes.get(cell_id).copied().unwrap_or(2);
279 cells.push(TableCell {
280 row: r,
281 col: c,
282 colspan,
283 rowspan,
284 tag,
285 class,
286 cx: b[0],
287 cy: b[1],
288 w: b[2],
289 h: b[3],
290 });
291 cell_id += 1;
292 }
293 }
294 cells
295}
296
297pub fn argmax(v: &[f32]) -> usize {
301 v.iter()
302 .enumerate()
303 .max_by(|a, b| a.1.total_cmp(b.1))
304 .map(|(i, _)| i)
305 .unwrap_or(0)
306}
307
308use crate::pdfium_backend::TextCell;
309use crate::tf_match::{PdfWord, TfCell};
310
311pub fn table_rows(cells: &[TableCell], region: [f32; 4], words: &[TextCell]) -> Option<TableGrid> {
318 let table_words: Vec<PdfWord> = words
322 .iter()
323 .enumerate()
324 .filter(|(_, w)| !w.text.trim().is_empty())
325 .filter_map(|(wi, w)| {
326 let (l, t, r, b) = (w.l as f64, w.t as f64, w.r as f64, w.b as f64);
327 let area = (r - l) * (b - t);
328 let iw = (r.min(region[2] as f64) - l.max(region[0] as f64)).max(0.0);
329 let ih = (b.min(region[3] as f64) - t.max(region[1] as f64)).max(0.0);
330 if area > 0.0 && iw * ih / area > 0.8 {
331 Some(PdfWord {
332 id: wi,
333 bbox: [l, t, r, b],
334 text: w.text.trim().to_string(),
335 })
336 } else {
337 None
338 }
339 })
340 .collect();
341
342 if !table_words.is_empty() && !simple_match() {
343 return docling_match_rows(cells, region, &table_words, words);
344 }
345
346 let (rw, rh) = (region[2] - region[0], region[3] - region[1]);
347
348 let boxes: Vec<[f32; 4]> = cells
350 .iter()
351 .map(|c| {
352 [
353 region[0] + (c.cx - c.w / 2.0) * rw,
354 region[1] + (c.cy - c.h / 2.0) * rh,
355 region[0] + (c.cx + c.w / 2.0) * rw,
356 region[1] + (c.cy + c.h / 2.0) * rh,
357 ]
358 })
359 .collect();
360
361 let mut cell_words: Vec<Vec<usize>> = vec![Vec::new(); cells.len()];
363 for (wi, w) in words.iter().enumerate() {
364 let wa = ((w.r - w.l) * (w.b - w.t)).max(1.0);
365 let mut best: Option<(f32, usize)> = None;
366 for (ci, b) in boxes.iter().enumerate() {
367 let ix = (w.r.min(b[2]) - w.l.max(b[0])).max(0.0);
368 let iy = (w.b.min(b[3]) - w.t.max(b[1])).max(0.0);
369 let io = ix * iy / wa;
370 if io > 0.0 && best.is_none_or(|(bo, _)| io > bo) {
371 best = Some((io, ci));
372 }
373 }
374 if let Some((_, ci)) = best {
375 cell_words[ci].push(wi);
376 }
377 }
378
379 let num_rows = cells.iter().map(|c| c.row + c.rowspan).max().unwrap_or(0);
380 let num_cols = cells.iter().map(|c| c.col + c.colspan).max().unwrap_or(0);
381 if num_rows == 0 || num_cols == 0 {
382 return None;
383 }
384 let mut grid = vec![vec![String::new(); num_cols]; num_rows];
385 let mut out_cells = Vec::with_capacity(cells.len());
386 for (ci, c) in cells.iter().enumerate() {
387 let wis = std::mem::take(&mut cell_words[ci]);
390 let text = wis
391 .iter()
392 .map(|&i| words[i].text.trim())
393 .collect::<Vec<_>>()
394 .join(" ");
395 let text = normalize_cell_text(text);
396 for row in grid.iter_mut().skip(c.row).take(c.rowspan) {
398 for cell in row.iter_mut().skip(c.col).take(c.colspan) {
399 *cell = text.clone();
400 }
401 }
402 out_cells.push(first_class_cell(text, Some(boxes[ci]), c, c.tag));
403 }
404 Some(TableGrid {
405 rows: grid,
406 cells: out_cells,
407 })
408}
409
410fn first_class_cell(
413 text: String,
414 bbox: Option<[f32; 4]>,
415 c: &TableCell,
416 tag: i64,
417) -> docling_core::TableCell {
418 docling_core::TableCell {
419 text,
420 bbox,
421 start_row: c.row,
422 start_col: c.col,
423 row_span: c.rowspan.max(1),
424 col_span: c.colspan.max(1),
425 column_header: tag == CHED,
426 row_header: tag == RHED,
427 row_section: tag == SROW,
428 }
429}
430
431fn simple_match() -> bool {
434 docling_core::env::flag("DOCLING_RS_TF_SIMPLE_MATCH")
435}
436
437fn normalize_cell_text(text: String) -> String {
442 text.replace("@ ", "@")
443}
444
445fn docling_match_rows(
453 cells: &[TableCell],
454 region: [f32; 4],
455 table_words: &[PdfWord],
456 words: &[TextCell],
457) -> Option<TableGrid> {
458 const SCALE: f64 = 2.0; let sl = (region[0] as f64).round_ties_even() * SCALE;
460 let st = (region[1] as f64).round_ties_even() * SCALE;
461 let sr = (region[2] as f64).round_ties_even() * SCALE;
462 let sb = (region[3] as f64).round_ties_even() * SCALE;
463 let (w2, h2) = (sr - sl, sb - st);
464
465 let tf_cells: Vec<TfCell> = cells
466 .iter()
467 .enumerate()
468 .map(|(i, c)| {
469 let (cx, cy) = (c.cx as f64, c.cy as f64);
470 let (w, h) = (c.w as f64, c.h as f64);
471 TfCell {
472 bbox: [
473 sl + (cx - w / 2.0) * w2,
474 st + (cy - h / 2.0) * h2,
475 sl + (cx + w / 2.0) * w2,
476 st + (cy + h / 2.0) * h2,
477 ],
478 cell_id: i,
479 row_id: c.row,
480 column_id: c.col,
481 cell_class: c.class,
482 colspan_val: if c.colspan > 1 { c.colspan } else { 0 },
483 rowspan_val: if c.rowspan > 1 { c.rowspan } else { 0 },
484 }
485 })
486 .collect();
487
488 let scaled_words: Vec<PdfWord> = table_words
489 .iter()
490 .map(|w| PdfWord {
491 id: w.id,
492 bbox: [
493 w.bbox[0] * SCALE,
494 w.bbox[1] * SCALE,
495 w.bbox[2] * SCALE,
496 w.bbox[3] * SCALE,
497 ],
498 text: w.text.clone(),
499 })
500 .collect();
501
502 #[cfg(feature = "ml")]
505 if let Some(dir) = docling_core::env::nonempty("DOCLING_RS_TF_MATCH_DUMP") {
506 dump_match_inputs(&dir, &tf_cells, &scaled_words);
507 }
508
509 let (cells_wo, final_matches) =
510 crate::tf_match::match_and_post_process(tf_cells, &scaled_words);
511
512 struct Merged {
515 start_row: usize,
516 start_col: usize,
517 row_span: usize,
518 col_span: usize,
519 word_ids: Vec<usize>,
520 bbox: [f32; 4],
523 tag: i64,
525 }
526 let mut merged: Vec<Merged> = Vec::new();
527 let mut key_ix: std::collections::HashMap<(usize, usize), usize> =
528 std::collections::HashMap::new();
529 for (&pdf_id, list) in &final_matches {
530 let tm = list[0].table_cell_id;
531 let Some(cell) = cells_wo.iter().find(|c| c.cell_id == tm) else {
532 continue;
533 };
534 match key_ix.entry((cell.column_id, cell.row_id)) {
535 std::collections::hash_map::Entry::Occupied(e) => {
536 merged[*e.get()].word_ids.push(pdf_id);
537 }
538 std::collections::hash_map::Entry::Vacant(e) => {
539 e.insert(merged.len());
540 merged.push(Merged {
541 start_row: cell.row_id,
542 start_col: cell.column_id,
543 row_span: cell.rowspan_val.max(1),
544 col_span: cell.colspan_val.max(1),
545 word_ids: vec![pdf_id],
546 bbox: [
547 (cell.bbox[0] / 2.0) as f32,
548 (cell.bbox[1] / 2.0) as f32,
549 (cell.bbox[2] / 2.0) as f32,
550 (cell.bbox[3] / 2.0) as f32,
551 ],
552 tag: cells.get(cell.cell_id).map_or(FCEL, |c| c.tag),
553 });
554 }
555 }
556 }
557 if merged.is_empty() {
558 return None;
559 }
560
561 let mut start_cols: Vec<usize> = merged.iter().map(|m| m.start_col).collect();
564 start_cols.sort_unstable();
565 start_cols.dedup();
566 let mut start_rows: Vec<usize> = merged.iter().map(|m| m.start_row).collect();
567 start_rows.sort_unstable();
568 start_rows.dedup();
569 let mut num_rows = 0;
570 let mut num_cols = 0;
571 for m in &mut merged {
572 m.start_col = start_cols.binary_search(&m.start_col).expect("own value");
573 m.start_row = start_rows.binary_search(&m.start_row).expect("own value");
574 num_cols = num_cols.max(m.start_col + m.col_span);
575 num_rows = num_rows.max(m.start_row + m.row_span);
576 }
577 if num_rows == 0 || num_cols == 0 {
578 return None;
579 }
580
581 let mut grid = vec![vec![String::new(); num_cols]; num_rows];
582 let mut out_cells = Vec::with_capacity(merged.len());
583 for m in &merged {
584 let text = m
585 .word_ids
586 .iter()
587 .map(|&i| words[i].text.trim())
588 .collect::<Vec<_>>()
589 .join(" ");
590 let text = normalize_cell_text(text);
591 for row in grid.iter_mut().skip(m.start_row).take(m.row_span) {
592 for cell in row.iter_mut().skip(m.start_col).take(m.col_span) {
593 *cell = text.clone();
594 }
595 }
596 out_cells.push(docling_core::TableCell {
597 text,
598 bbox: Some(m.bbox),
599 start_row: m.start_row,
600 start_col: m.start_col,
601 row_span: m.row_span,
602 col_span: m.col_span,
603 column_header: m.tag == CHED,
604 row_header: m.tag == RHED,
605 row_section: m.tag == SROW,
606 });
607 }
608 Some(TableGrid {
609 rows: grid,
610 cells: out_cells,
611 })
612}
613
614#[cfg(feature = "ml")]
617fn dump_match_inputs(dir: &str, tf_cells: &[TfCell], words: &[PdfWord]) {
618 use std::io::Write;
619 let cells: Vec<String> = tf_cells
620 .iter()
621 .map(|c| {
622 format!(
623 r#"{{"bbox":[{},{},{},{}],"cell_id":{},"row_id":{},"column_id":{},"cell_class":{},"colspan_val":{},"rowspan_val":{}}}"#,
624 c.bbox[0], c.bbox[1], c.bbox[2], c.bbox[3],
625 c.cell_id, c.row_id, c.column_id, c.cell_class,
626 c.colspan_val, c.rowspan_val
627 )
628 })
629 .collect();
630 let ws: Vec<String> = words
631 .iter()
632 .map(|w| {
633 format!(
634 r#"{{"id":{},"bbox":[{},{},{},{}],"text":{}}}"#,
635 w.id,
636 w.bbox[0],
637 w.bbox[1],
638 w.bbox[2],
639 w.bbox[3],
640 serde_json_escape(&w.text)
641 )
642 })
643 .collect();
644 let line = format!(
645 r#"{{"table_cells":[{}],"pdf_cells":[{}]}}"#,
646 cells.join(","),
647 ws.join(",")
648 );
649 if let Ok(mut f) = std::fs::OpenOptions::new()
650 .create(true)
651 .append(true)
652 .open(format!("{dir}/tf_match_dump.jsonl"))
653 {
654 let _ = writeln!(f, "{line}");
655 }
656}
657
658#[cfg(feature = "ml")]
660fn serde_json_escape(s: &str) -> String {
661 let mut out = String::with_capacity(s.len() + 2);
662 out.push('"');
663 for ch in s.chars() {
664 match ch {
665 '"' => out.push_str("\\\""),
666 '\\' => out.push_str("\\\\"),
667 '\n' => out.push_str("\\n"),
668 '\r' => out.push_str("\\r"),
669 '\t' => out.push_str("\\t"),
670 c if (c as u32) < 0x20 => out.push_str(&format!("\\u{:04x}", c as u32)),
671 c => out.push(c),
672 }
673 }
674 out.push('"');
675 out
676}
677
678#[cfg(test)]
679mod tests {
680 use super::*;
681
682 #[test]
683 fn corrections() {
684 assert_eq!(correct(XCEL, false), LCEL); assert_eq!(correct(LCEL, true), FCEL); assert_eq!(correct(XCEL, true), FCEL); assert_eq!(correct(FCEL, false), FCEL);
688 assert_eq!(correct(LCEL, false), LCEL);
689 }
690
691 #[test]
692 fn argmax_behaviour() {
693 assert_eq!(argmax(&[0.1, 0.9, 0.3]), 1);
694 assert_eq!(argmax(&[0.5, 0.5]), 1); assert_eq!(argmax(&[]), 0);
696 }
697
698 #[test]
699 fn book_skips_first_and_collects_hiddens() {
700 let mut b = BboxBook::new();
704 let h = [1.0f32; EMBED_DIM];
705 assert!(b.step(FCEL, &h)); assert!(b.step(FCEL, &h));
707 assert!(b.step(NL, &h));
708 assert!(!b.step(END, &h)); assert_eq!(b.otsl, vec![FCEL, FCEL, NL]);
710 assert_eq!(b.n, 2);
712 assert_eq!(b.hiddens.len(), 2 * EMBED_DIM);
713 assert!(b.merge.is_empty());
714 }
715
716 #[test]
717 fn book_merges_horizontal_span() {
718 let mut b = BboxBook::new();
721 let h = [0.0f32; EMBED_DIM];
722 b.step(FCEL, &h); b.step(FCEL, &h); b.step(LCEL, &h); assert_eq!(b.merge.get(&1), Some(&-1));
726 }
727
728 #[test]
729 fn book_stops_runaway_row() {
730 let mut b = BboxBook::new();
733 let h = [0.0f32; EMBED_DIM];
734 assert!(b.step(CHED, &h));
735 assert!(b.step(CHED, &h));
736 while b.step(LCEL, &h) {
737 assert!(b.otsl.len() < MAX_STEPS);
738 }
739 assert_eq!(b.otsl.len(), MAX_ROW_TAGS);
740 assert!(b.runaway());
741 }
742
743 #[test]
744 fn book_long_table_is_not_runaway() {
745 let mut b = BboxBook::new();
749 let h = [0.0f32; EMBED_DIM];
750 for i in 0..MAX_STEPS {
751 assert!(b.step(if i % 22 == 21 { NL } else { FCEL }, &h));
752 }
753 assert!(!b.runaway());
754 let mut b = BboxBook::new();
755 for t in [FCEL, LCEL, NL, FCEL, FCEL, NL] {
756 assert!(b.step(t, &h));
757 }
758 assert!(!b.step(END, &h));
759 assert!(!b.runaway());
760 }
761
762 #[test]
763 fn build_cells_spans() {
764 let otsl = vec![FCEL, LCEL, NL, FCEL, ECEL];
767 let boxes = vec![[0.0; 4]; 3];
768 let classes = vec![2, 2, 2];
769 let cells = build_table_cells(&otsl, &boxes, &classes);
770 assert_eq!(cells.len(), 3);
771 assert_eq!((cells[0].colspan, cells[0].rowspan), (2, 1));
772 assert_eq!((cells[0].row, cells[0].col), (0, 0));
773 assert_eq!((cells[1].row, cells[1].col), (1, 0));
774 }
775}