docling_pdf/tableformer.rs
1//! TableFormer: table-structure recovery via docling-ibm-models, exported to
2//! ONNX by `scripts/install/export_tableformer.py`. The image encoder + tag-transformer
3//! encoder run once to a memory tensor; the decoder is then stepped
4//! autoregressively to emit an OTSL structure-token sequence (the same model
5//! docling runs). See docs/PDF_CONFORMANCE.md.
6
7use crate::pdfium_backend::TextCell;
8// The ONNX-free half (preprocessing, structure corrections, bbox bookkeeping,
9// span merge, OTSL→grid) lives in tf_core so the browser build (#157 stage 3)
10// runs the same logic; this file owns the three `ort` sessions and the
11// owned-value KV-cache fast path.
12use crate::tf_core::{
13 argmax, build_table_cells, correct, merge_spans, preprocess_input, BboxBook, TableCell, END,
14 MAX_ROW_TAGS, MAX_STEPS, START, UCEL,
15};
16use image::RgbImage;
17use ort::session::Session;
18use ort::value::{DynValue, Tensor};
19
20const SIDE: usize = crate::tf_core::SIDE as usize;
21const EMBED_DIM: usize = crate::tf_core::EMBED_DIM;
22/// Decoder geometry, fixed by the exported TableModel04_rs graph: the cached
23/// decoder threads a `[N_LAYERS, past, 1, EMBED_DIM]` per-layer state cache.
24const N_LAYERS: usize = 6;
25
26/// Resolve the encoder / decoder / bbox files exactly as [`TableFormer::load`]
27/// will (shared with `model_inventory`, so diagnostics can never drift from
28/// what actually loads). Explicit `DOCLING_TABLEFORMER_*` overrides win; the
29/// decoder otherwise picks by preference — INT8 variants first unless
30/// `DOCLING_RS_FP32` opts out, and within a precision the true-KV-cache
31/// export (`decoder_kv*`, one token per step, O(past) step cost) ranks ahead
32/// of the legacy layer-output-cache graph it matches byte-for-byte (91/91
33/// snapshot corpus exact with either; the KV graph re-measured ~13–17% faster
34/// warm, so speed wins the default and the legacy file stays as the smaller
35/// fallback). `decoder_kv` ranks ABOVE `decoder_int8`: the #97 hoisted fp32
36/// KV graph is faster than the quantized legacy graph on every machine
37/// measured, and it is byte-exact (its own int8 variant is not produced — see
38/// quantize_models.py).
39pub fn resolved_paths() -> (String, String, String) {
40 // The encoder ranks its fp16-weight repack (`encoder_fp16.onnx`, #374 —
41 // the same graph with the weights stored as fp16 and cast back to fp32
42 // at load, ~half the download, fp32 compute) ahead of the fp32 file
43 // unless `DOCLING_RS_FP32` opts out; an explicit override wins.
44 let enc = docling_core::env::nonempty("DOCLING_TABLEFORMER_ENCODER").unwrap_or_else(|| {
45 let candidates: &[&str] = if crate::prefer_fp32() {
46 &[".models/tableformer/encoder.onnx"]
47 } else {
48 &[
49 ".models/tableformer/encoder_fp16.onnx",
50 ".models/tableformer/encoder.onnx",
51 ]
52 };
53 candidates
54 .iter()
55 .map(|p| crate::resolve_asset(p))
56 .find(|p| std::path::Path::new(p).exists())
57 .unwrap_or_else(|| crate::resolve_asset(".models/tableformer/encoder.onnx"))
58 });
59 let dec = docling_core::env::nonempty("DOCLING_TABLEFORMER_DECODER").unwrap_or_else(|| {
60 let candidates: &[&str] = if crate::prefer_fp32() {
61 &[
62 ".models/tableformer/decoder_kv.onnx",
63 ".models/tableformer/decoder.onnx",
64 ]
65 } else {
66 &[
67 ".models/tableformer/decoder_kv_int8.onnx",
68 ".models/tableformer/decoder_kv.onnx",
69 ".models/tableformer/decoder_int8.onnx",
70 ".models/tableformer/decoder.onnx",
71 ]
72 };
73 candidates
74 .iter()
75 .map(|p| crate::resolve_asset(p))
76 .find(|p| std::path::Path::new(p).exists())
77 .unwrap_or_else(|| ".models/tableformer/decoder.onnx".to_string())
78 });
79 let bbx = docling_core::env::nonempty("DOCLING_TABLEFORMER_BBOX")
80 .unwrap_or_else(|| crate::resolve_asset(".models/tableformer/bbox.onnx"));
81 (enc, dec, bbx)
82}
83
84pub struct TableFormer {
85 encoder: Session,
86 decoder: Session,
87 bbox: Session,
88 /// Which decoder graph flavour is loaded, detected from the session's
89 /// input names (so an explicit `DOCLING_TABLEFORMER_DECODER` override
90 /// works with any of them).
91 style: DecoderStyle,
92 /// The `KvHoisted` decoder's `tag` input has a symbolic batch axis (the
93 /// dynamic-batch `decoder_kv.onnx` export): a page's tables decode
94 /// together, one step for all of them — see [`Self::predict_tables_on`].
95 /// The older fixed-`[1,1]` export decodes the tables one after another.
96 batched: bool,
97}
98
99/// The three decoder-graph generations the loop supports.
100#[derive(Clone, Copy, PartialEq, Eq)]
101enum DecoderStyle {
102 /// `decoder.onnx`: layer-output cache; feeds the full `tags` prefix and a
103 /// single `cache` every step.
104 Legacy,
105 /// The pre-#97 `decoder_kv.onnx`: one tag per step, `cache_k`/`cache_v`,
106 /// with the stacked `cross_k`/`cross_v` re-split inside every step.
107 KvStacked,
108 /// The #97 `decoder_kv.onnx`: one tag per step, and the constant cross
109 /// tensors arrive as 2×`N_LAYERS` per-layer inputs (`cross_kt_i` already
110 /// transposed for q·Kᵀ, `cross_v_i`), computed once per table by the
111 /// encoder — the step graph does no work proportional to their size.
112 KvHoisted,
113}
114
115/// KV-cache geometry fixed by the `decoder_kv.onnx` export
116/// (`[N_LAYERS, 1, KV_HEADS, past, KV_HEAD_DIM]`, `KV_HEADS × KV_HEAD_DIM = EMBED_DIM`).
117const KV_HEADS: usize = 8;
118const KV_HEAD_DIM: usize = 64;
119
120/// The autoregressive decode state: `a` is the legacy layer-output cache, or
121/// `cache_k` for the KV graph; `b` is `cache_v` (KV graph only). `None` = first
122/// step (the zero-`past` empties are allocated per table by [`TableFormer::empty_cache`]).
123#[derive(Default)]
124struct DecodeCache {
125 a: Option<DynValue>,
126 b: Option<DynValue>,
127}
128
129/// Zero-`past` first-step cache tensors: `(cache, None)` for the legacy graph,
130/// `(cache_k, Some(cache_v))` for the KV graph.
131type EmptyCache = (Tensor<f32>, Option<Tensor<f32>>);
132
133/// Encoder outputs that drive the cached decode loop: the per-layer cross-attention
134/// K/V (projected from the image memory once, constant across decode steps) and
135/// `enc_out` for the bbox decoder. Kept as owned `ort` values so each decode step
136/// (and the bbox run) borrows them directly — no per-step extract/copy/re-wrap.
137struct EncodeOut {
138 /// Stacked `[N_LAYERS,1,H,S,hd]` cross K/V — the `Legacy`/`KvStacked`
139 /// decoders' inputs. `None` for `KvHoisted`, which reads the per-layer
140 /// tensors instead: the stacked pair is 2×9.6 MB per table, and a page's
141 /// tables are now all held encoded at once for the batched loop.
142 ck: Option<DynValue>,
143 cv: Option<DynValue>,
144 eo: DynValue,
145 /// `KvHoisted` only: per-layer `[cross_kt_0..N, cross_v_0..N]`, index-aligned
146 /// with the decoder's input names, borrowed by every decode step.
147 per_layer: Vec<(String, DynValue)>,
148}
149
150impl TableFormer {
151 /// Load the exported encoder/decoder/bbox ONNX graphs (env overrides, else
152 /// `.models/tableformer/{encoder,decoder,bbox}.onnx`). Returns `None` if any is
153 /// absent, so the pipeline falls back to geometric reconstruction.
154 pub fn load() -> Option<Self> {
155 Self::load_with(crate::intra_threads())
156 }
157
158 /// Like [`load`](Self::load) but with an explicit intra-op thread count, so a
159 /// parallel page-worker pool can run each table model on fewer threads (the
160 /// throughput comes from running pages concurrently, not from one fat model).
161 ///
162 /// See [`resolved_paths`] for the encoder/decoder/bbox file selection.
163 pub fn load_with(intra: usize) -> Option<Self> {
164 // (resolution shared with the model inventory — see resolved_paths)
165 let (enc, dec, bbx) = resolved_paths();
166 if crate::timing::enabled() {
167 eprintln!("docling-pdf: tableformer decoder: {dec}");
168 }
169 if [&enc, &dec, &bbx]
170 .iter()
171 .any(|p| !std::path::Path::new(p).exists())
172 {
173 // The geometric fallback is a supported, intentional configuration
174 // (docling has no ML table-structure equivalent baked in either), so
175 // this stays a single quiet stderr note rather than an error — but it
176 // fires every process (not per-worker) so a CWD-relative default that
177 // silently misses its files (a very easy mistake for anything not run
178 // from the repo root, e.g. an embedding app) is at least visible once.
179 warn_missing_once(&enc, &dec, &bbx);
180 return None;
181 }
182 // The decoder's KV-cache grows by one entry every autoregressive step, so
183 // its input shapes differ on every `run()` call. ONNX Runtime's memory
184 // pattern optimizer assumes stable shapes to plan buffer reuse; disabling
185 // it for this session avoids repeatedly re-validating/re-touching that
186 // plan (and the external-weights file) on each step. The bbox head has
187 // the same problem one level up: its `tag_h` input is `[ncells, 512]`
188 // and every table has a different cell count, so with the pattern
189 // planner on each run re-plans — and on this graph the plan is *worse*
190 // than none: 290 ms vs 54 ms for a 100-cell table, 560 vs 94 ms for
191 // 200 cells (ORT 1.22, 4 threads). It was 0.26 s per table on the
192 // corpus, more than the encoder.
193 //
194 // The decoder runs on ONE intra-op thread. A step is 49 small GEMMs
195 // over a single token — it streams the layer weights, it does not
196 // compute — so extra threads only add synchronisation: measured 4.1 ms
197 // per step on 1 thread vs 5.5 on 4 (7.1 vs 4.9 once the cache is 100+
198 // long). In the pool it also stops a table decode from taking all the
199 // cores away from the other workers' layout inference. And a
200 // single-thread session has a fixed reduction order, so table
201 // structure no longer varies run-to-run on near-tie tokens the way
202 // multi-threaded float sums let it (the conformance scripts pin one
203 // thread for exactly that reason; the default now matches them). The
204 // encoder keeps the shared budget: one 448×448 CNN + transformer pass
205 // per table, 680 ms single-threaded vs 165 on four.
206 let build = |path: &str, mem_pattern: bool, threads: usize| -> Result<Session, String> {
207 let builder = docling_onnx::session_builder()?
208 .with_intra_threads(threads)
209 .map_err(|e| e.to_string())?
210 .with_memory_pattern(mem_pattern)
211 .map_err(|e| e.to_string())?;
212 let variant = if mem_pattern {
213 "mem_pattern"
214 } else {
215 "no_mem_pattern"
216 };
217 docling_onnx::commit(docling_onnx::apply(builder)?, path, variant)
218 .map_err(|e| format!("tableformer load {path}: {e}"))
219 };
220 match (
221 build(&enc, true, intra),
222 build(&dec, false, 1),
223 build(&bbx, false, intra),
224 ) {
225 (Ok(encoder), Ok(decoder), Ok(bbox)) => {
226 let has = |n: &str| decoder.inputs().iter().any(|i| i.name() == n);
227 let style = if has("cross_kt_0") {
228 DecoderStyle::KvHoisted
229 } else if has("cache_k") {
230 DecoderStyle::KvStacked
231 } else {
232 DecoderStyle::Legacy
233 };
234 if style == DecoderStyle::KvHoisted
235 && !encoder.outputs().iter().any(|o| o.name() == "cross_kt_0")
236 {
237 eprintln!(
238 "docling-pdf: tableformer decoder needs per-layer cross tensors \
239 (cross_kt_*) the encoder doesn't emit — re-download or re-export \
240 the model set (scripts/install/export_tableformer.py); \
241 falling back to geometric tables"
242 );
243 return None;
244 }
245 // Dynamic batch axis on `tag` ⇒ the export batches decode
246 // steps across tables (ort reports a symbolic dim as -1).
247 let batched = style == DecoderStyle::KvHoisted
248 && decoder.inputs().iter().any(|i| {
249 i.name() == "tag"
250 && matches!(i.dtype(), ort::value::ValueType::Tensor { shape, .. }
251 if shape.first().is_some_and(|d| *d < 0))
252 });
253 if crate::timing::enabled() && batched {
254 eprintln!("docling-pdf: tableformer decoder batches a page's tables per step");
255 }
256 Some(Self {
257 encoder,
258 decoder,
259 bbox,
260 style,
261 batched,
262 })
263 }
264 _ => None,
265 }
266 }
267
268 /// Run the image encoder and capture what the cached decoder loop needs: each
269 /// decoder layer's cross-attention K/V (projected from the image memory once,
270 /// shape `[N_LAYERS,1,H,S,head_dim]`) and `enc_out` for the bbox decoder.
271 fn encode(&mut self, img: &RgbImage) -> Result<EncodeOut, String> {
272 let input = crate::timing::timed("tf.preprocess", || preprocess(img))?;
273 let mut enc_out = crate::timing::timed("tf.encoder", || {
274 self.encoder
275 .run(ort::inputs!["image" => input])
276 .map_err(|e| format!("tableformer: encode: {e}"))
277 })?;
278 let mut per_layer = Vec::new();
279 if self.style == DecoderStyle::KvHoisted {
280 for prefix in ["cross_kt_", "cross_v_"] {
281 for i in 0.. {
282 let name = format!("{prefix}{i}");
283 match enc_out.remove(&name) {
284 Some(v) => per_layer.push((name, v)),
285 None => break,
286 }
287 }
288 }
289 if per_layer.is_empty() {
290 return Err("tableformer: encoder emitted no cross_kt_* outputs".into());
291 }
292 }
293 let mut grab = |name: &str| -> Result<DynValue, String> {
294 enc_out
295 .remove(name)
296 .ok_or_else(|| format!("tableformer: encoder output {name} missing"))
297 };
298 let hoisted = self.style == DecoderStyle::KvHoisted;
299 Ok(EncodeOut {
300 ck: if hoisted {
301 None
302 } else {
303 Some(grab("cross_k")?)
304 },
305 cv: if hoisted {
306 None
307 } else {
308 Some(grab("cross_v")?)
309 },
310 eo: grab("enc_out")?,
311 per_layer,
312 })
313 }
314
315 /// One doubly-cached decode step: feed the current `tags`, the constant cross
316 /// K/V, and the growing self-attention `cache`; return the raw argmax tag and
317 /// the last token's hidden state, advancing the cache. The cache stays an owned
318 /// `ort` value — the previous step's `out_cache` output is fed back directly,
319 /// never extracted or copied (it grows every step, so per-step copies were
320 /// O(steps²) float traffic). `empty_cache` is the zero-`past` value used on the
321 /// first step (ort's array constructors reject a 0-length dim, so it is
322 /// allocated through the session allocator by the caller).
323 fn decode_step(
324 &mut self,
325 tags: &[i64],
326 enc: &EncodeOut,
327 cache: &mut DecodeCache,
328 empty: &EmptyCache,
329 ) -> Result<(i64, Vec<f32>), String> {
330 crate::timing::timed("tf.decode_step", || {
331 self.decode_step_inner(tags, enc, cache, empty)
332 })
333 }
334
335 fn decode_step_inner(
336 &mut self,
337 tags: &[i64],
338 enc: &EncodeOut,
339 cache: &mut DecodeCache,
340 empty: &EmptyCache,
341 ) -> Result<(i64, Vec<f32>), String> {
342 if self.style == DecoderStyle::KvHoisted {
343 // #97 graph: one tag; the constant per-layer cross tensors are
344 // borrowed views — the step pays nothing proportional to them.
345 let last = *tags.last().expect("decode starts from <start>");
346 let (raws, hidden) = self.step_kv_hoisted(&[last], &enc.per_layer, cache, empty)?;
347 return Ok((raws[0], hidden));
348 }
349 let (ck, cv) = match (enc.ck.as_ref(), enc.cv.as_ref()) {
350 (Some(k), Some(v)) => (k, v),
351 _ => return Err("tableformer: stacked cross K/V missing".into()),
352 };
353 let mut dout = match self.style {
354 DecoderStyle::KvHoisted => unreachable!("handled above"),
355 DecoderStyle::KvStacked => {
356 // Pre-#97 KV graph: feed only the newly emitted tag; the projected
357 // K/V for the whole prefix live in cache_k/cache_v and are fed
358 // back as-is.
359 let last = *tags.last().expect("decode starts from <start>");
360 let tag_t = Tensor::from_array(([1usize, 1usize], vec![last]))
361 .map_err(|e| format!("tableformer: tag: {e}"))?;
362 match (cache.a.as_ref(), cache.b.as_ref()) {
363 (Some(k), Some(v)) => self.decoder.run(ort::inputs![
364 "tag" => tag_t, "cross_k" => ck, "cross_v" => cv,
365 "cache_k" => k, "cache_v" => v]),
366 _ => self.decoder.run(ort::inputs![
367 "tag" => tag_t, "cross_k" => ck, "cross_v" => cv,
368 "cache_k" => &empty.0,
369 "cache_v" => empty.1.as_ref().expect("kv empty cache has both halves")]),
370 }
371 }
372 DecoderStyle::Legacy => {
373 let tags_t = Tensor::from_array(([tags.len(), 1usize], tags.to_vec()))
374 .map_err(|e| format!("tableformer: tags: {e}"))?;
375 match cache.a.as_ref() {
376 None => self.decoder.run(ort::inputs![
377 "tags" => tags_t, "cross_k" => ck, "cross_v" => cv,
378 "cache" => &empty.0]),
379 Some(c) => self.decoder.run(ort::inputs![
380 "tags" => tags_t, "cross_k" => ck, "cross_v" => cv,
381 "cache" => c]),
382 }
383 }
384 }
385 .map_err(|e| format!("tableformer: decode: {e}"))?;
386 let (_, logits) = dout["logits"]
387 .try_extract_tensor::<f32>()
388 .map_err(|e| format!("tableformer: logits: {e}"))?;
389 let raw = argmax(logits) as i64;
390 let (_, hidden) = dout["hidden"]
391 .try_extract_tensor::<f32>()
392 .map_err(|e| format!("tableformer: hidden: {e}"))?;
393 let hidden = hidden.to_vec();
394 if self.style != DecoderStyle::Legacy {
395 cache.a = Some(
396 dout.remove("out_cache_k")
397 .ok_or_else(|| "tableformer: out_cache_k missing".to_string())?,
398 );
399 cache.b = Some(
400 dout.remove("out_cache_v")
401 .ok_or_else(|| "tableformer: out_cache_v missing".to_string())?,
402 );
403 } else {
404 cache.a = Some(
405 dout.remove("out_cache")
406 .ok_or_else(|| "tableformer: decoder output out_cache missing".to_string())?,
407 );
408 }
409 Ok((raw, hidden))
410 }
411
412 /// One `KvHoisted` step over `tags.len()` rows — one table per row. `tags`
413 /// holds each row's last emitted tag, `per_layer` the cross tensors with a
414 /// matching leading batch axis (the encoder's own `[1,…]` outputs for a
415 /// single table, or [`Self::batch_cross`]'s concatenation), and the cache
416 /// grows `[N_LAYERS, rows, H, past, hd]` in lockstep. Returns each row's raw
417 /// argmax tag and the `[rows, EMBED_DIM]` hidden states, flattened.
418 fn step_kv_hoisted(
419 &mut self,
420 tags: &[i64],
421 per_layer: &[(String, DynValue)],
422 cache: &mut DecodeCache,
423 empty: &EmptyCache,
424 ) -> Result<(Vec<i64>, Vec<f32>), String> {
425 let rows = tags.len();
426 let tag_t = Tensor::from_array(([rows, 1usize], tags.to_vec()))
427 .map_err(|e| format!("tableformer: tag: {e}"))?;
428 let mut inputs: Vec<(
429 std::borrow::Cow<'_, str>,
430 ort::session::SessionInputValue<'_>,
431 )> = Vec::with_capacity(3 + per_layer.len());
432 inputs.push(("tag".into(), tag_t.into()));
433 match (cache.a.as_ref(), cache.b.as_ref()) {
434 (Some(k), Some(v)) => {
435 inputs.push(("cache_k".into(), k.into()));
436 inputs.push(("cache_v".into(), v.into()));
437 }
438 _ => {
439 inputs.push(("cache_k".into(), (&empty.0).into()));
440 inputs.push((
441 "cache_v".into(),
442 empty
443 .1
444 .as_ref()
445 .expect("kv empty cache has both halves")
446 .into(),
447 ));
448 }
449 }
450 for (name, v) in per_layer {
451 inputs.push((name.as_str().into(), v.into()));
452 }
453 let mut dout = self
454 .decoder
455 .run(inputs)
456 .map_err(|e| format!("tableformer: decode: {e}"))?;
457 let (_, logits) = dout["logits"]
458 .try_extract_tensor::<f32>()
459 .map_err(|e| format!("tableformer: logits: {e}"))?;
460 let vocab = logits.len() / rows;
461 let raws: Vec<i64> = logits
462 .chunks_exact(vocab)
463 .map(|row| argmax(row) as i64)
464 .collect();
465 let (_, hidden) = dout["hidden"]
466 .try_extract_tensor::<f32>()
467 .map_err(|e| format!("tableformer: hidden: {e}"))?;
468 let hidden = hidden.to_vec();
469 cache.a = Some(
470 dout.remove("out_cache_k")
471 .ok_or_else(|| "tableformer: out_cache_k missing".to_string())?,
472 );
473 cache.b = Some(
474 dout.remove("out_cache_v")
475 .ok_or_else(|| "tableformer: out_cache_v missing".to_string())?,
476 );
477 Ok((raws, hidden))
478 }
479
480 /// Stack the per-layer cross tensors of several encoded tables along the
481 /// batch axis (`[1,H,hd,S]` × B → `[B,H,hd,S]`, same for `cross_v`), index-
482 /// aligned with the decoder's input names. One copy per page — ~20 MB per
483 /// table, nothing next to the decode steps it lets the tables share.
484 fn batch_cross(encs: &[EncodeOut]) -> Result<Vec<(String, DynValue)>, String> {
485 let b = encs.len();
486 let mut out = Vec::with_capacity(encs[0].per_layer.len());
487 for j in 0..encs[0].per_layer.len() {
488 let name = encs[0].per_layer[j].0.clone();
489 let mut data: Vec<f32> = Vec::new();
490 let mut dims = [b, 0, 0, 0];
491 for enc in encs {
492 let (shape, v) = enc.per_layer[j]
493 .1
494 .try_extract_tensor::<f32>()
495 .map_err(|e| format!("tableformer: {name}: {e}"))?;
496 if shape.len() != 4 || shape[0] != 1 {
497 return Err(format!("tableformer: {name}: unexpected shape {shape:?}"));
498 }
499 dims[1..].copy_from_slice(&[
500 shape[1] as usize,
501 shape[2] as usize,
502 shape[3] as usize,
503 ]);
504 data.reserve(v.len() * b);
505 data.extend_from_slice(v);
506 }
507 let t = Tensor::from_array((dims, data))
508 .map_err(|e| format!("tableformer: {name}: {e}"))?;
509 out.push((name, t.into_dyn()));
510 }
511 Ok(out)
512 }
513
514 /// Decode `encs.len()` tables in lockstep: every step runs the decoder once
515 /// over all of them (a step is 49 weight-streaming GEMMs over one token per
516 /// row — B rows cost about what one does). The caches start empty for
517 /// every row and grow together, so nothing is ever padded or masked; a
518 /// table that emits `<end>` simply keeps its row (fed `END`, output
519 /// ignored) until the last one finishes. Row b of every op is exactly the
520 /// single-table computation, so each table's tokens and hidden states are
521 /// bit-identical to decoding it alone (asserted by the export script's
522 /// batching gate; the corpus snapshots pin it end-to-end).
523 fn decode_batch(&mut self, encs: &[EncodeOut]) -> Result<Vec<BboxBook>, String> {
524 let b = encs.len();
525 let cross = Self::batch_cross(encs)?;
526 let mut books: Vec<BboxBook> = (0..b).map(|_| BboxBook::new()).collect();
527 let mut active = vec![true; b];
528 let mut last = vec![START; b];
529 let mut cache = DecodeCache::default();
530 let empty = self.empty_cache(b)?;
531 crate::timing::timed("tf.decode_loop", || -> Result<(), String> {
532 // Each active table's `otsl` grows by one per step, so a shared
533 // step counter is the per-table `otsl.len() < MAX_STEPS` bound.
534 for _ in 0..MAX_STEPS {
535 if !active.iter().any(|a| *a) {
536 break;
537 }
538 let (raws, hidden) = crate::timing::timed("tf.decode_step", || {
539 self.step_kv_hoisted(&last, &cross, &mut cache, &empty)
540 })?;
541 for t in 0..b {
542 if !active[t] {
543 continue;
544 }
545 let h = &hidden[t * EMBED_DIM..(t + 1) * EMBED_DIM];
546 if books[t].step(raws[t], h) {
547 last[t] = *books[t].tags.last().expect("step pushed a tag");
548 } else {
549 active[t] = false;
550 last[t] = END;
551 }
552 }
553 }
554 Ok(())
555 })?;
556 Ok(books)
557 }
558
559 /// The zero-`past` first-step cache(s) for `rows` tables, allocated through
560 /// the session allocator (ort's array constructors reject a 0-length dim;
561 /// the C API does allow it).
562 fn empty_cache(&self, rows: usize) -> Result<EmptyCache, String> {
563 let alloc = self.decoder.allocator();
564 if self.style != DecoderStyle::Legacy {
565 let mk = || {
566 Tensor::<f32>::new(alloc, [N_LAYERS, rows, KV_HEADS, 0usize, KV_HEAD_DIM])
567 .map_err(|e| format!("tableformer: empty kv cache: {e}"))
568 };
569 Ok((mk()?, Some(mk()?)))
570 } else {
571 let c = Tensor::<f32>::new(alloc, [N_LAYERS, 0usize, 1, EMBED_DIM])
572 .map_err(|e| format!("tableformer: empty cache: {e}"))?;
573 Ok((c, None))
574 }
575 }
576
577 /// Predict the OTSL structure-token sequence for a table-region image.
578 pub fn predict_otsl(&mut self, img: &RgbImage) -> Result<Vec<i64>, String> {
579 let enc = self.encode(img)?;
580 // Structure corrections live in tf_core::correct (shared with the wasm
581 // path); docling's line_num is never incremented, so xcel→lcel fires on
582 // every row.
583 let mut tags: Vec<i64> = vec![START];
584 let mut out: Vec<i64> = Vec::new();
585 let mut prev_ucel = false;
586 let mut cache = DecodeCache::default();
587 let empty = self.empty_cache(1)?;
588 while out.len() < MAX_STEPS {
589 let (raw, _hidden) = self.decode_step(&tags, &enc, &mut cache, &empty)?;
590 let tag = correct(raw, prev_ucel);
591 if tag == END {
592 break;
593 }
594 out.push(tag);
595 tags.push(tag);
596 prev_ucel = tag == UCEL;
597 }
598 Ok(out)
599 }
600
601 /// Full structure prediction: OTSL grid cells with per-cell boxes (in the 448
602 /// image, normalized cxcywh). Collects per-cell decoder hidden states using
603 /// docling's exact bbox bookkeeping (skip-after-row-break, first-lcel of a
604 /// horizontal span), runs the bbox decoder, merges span boxes, then lays the
605 /// cells onto the OTSL grid with row/col spans.
606 pub fn predict_table_structure(&mut self, img: &RgbImage) -> Result<Vec<TableCell>, String> {
607 let enc = self.encode(img)?;
608
609 // The autoregressive loop's bbox bookkeeping lives in tf_core::BboxBook
610 // (shared with the wasm path); this loop only steps the decoder.
611 let mut book = BboxBook::new();
612 let mut cache = DecodeCache::default();
613 let empty = self.empty_cache(1)?;
614 crate::timing::timed("tf.decode_loop", || -> Result<(), String> {
615 while book.otsl.len() < MAX_STEPS {
616 let (raw, hidden) = self.decode_step(&book.tags, &enc, &mut cache, &empty)?;
617 if !book.step(raw, &hidden) {
618 break;
619 }
620 }
621 Ok(())
622 })?;
623 self.finish_table(book, &enc.eo)
624 }
625
626 /// The bbox stage after a table's decode loop: run the bbox decoder over
627 /// the collected per-cell hidden states, merge span boxes, lay the cells
628 /// onto the OTSL grid.
629 fn finish_table(
630 &mut self,
631 mut book: BboxBook,
632 eo: &DynValue,
633 ) -> Result<Vec<TableCell>, String> {
634 if book.runaway() {
635 docling_core::debug_log!(
636 "docling-pdf: tableformer: no row break in {MAX_ROW_TAGS} tags; \
637 geometric table fallback"
638 );
639 return Ok(Vec::new());
640 }
641 if book.n == 0 {
642 return Ok(Vec::new());
643 }
644 let tag_h = Tensor::from_array(([book.n, EMBED_DIM], std::mem::take(&mut book.hiddens)))
645 .map_err(|e| format!("tableformer: tag_h: {e}"))?;
646 let bout = crate::timing::timed("tf.bbox", || {
647 self.bbox
648 .run(ort::inputs!["enc_out" => eo, "tag_h" => tag_h])
649 .map_err(|e| format!("tableformer: bbox: {e}"))
650 })?;
651 let (_, raw) = bout["boxes"]
652 .try_extract_tensor::<f32>()
653 .map_err(|e| format!("tableformer: boxes: {e}"))?;
654 let boxes: Vec<[f32; 4]> = raw
655 .chunks_exact(4)
656 .map(|c| [c[0], c[1], c[2], c[3]])
657 .collect();
658 // Per-cell class logits [n, 3] → argmax (docling's `outputs_class`).
659 let (_, craw) = bout["classes"]
660 .try_extract_tensor::<f32>()
661 .map_err(|e| format!("tableformer: classes: {e}"))?;
662 let classes: Vec<i64> = craw.chunks_exact(3).map(|c| argmax(c) as i64).collect();
663 let (merged, merged_classes) = merge_spans(&boxes, &classes, &book.merge);
664 Ok(build_table_cells(&book.otsl, &merged, &merged_classes))
665 }
666
667 /// Predict a table region's Markdown grid: crop the region (docling's
668 /// page→1024px box-average then bbox crop), run the structure model, then
669 /// match the page's word cells into the predicted cells with docling's
670 /// matching post-processor ([`crate::tf_match`]) and expand spans into a
671 /// dense `rows × cols` grid. `region` is `(l, t, r, b)` in page points
672 /// (top-left). Returns `None` if no structure is predicted.
673 pub fn predict_table_rows(
674 &mut self,
675 page_image: &RgbImage,
676 region: [f32; 4],
677 words: &[TextCell],
678 ) -> Option<crate::tf_core::TableGrid> {
679 let page1024 = Self::page_1024(page_image);
680 self.predict_table_rows_on(page_image.height(), &page1024, region, words)
681 }
682
683 /// The page rendered at 1024 px height (cv2.INTER_AREA), the frame every
684 /// table crop of that page is cut from. Computed once per page by the
685 /// pipeline and shared across its tables — the resample is a full-page
686 /// f64 box filter, 110–170 ms on the corpus pages, and it used to run
687 /// again for every table on the page.
688 pub fn page_1024(page_image: &RgbImage) -> RgbImage {
689 let sf = 1024.0 / page_image.height() as f32;
690 let pw = (page_image.width() as f32 * sf) as u32;
691 crate::timing::timed("tableformer.inter_area", || {
692 crate::resample::inter_area(page_image, pw, 1024)
693 })
694 }
695
696 /// [`predict_table_rows`](Self::predict_table_rows) with the page's
697 /// 1024-px frame already built ([`page_1024`](Self::page_1024));
698 /// `page_h` is the source page image's pixel height.
699 pub fn predict_table_rows_on(
700 &mut self,
701 page_h: u32,
702 page1024: &RgbImage,
703 region: [f32; 4],
704 words: &[TextCell],
705 ) -> Option<crate::tf_core::TableGrid> {
706 let crop = Self::crop_region(page_h, page1024, region)?;
707 let cells = crate::timing::timed("tableformer.structure", || {
708 self.predict_table_structure(&crop)
709 })
710 .ok()?;
711 if cells.is_empty() {
712 return None;
713 }
714 // The ort-free tail (word matching + grid assembly) is shared with the
715 // browser path in tf_core.
716 crate::tf_core::table_rows(&cells, region, words)
717 }
718
719 /// Every table of a page at once: [`predict_table_rows_on`](Self::predict_table_rows_on)
720 /// per region, except that with the dynamic-batch decoder the tables'
721 /// decode steps are shared — each table is encoded on its own, then one
722 /// decode loop steps all of them together ([`Self::decode_batch`]), then
723 /// each runs its own bbox head. Per table the result is bit-identical to
724 /// the one-at-a-time path; a page with a single table takes exactly that
725 /// path. Should the batched run fail (an ort error), the tables are
726 /// retried one by one so a page never loses all of its tables to one
727 /// shared step.
728 pub fn predict_tables_on(
729 &mut self,
730 page_h: u32,
731 page1024: &RgbImage,
732 regions: &[[f32; 4]],
733 words: &[TextCell],
734 ) -> Vec<Option<crate::tf_core::TableGrid>> {
735 let mut out: Vec<Option<crate::tf_core::TableGrid>> = vec![None; regions.len()];
736 let crops: Vec<(usize, RgbImage)> = regions
737 .iter()
738 .enumerate()
739 .filter_map(|(i, r)| Self::crop_region(page_h, page1024, *r).map(|c| (i, c)))
740 .collect();
741 if self.batched && crops.len() > 1 {
742 let batched = crate::timing::timed("tableformer.structure", || {
743 self.predict_structures_batched(crops.iter().map(|(_, c)| c))
744 });
745 match batched {
746 Ok(cells) => {
747 for ((i, _), cells) in crops.iter().zip(cells) {
748 if !cells.is_empty() {
749 out[*i] = crate::tf_core::table_rows(&cells, regions[*i], words);
750 }
751 }
752 return out;
753 }
754 Err(e) => docling_core::debug_log!(
755 "docling-pdf: tableformer batched decode failed ({e}); decoding tables one by one"
756 ),
757 }
758 }
759 for (i, crop) in &crops {
760 let cells = crate::timing::timed("tableformer.structure", || {
761 self.predict_table_structure(crop)
762 });
763 if let Ok(cells) = cells {
764 if !cells.is_empty() {
765 out[*i] = crate::tf_core::table_rows(&cells, regions[*i], words);
766 }
767 }
768 }
769 out
770 }
771
772 /// [`predict_table_structure`](Self::predict_table_structure) for several
773 /// crops with the decode steps shared across them.
774 fn predict_structures_batched<'a>(
775 &mut self,
776 crops: impl Iterator<Item = &'a RgbImage>,
777 ) -> Result<Vec<Vec<TableCell>>, String> {
778 let mut encs = Vec::new();
779 for crop in crops {
780 encs.push(self.encode(crop)?);
781 }
782 let books = self.decode_batch(&encs)?;
783 books
784 .into_iter()
785 .zip(&encs)
786 .map(|(book, enc)| self.finish_table(book, &enc.eo))
787 .collect()
788 }
789
790 /// Crop the table bbox out of the 1024px frame. docling's coordinate
791 /// chain, rounding included: the cluster bbox is rounded to integer page
792 /// points *first* (`round(cluster.bbox.l) * scale`, banker's rounding),
793 /// scaled by 2 (its table-structure page scale), then by `1024 / <2x
794 /// page-image height>`, and the crop indices round again. Rounding after
795 /// scaling instead shifts some crops by a pixel — enough to change
796 /// TableFormer's cell boxes on tall tables (redp5110's TOC). `None` for a
797 /// region that collapses to an empty crop.
798 fn crop_region(page_h: u32, page1024: &RgbImage, region: [f32; 4]) -> Option<RgbImage> {
799 let k = 2.0 * 1024.0 / page_h as f64;
800 let px = |v: f32| (v as f64).round_ties_even() * k;
801 let x = (px(region[0]).round_ties_even()).max(0.0) as u32;
802 let y = (px(region[1]).round_ties_even()).max(0.0) as u32;
803 let x2 = (px(region[2]).round_ties_even() as u32).min(page1024.width());
804 let y2 = (px(region[3]).round_ties_even() as u32).min(page1024.height());
805 if x2 <= x || y2 <= y {
806 return None;
807 }
808 Some(image::imageops::crop_imm(page1024, x, y, x2 - x, y2 - y).to_image())
809 }
810}
811
812/// Note once per process that TableFormer's ONNX graphs weren't found, so tables
813/// fall back to geometric reconstruction. The default paths are relative
814/// (`.models/tableformer/*.onnx`), which only resolves when the process's current
815/// directory happens to be the repo root — a very easy miss for anything else
816/// (an embedding app, a binding invoked from a different working directory, …),
817/// and previously failed with no signal at all.
818fn warn_missing_once(enc: &str, dec: &str, bbx: &str) {
819 static WARNED: std::sync::Once = std::sync::Once::new();
820 WARNED.call_once(|| {
821 eprintln!(
822 "docling.rs: TableFormer models not found (checked {enc}, {dec}, {bbx}); \
823 tables will use geometric reconstruction instead of ML table-structure \
824 recognition. Set DOCLING_TABLEFORMER_ENCODER / DOCLING_TABLEFORMER_DECODER \
825 / DOCLING_TABLEFORMER_BBOX to enable it (see README.md)."
826 );
827 });
828}
829
830/// docling's preprocessing: bilinear (cv2.INTER_LINEAR) resize the crop to 448²,
831/// normalize `(x/255 − mean)/std`, laid out as (C, W, H) — docling transposes
832/// (2,1,0), so width is the major spatial axis. The page→1024px box-average
833/// (cv2.INTER_AREA) is the caller's job.
834fn preprocess(img: &RgbImage) -> Result<Tensor<f32>, String> {
835 Tensor::from_array(([1usize, 3, SIDE, SIDE], preprocess_input(img)))
836 .map_err(|e| format!("tableformer: input: {e}"))
837}