use crate::pdfium_backend::TextCell;
use crate::tf_core::{
argmax, build_table_cells, correct, merge_spans, preprocess_input, BboxBook, TableCell, END,
MAX_STEPS, START, UCEL,
};
use image::RgbImage;
use ort::session::Session;
use ort::value::{DynValue, Tensor};
const SIDE: usize = crate::tf_core::SIDE as usize;
const EMBED_DIM: usize = crate::tf_core::EMBED_DIM;
const N_LAYERS: usize = 6;
pub struct TableFormer {
encoder: Session,
decoder: Session,
bbox: Session,
style: DecoderStyle,
}
#[derive(Clone, Copy, PartialEq, Eq)]
enum DecoderStyle {
Legacy,
KvStacked,
KvHoisted,
}
const KV_HEADS: usize = 8;
const KV_HEAD_DIM: usize = 64;
#[derive(Default)]
struct DecodeCache {
a: Option<DynValue>,
b: Option<DynValue>,
}
type EmptyCache = (Tensor<f32>, Option<Tensor<f32>>);
struct EncodeOut {
ck: DynValue,
cv: DynValue,
eo: DynValue,
per_layer: Vec<(String, DynValue)>,
}
impl TableFormer {
pub fn load() -> Option<Self> {
Self::load_with(crate::intra_threads())
}
pub fn load_with(intra: usize) -> Option<Self> {
let enc = std::env::var("DOCLING_TABLEFORMER_ENCODER")
.unwrap_or_else(|_| crate::resolve_asset("models/tableformer/encoder.onnx"));
let dec = std::env::var("DOCLING_TABLEFORMER_DECODER").unwrap_or_else(|_| {
let candidates: &[&str] = if crate::prefer_fp32() {
&[
"models/tableformer/decoder_kv.onnx",
"models/tableformer/decoder.onnx",
]
} else {
&[
"models/tableformer/decoder_kv_int8.onnx",
"models/tableformer/decoder_kv.onnx",
"models/tableformer/decoder_int8.onnx",
"models/tableformer/decoder.onnx",
]
};
candidates
.iter()
.map(|p| crate::resolve_asset(p))
.find(|p| std::path::Path::new(p).exists())
.unwrap_or_else(|| "models/tableformer/decoder.onnx".to_string())
});
let bbx = std::env::var("DOCLING_TABLEFORMER_BBOX")
.unwrap_or_else(|_| crate::resolve_asset("models/tableformer/bbox.onnx"));
if crate::timing::enabled() {
eprintln!("docling-pdf: tableformer decoder: {dec}");
}
if [&enc, &dec, &bbx]
.iter()
.any(|p| !std::path::Path::new(p).exists())
{
warn_missing_once(&enc, &dec, &bbx);
return None;
}
let build = |path: &str, mem_pattern: bool| -> Result<Session, String> {
let builder = Session::builder()
.map_err(|e| e.to_string())?
.with_intra_threads(intra)
.map_err(|e| e.to_string())?
.with_memory_pattern(mem_pattern)
.map_err(|e| e.to_string())?;
crate::ep::apply(builder)?
.commit_from_file(path)
.map_err(|e| format!("tableformer load {path}: {e}"))
};
match (build(&enc, true), build(&dec, false), build(&bbx, true)) {
(Ok(encoder), Ok(decoder), Ok(bbox)) => {
let has = |n: &str| decoder.inputs().iter().any(|i| i.name() == n);
let style = if has("cross_kt_0") {
DecoderStyle::KvHoisted
} else if has("cache_k") {
DecoderStyle::KvStacked
} else {
DecoderStyle::Legacy
};
if style == DecoderStyle::KvHoisted
&& !encoder.outputs().iter().any(|o| o.name() == "cross_kt_0")
{
eprintln!(
"docling-pdf: tableformer decoder needs per-layer cross tensors \
(cross_kt_*) the encoder doesn't emit — re-download or re-export \
the model set (scripts/install/export_tableformer.py); \
falling back to geometric tables"
);
return None;
}
Some(Self {
encoder,
decoder,
bbox,
style,
})
}
_ => None,
}
}
fn encode(&mut self, img: &RgbImage) -> Result<EncodeOut, String> {
let input = preprocess(img)?;
let mut enc_out = self
.encoder
.run(ort::inputs!["image" => input])
.map_err(|e| format!("tableformer: encode: {e}"))?;
let mut per_layer = Vec::new();
if self.style == DecoderStyle::KvHoisted {
for prefix in ["cross_kt_", "cross_v_"] {
for i in 0.. {
let name = format!("{prefix}{i}");
match enc_out.remove(&name) {
Some(v) => per_layer.push((name, v)),
None => break,
}
}
}
if per_layer.is_empty() {
return Err("tableformer: encoder emitted no cross_kt_* outputs".into());
}
}
let mut grab = |name: &str| -> Result<DynValue, String> {
enc_out
.remove(name)
.ok_or_else(|| format!("tableformer: encoder output {name} missing"))
};
Ok(EncodeOut {
ck: grab("cross_k")?,
cv: grab("cross_v")?,
eo: grab("enc_out")?,
per_layer,
})
}
fn decode_step(
&mut self,
tags: &[i64],
enc: &EncodeOut,
cache: &mut DecodeCache,
empty: &EmptyCache,
) -> Result<(i64, Vec<f32>), String> {
crate::timing::timed("tf.decode_step", || {
self.decode_step_inner(tags, enc, cache, empty)
})
}
fn decode_step_inner(
&mut self,
tags: &[i64],
enc: &EncodeOut,
cache: &mut DecodeCache,
empty: &EmptyCache,
) -> Result<(i64, Vec<f32>), String> {
let mut dout = match self.style {
DecoderStyle::KvHoisted => {
let last = *tags.last().expect("decode starts from <start>");
let tag_t = Tensor::from_array(([1usize, 1usize], vec![last]))
.map_err(|e| format!("tableformer: tag: {e}"))?;
let mut inputs: Vec<(
std::borrow::Cow<'_, str>,
ort::session::SessionInputValue<'_>,
)> = Vec::with_capacity(3 + enc.per_layer.len());
inputs.push(("tag".into(), tag_t.into()));
match (cache.a.as_ref(), cache.b.as_ref()) {
(Some(k), Some(v)) => {
inputs.push(("cache_k".into(), k.into()));
inputs.push(("cache_v".into(), v.into()));
}
_ => {
inputs.push(("cache_k".into(), (&empty.0).into()));
inputs.push((
"cache_v".into(),
empty
.1
.as_ref()
.expect("kv empty cache has both halves")
.into(),
));
}
}
for (name, v) in &enc.per_layer {
inputs.push((name.as_str().into(), v.into()));
}
self.decoder.run(inputs)
}
DecoderStyle::KvStacked => {
let last = *tags.last().expect("decode starts from <start>");
let tag_t = Tensor::from_array(([1usize, 1usize], vec![last]))
.map_err(|e| format!("tableformer: tag: {e}"))?;
match (cache.a.as_ref(), cache.b.as_ref()) {
(Some(k), Some(v)) => self.decoder.run(ort::inputs![
"tag" => tag_t, "cross_k" => &enc.ck, "cross_v" => &enc.cv,
"cache_k" => k, "cache_v" => v]),
_ => self.decoder.run(ort::inputs![
"tag" => tag_t, "cross_k" => &enc.ck, "cross_v" => &enc.cv,
"cache_k" => &empty.0,
"cache_v" => empty.1.as_ref().expect("kv empty cache has both halves")]),
}
}
DecoderStyle::Legacy => {
let tags_t = Tensor::from_array(([tags.len(), 1usize], tags.to_vec()))
.map_err(|e| format!("tableformer: tags: {e}"))?;
match cache.a.as_ref() {
None => self.decoder.run(ort::inputs![
"tags" => tags_t, "cross_k" => &enc.ck, "cross_v" => &enc.cv,
"cache" => &empty.0]),
Some(c) => self.decoder.run(ort::inputs![
"tags" => tags_t, "cross_k" => &enc.ck, "cross_v" => &enc.cv,
"cache" => c]),
}
}
}
.map_err(|e| format!("tableformer: decode: {e}"))?;
let (_, logits) = dout["logits"]
.try_extract_tensor::<f32>()
.map_err(|e| format!("tableformer: logits: {e}"))?;
let raw = argmax(logits) as i64;
let (_, hidden) = dout["hidden"]
.try_extract_tensor::<f32>()
.map_err(|e| format!("tableformer: hidden: {e}"))?;
let hidden = hidden.to_vec();
if self.style != DecoderStyle::Legacy {
cache.a = Some(
dout.remove("out_cache_k")
.ok_or_else(|| "tableformer: out_cache_k missing".to_string())?,
);
cache.b = Some(
dout.remove("out_cache_v")
.ok_or_else(|| "tableformer: out_cache_v missing".to_string())?,
);
} else {
cache.a = Some(
dout.remove("out_cache")
.ok_or_else(|| "tableformer: decoder output out_cache missing".to_string())?,
);
}
Ok((raw, hidden))
}
fn empty_cache(&self) -> Result<EmptyCache, String> {
let alloc = self.decoder.allocator();
if self.style != DecoderStyle::Legacy {
let mk = || {
Tensor::<f32>::new(alloc, [N_LAYERS, 1, KV_HEADS, 0usize, KV_HEAD_DIM])
.map_err(|e| format!("tableformer: empty kv cache: {e}"))
};
Ok((mk()?, Some(mk()?)))
} else {
let c = Tensor::<f32>::new(alloc, [N_LAYERS, 0usize, 1, EMBED_DIM])
.map_err(|e| format!("tableformer: empty cache: {e}"))?;
Ok((c, None))
}
}
pub fn predict_otsl(&mut self, img: &RgbImage) -> Result<Vec<i64>, String> {
let enc = self.encode(img)?;
let mut tags: Vec<i64> = vec![START];
let mut out: Vec<i64> = Vec::new();
let mut prev_ucel = false;
let mut cache = DecodeCache::default();
let empty = self.empty_cache()?;
while out.len() < MAX_STEPS {
let (raw, _hidden) = self.decode_step(&tags, &enc, &mut cache, &empty)?;
let tag = correct(raw, prev_ucel);
if tag == END {
break;
}
out.push(tag);
tags.push(tag);
prev_ucel = tag == UCEL;
}
Ok(out)
}
pub fn predict_table_structure(&mut self, img: &RgbImage) -> Result<Vec<TableCell>, String> {
let enc = self.encode(img)?;
let mut book = BboxBook::new();
let mut cache = DecodeCache::default();
let empty = self.empty_cache()?;
while book.otsl.len() < MAX_STEPS {
let (raw, hidden) = self.decode_step(&book.tags, &enc, &mut cache, &empty)?;
if !book.step(raw, &hidden) {
break;
}
}
if book.n == 0 {
return Ok(Vec::new());
}
let tag_h = Tensor::from_array(([book.n, EMBED_DIM], std::mem::take(&mut book.hiddens)))
.map_err(|e| format!("tableformer: tag_h: {e}"))?;
let bout = self
.bbox
.run(ort::inputs!["enc_out" => &enc.eo, "tag_h" => tag_h])
.map_err(|e| format!("tableformer: bbox: {e}"))?;
let (_, raw) = bout["boxes"]
.try_extract_tensor::<f32>()
.map_err(|e| format!("tableformer: boxes: {e}"))?;
let boxes: Vec<[f32; 4]> = raw
.chunks_exact(4)
.map(|c| [c[0], c[1], c[2], c[3]])
.collect();
let (_, craw) = bout["classes"]
.try_extract_tensor::<f32>()
.map_err(|e| format!("tableformer: classes: {e}"))?;
let classes: Vec<i64> = craw.chunks_exact(3).map(|c| argmax(c) as i64).collect();
let (merged, merged_classes) = merge_spans(&boxes, &classes, &book.merge);
Ok(build_table_cells(&book.otsl, &merged, &merged_classes))
}
pub fn predict_table_rows(
&mut self,
page_image: &RgbImage,
region: [f32; 4],
words: &[TextCell],
) -> Option<Vec<Vec<String>>> {
let sf = 1024.0 / page_image.height() as f32;
let pw = (page_image.width() as f32 * sf) as u32;
let page1024 = crate::timing::timed("tableformer.inter_area", || {
crate::resample::inter_area(page_image, pw, 1024)
});
let k = 2.0 * 1024.0 / page_image.height() as f64;
let px = |v: f32| (v as f64).round_ties_even() * k;
let x = (px(region[0]).round_ties_even()).max(0.0) as u32;
let y = (px(region[1]).round_ties_even()).max(0.0) as u32;
let x2 = (px(region[2]).round_ties_even() as u32).min(page1024.width());
let y2 = (px(region[3]).round_ties_even() as u32).min(page1024.height());
if x2 <= x || y2 <= y {
return None;
}
let crop = image::imageops::crop_imm(&page1024, x, y, x2 - x, y2 - y).to_image();
let cells = crate::timing::timed("tableformer.structure", || {
self.predict_table_structure(&crop)
})
.ok()?;
if cells.is_empty() {
return None;
}
crate::tf_core::table_rows(&cells, region, words)
}
}
fn warn_missing_once(enc: &str, dec: &str, bbx: &str) {
static WARNED: std::sync::Once = std::sync::Once::new();
WARNED.call_once(|| {
eprintln!(
"docling.rs: TableFormer models not found (checked {enc}, {dec}, {bbx}); \
tables will use geometric reconstruction instead of ML table-structure \
recognition. Set DOCLING_TABLEFORMER_ENCODER / DOCLING_TABLEFORMER_DECODER \
/ DOCLING_TABLEFORMER_BBOX to enable it (see README.md)."
);
});
}
fn preprocess(img: &RgbImage) -> Result<Tensor<f32>, String> {
Tensor::from_array(([1usize, 3, SIDE, SIDE], preprocess_input(img)))
.map_err(|e| format!("tableformer: input: {e}"))
}