docling_pdf/layout.rs
1//! Layout detection via the RT-DETR (`docling-layout-heron`) model exported to
2//! ONNX, run with `ort`. A port of docling-ibm-models' `LayoutPredictor`:
3//! resize the page image to 640×640 and rescale to `[0,1]` (the heron processor
4//! has `do_normalize=false`), run the model, then RT-DETR
5//! `post_process_object_detection` (sigmoid → top-k over query×class →
6//! center-to-corners boxes scaled to the page).
7
8#[cfg(feature = "ml")]
9use image::imageops::FilterType;
10#[cfg(feature = "ml")]
11use ort::session::Session;
12#[cfg(feature = "ml")]
13use ort::value::Tensor;
14
15/// The 17 canonical layout classes, indexed by the model's class id
16/// (`config.json` `id2label`).
17pub const LABELS: [&str; 17] = [
18 "caption",
19 "footnote",
20 "formula",
21 "list_item",
22 "page_footer",
23 "page_header",
24 "picture",
25 "section_header",
26 "table",
27 "text",
28 "title",
29 "document_index",
30 "code",
31 "checkbox_selected",
32 "checkbox_unselected",
33 "form",
34 "key_value_region",
35];
36
37/// One detected region, in page points (top-left origin).
38#[derive(Debug, Clone)]
39pub struct Region {
40 pub label: &'static str,
41 pub score: f32,
42 pub l: f32,
43 pub t: f32,
44 pub r: f32,
45 pub b: f32,
46}
47
48/// What a layout inference call receives per page — which resize kernel packs
49/// the 640×640 model input depends on it (docling parity, #58-branch):
50///
51/// docling's layout stage runs on `page.get_image(scale=1.0)` — the
52/// point-sized page image (pdfium at 1.5×, PIL-BICUBIC down) — which its
53/// RT-DETR processor then stretches to 640×640 with **PIL BILINEAR**
54/// (`preprocessor_config.json`: `do_pad: false`, `resample: 2`; no letterbox,
55/// no normalize beyond `/255`). [`PageImage`](LayoutSrc::PageImage) is that
56/// image and goes through the byte-exact PIL kernel. [`Raw`](LayoutSrc::Raw)
57/// is any other bitmap (the browser path's canvas render, METS/TIFF page
58/// scans) and keeps the legacy Triangle stretch.
59#[cfg(feature = "ocr-prep")]
60#[derive(Clone, Copy)]
61pub enum LayoutSrc<'a> {
62 /// The scale-1.0 page image (`PdfPage::image_layout`), exact against
63 /// docling's pypdfium2 backend (#478).
64 PageImage(&'a image::RgbImage),
65 /// Any other page bitmap — legacy stretch.
66 Raw(&'a image::RgbImage),
67}
68
69/// Base confidence threshold (docling-ibm-models `base_threshold`): the raw
70/// RT-DETR floor before docling's `LayoutPostprocessor` applies its stricter
71/// per-label thresholds ([`label_threshold`]).
72const THRESHOLD: f32 = 0.3;
73/// RT-DETR's fixed square input side.
74pub const SIDE: u32 = 640;
75
76/// Per-label confidence threshold, ported from docling's
77/// `LayoutPostprocessor.CONFIDENCE_THRESHOLDS`. The raw predictor keeps every
78/// detection above the 0.3 base; the postprocessor then drops a cluster whose
79/// score is below its label's threshold. Applying it here (equivalent, since
80/// every per-label threshold is ≥ the 0.3 base) keeps low-confidence pictures /
81/// tables / list-items out of the assembly, matching docling.
82pub fn label_threshold(label: &str) -> f32 {
83 match label {
84 "section_header"
85 | "title"
86 | "code"
87 | "checkbox_selected"
88 | "checkbox_unselected"
89 | "form"
90 | "key_value_region"
91 | "document_index" => 0.45,
92 // caption, footnote, formula, list_item, page_footer, page_header,
93 // picture, table, text — all 0.5 in docling.
94 _ => 0.5,
95 }
96}
97
98#[cfg(feature = "ml")]
99pub struct LayoutModel {
100 session: Session,
101 /// Set when a multi-page inference fails — e.g. a locally built pre-#73
102 /// static graph (fixed batch=1) via `DOCLING_LAYOUT_ONNX` or a stale
103 /// `layout_heron_int8.onnx`. Batched calls then fall back to per-page runs
104 /// instead of failing the conversion.
105 batch_unsupported: bool,
106 /// The fp32 graph to escalate a suspicious page to, set only when the
107 /// *auto-selected* int8 graph loaded (an explicit `DOCLING_LAYOUT_ONNX` /
108 /// `DOCLING_RS_FP32` choice is respected). Int8 confidences sit close
109 /// enough to the 0.5 label thresholds that a different CPU's quantized
110 /// kernels (AVX-VNNI vs AVX2, CUDA's fallback mix) can flip a whole page's
111 /// detections — observed as a bill page whose tables all dissolved into
112 /// orphan lines on one machine while converting perfectly on another.
113 fp32_path: Option<String>,
114 /// Lazily-loaded session over `fp32_path` — most documents never pay for it.
115 fp32: Option<Session>,
116 /// Intra-op threads, kept for the lazy fp32 load.
117 intra: usize,
118}
119
120#[cfg(feature = "ml")]
121impl LayoutModel {
122 /// Load the ONNX model from `DOCLING_LAYOUT_ONNX`. Without the override,
123 /// prefers `.models/layout_heron_int8.onnx` when present (the quantized
124 /// default; `DOCLING_RS_FP32=1` opts out), else `.models/layout_heron.onnx`.
125 pub fn load() -> Result<Self, String> {
126 Self::load_with(crate::intra_threads())
127 }
128
129 /// Like [`load`](Self::load) but with an explicit intra-op thread count. A
130 /// parallel page-worker pool loads its helper models on a single thread each
131 /// and gets its speed-up from running pages concurrently instead.
132 pub fn load_with(intra: usize) -> Result<Self, String> {
133 let path = crate::model_path(
134 "DOCLING_LAYOUT_ONNX",
135 ".models/layout_heron.onnx",
136 ".models/layout_heron_int8.onnx",
137 );
138 if crate::timing::enabled() {
139 eprintln!("docling-pdf: layout model: {path}");
140 }
141 // Escalation target for the quant-robustness guard: only when the
142 // int8 graph was picked automatically and the fp32 one is also there.
143 let fp32_path = if docling_core::env::nonempty("DOCLING_LAYOUT_ONNX").is_none() {
144 let fp32 = crate::resolve_asset(".models/layout_heron.onnx");
145 (path != fp32 && std::path::Path::new(&fp32).exists()).then_some(fp32)
146 } else {
147 None
148 };
149 let session = Self::open_session(&path, intra)?;
150 Ok(Self {
151 session,
152 batch_unsupported: false,
153 fp32_path,
154 fp32: None,
155 intra,
156 })
157 }
158
159 fn open_session(path: &str, intra: usize) -> Result<Session, String> {
160 // The layout model is the pipeline's first hard model dependency; a
161 // missing file here almost always means the models were never
162 // downloaded (`cargo install` ships none) — say what to do.
163 if !std::path::Path::new(path).exists() {
164 return Err(format!(
165 "layout: model not found at {path} — PDF/image conversion needs \
166 the ONNX models: fetch them with \
167 scripts/install/download_dependencies.sh from a docling.rs \
168 checkout (https://github.com/docling-project/docling.rs), or \
169 set DOCLING_LAYOUT_ONNX. A digital PDF's embedded text layer \
170 converts without models in no-OCR mode (CLI: --no-ocr)"
171 ));
172 }
173 let mut builder = Session::builder()
174 .map_err(|e| format!("layout: builder: {e}"))?
175 // Let inference use the available cores (ort otherwise defaults low);
176 // a large PDF runs this model once per page.
177 .with_intra_threads(intra)
178 .map_err(|e| format!("layout: intra_threads: {e}"))?;
179 // Per-page mode pins the model's dynamic `batch` axis to 1 (#339):
180 // the free dimension blocks ONNX Runtime's channels-last conv
181 // transform, so the graph runs NCHW `FusedConv` instead of
182 // `NhwcFusedConv` — the issue measured ~1.4× on Apple-silicon CPU
183 // for the same weights re-exported static. Overriding the dimension
184 // at session creation gets the static graph without a re-export; it
185 // also leaves the whole graph static-shaped, which is what the
186 // CoreML provider's static-partitions default (#324) wants. Batched
187 // mode keeps the axis free — those sessions must accept N pages.
188 if crate::pdf_layout_batch() == 1 {
189 builder = builder
190 .with_dimension_override("batch", 1)
191 .map_err(|e| format!("layout: dimension override: {e}"))?;
192 }
193 let builder = docling_onnx::apply(builder).map_err(|e| format!("layout: {e}"))?;
194 // The pinned batch axis changes the optimized graph — separate cache entry.
195 let variant = if crate::pdf_layout_batch() == 1 {
196 "batch=1"
197 } else {
198 "batch=dyn"
199 };
200 docling_onnx::commit(builder, path, variant)
201 .map_err(|e| format!("layout: load {path}: {e}"))
202 }
203
204 /// Re-run one page through the fp32 graph — the escape hatch for a page
205 /// whose int8 detections look implausible (see `fp32_path`). `Ok(None)`
206 /// when there is nothing to escalate to: fp32 already loaded, an explicit
207 /// model override, or no fp32 file on disk.
208 pub fn predict_fp32_fallback(
209 &mut self,
210 img: LayoutSrc<'_>,
211 page_w: f32,
212 page_h: f32,
213 ) -> Result<Option<Vec<Region>>, String> {
214 let Some(path) = self.fp32_path.clone() else {
215 return Ok(None);
216 };
217 if self.fp32.is_none() {
218 if crate::timing::enabled() {
219 eprintln!("docling-pdf: loading fp32 layout fallback: {path}");
220 }
221 self.fp32 = Some(Self::open_session(&path, self.intra)?);
222 }
223 let session = self.fp32.as_mut().expect("just loaded");
224 Ok(Some(
225 Self::run_on(session, &[(img, page_w, page_h)])?
226 .pop()
227 .expect("one result per input page"),
228 ))
229 }
230
231 /// Detect layout regions on a page image. `page_w`/`page_h` are the page size
232 /// in points; returned boxes are in those coordinates.
233 pub fn predict(
234 &mut self,
235 img: LayoutSrc<'_>,
236 page_w: f32,
237 page_h: f32,
238 ) -> Result<Vec<Region>, String> {
239 Ok(self
240 .predict_batch(&[(img, page_w, page_h)])?
241 .pop()
242 .expect("one result per input page"))
243 }
244
245 /// Detect layout regions on several page images with **one** inference call
246 /// (issue #73). The ONNX export has a dynamic batch dimension, so a worker
247 /// can amortize the per-run framework overhead and keep its cores busier on
248 /// multi-page documents. Results are per-image, index-aligned with `pages`,
249 /// and identical to calling [`predict`](Self::predict) per page.
250 pub fn predict_batch(
251 &mut self,
252 pages: &[(LayoutSrc<'_>, f32, f32)],
253 ) -> Result<Vec<Vec<Region>>, String> {
254 if pages.len() > 1 && self.batch_unsupported {
255 return self.predict_singly(pages);
256 }
257 match self.run_batch(pages) {
258 Err(e) if pages.len() > 1 => {
259 // A graph without the dynamic batch dim (pre-#73 export) fails
260 // only for batch > 1 — remember and recover per page. Warn once
261 // per process, not per worker: every worker owns a LayoutModel
262 // over the same graph file, so repeats carry no information.
263 static WARNED: std::sync::atomic::AtomicBool =
264 std::sync::atomic::AtomicBool::new(false);
265 if !WARNED.swap(true, std::sync::atomic::Ordering::Relaxed) {
266 eprintln!(
267 "docling-pdf: layout model rejected a {}-page batch ({e}); \
268 falling back to per-page inference — re-export with \
269 scripts/install/export_layout.py for batched layout",
270 pages.len()
271 );
272 }
273 self.batch_unsupported = true;
274 self.predict_singly(pages)
275 }
276 other => other,
277 }
278 }
279
280 fn predict_singly(
281 &mut self,
282 pages: &[(LayoutSrc<'_>, f32, f32)],
283 ) -> Result<Vec<Vec<Region>>, String> {
284 pages
285 .iter()
286 .map(|p| Ok(self.run_batch(&[*p])?.pop().expect("one result")))
287 .collect()
288 }
289
290 fn run_batch(
291 &mut self,
292 pages: &[(LayoutSrc<'_>, f32, f32)],
293 ) -> Result<Vec<Vec<Region>>, String> {
294 Self::run_on(&mut self.session, pages)
295 }
296
297 fn run_on(
298 session: &mut Session,
299 pages: &[(LayoutSrc<'_>, f32, f32)],
300 ) -> Result<Vec<Vec<Region>>, String> {
301 if pages.is_empty() {
302 return Ok(Vec::new());
303 }
304 // Resize each page to 640×640 (RT-DETR ignores aspect ratio), rescale to
305 // [0,1], lay out as NCHW. The kernel depends on the source (see
306 // [`LayoutSrc`]): the pypdfium2-exact page image goes through Pillow's
307 // BILINEAR (the RT-DETR processor's kernel, byte-for-byte), raw
308 // bitmaps keep the legacy Triangle stretch.
309 let n = (SIDE * SIDE) as usize;
310 let batch = pages.len();
311 let mut data = vec![0f32; batch * 3 * n];
312 for (p, (src, _, _)) in pages.iter().enumerate() {
313 let resized = match src {
314 LayoutSrc::PageImage(img) => crate::resample::pil_resize(
315 img,
316 SIDE,
317 SIDE,
318 crate::resample::PilFilter::Bilinear,
319 ),
320 LayoutSrc::Raw(img) => {
321 image::imageops::resize(*img, SIDE, SIDE, FilterType::Triangle)
322 }
323 };
324 let page_off = p * 3 * n;
325 for (i, px) in resized.pixels().enumerate() {
326 data[page_off + i] = px[0] as f32 / 255.0;
327 data[page_off + n + i] = px[1] as f32 / 255.0;
328 data[page_off + 2 * n + i] = px[2] as f32 / 255.0;
329 }
330 }
331 let input = Tensor::from_array(([batch, 3, SIDE as usize, SIDE as usize], data))
332 .map_err(|e| format!("layout: input tensor: {e}"))?;
333 let outputs = session
334 .run(ort::inputs!["pixel_values" => input])
335 .map_err(|e| format!("layout: inference: {e}"))?;
336 let (lshape, logits) = outputs["logits"]
337 .try_extract_tensor::<f32>()
338 .map_err(|e| format!("layout: extract logits: {e}"))?;
339 let (_, boxes) = outputs["pred_boxes"]
340 .try_extract_tensor::<f32>()
341 .map_err(|e| format!("layout: extract boxes: {e}"))?;
342
343 let num_queries = lshape[1] as usize;
344 let num_classes = lshape[2] as usize;
345
346 let mut all = Vec::with_capacity(batch);
347 for (p, (_, page_w, page_h)) in pages.iter().enumerate() {
348 let logits =
349 &logits[p * num_queries * num_classes..(p + 1) * num_queries * num_classes];
350 let boxes = &boxes[p * num_queries * 4..(p + 1) * num_queries * 4];
351 all.push(decode_layout(
352 logits,
353 boxes,
354 num_queries,
355 num_classes,
356 *page_w,
357 *page_h,
358 ));
359 }
360 Ok(all)
361 }
362}
363
364fn sigmoid(x: f32) -> f32 {
365 1.0 / (1.0 + (-x).exp())
366}
367
368/// Pack one page image into the model's `(1, 3, SIDE, SIDE)` input: resize
369/// (aspect ignored, RT-DETR convention), rescale to `[0,1]`, CHW. Shared
370/// with the browser build (#157), which delegates only the session call.
371#[cfg(feature = "ocr-prep")]
372pub fn layout_input(img: &image::RgbImage) -> Vec<f32> {
373 let n = (SIDE * SIDE) as usize;
374 let mut data = vec![0f32; 3 * n];
375 let resized = image::imageops::resize(img, SIDE, SIDE, image::imageops::FilterType::Triangle);
376 for (i, px) in resized.pixels().enumerate() {
377 data[i] = px[0] as f32 / 255.0;
378 data[n + i] = px[1] as f32 / 255.0;
379 data[2 * n + i] = px[2] as f32 / 255.0;
380 }
381 data
382}
383
384/// Decode one page's raw RT-DETR outputs into scored [`Region`]s in page
385/// points — sigmoid over every (query, class), top-`num_queries` kept, boxes
386/// converted center→corners and scaled. Shared with the browser build; the
387/// native batch path calls it per page, so both decode identically.
388pub fn decode_layout(
389 logits: &[f32],
390 boxes: &[f32],
391 num_queries: usize,
392 num_classes: usize,
393 page_w: f32,
394 page_h: f32,
395) -> Vec<Region> {
396 let mut scored: Vec<(f32, usize)> = (0..num_queries * num_classes)
397 .map(|idx| (sigmoid(logits[idx]), idx))
398 .collect();
399 scored.sort_unstable_by(|a, b| b.0.total_cmp(&a.0));
400 scored.truncate(num_queries);
401
402 let mut regions = Vec::new();
403 for (score, idx) in scored {
404 if score <= THRESHOLD {
405 continue;
406 }
407 let label_id = idx % num_classes;
408 let q = idx / num_classes;
409 let cx = boxes[q * 4];
410 let cy = boxes[q * 4 + 1];
411 let w = boxes[q * 4 + 2];
412 let h = boxes[q * 4 + 3];
413 // center_to_corners, then scale normalized coords to page points.
414 let l = (cx - w / 2.0) * page_w;
415 let t = (cy - h / 2.0) * page_h;
416 let r = (cx + w / 2.0) * page_w;
417 let b = (cy + h / 2.0) * page_h;
418 regions.push(Region {
419 label: LABELS.get(label_id).copied().unwrap_or("text"),
420 score,
421 l,
422 t,
423 r,
424 b,
425 });
426 }
427 regions
428}