use image::RgbImage;
pub const LIMIT_SIDE_LEN: u32 = 736;
pub const THRESH: f32 = 0.3;
pub const BOX_THRESH: f32 = 0.5;
pub const UNCLIP_RATIO: f32 = 1.6;
const MIN_SIZE: f32 = 3.0;
const BOX_SORT_Y_THRESHOLD: f32 = 10.0;
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct DetBox {
pub l: f32,
pub t: f32,
pub r: f32,
pub b: f32,
pub score: f32,
}
pub fn det_input_size(w: u32, h: u32) -> Option<(u32, u32)> {
det_input_size_capped(w, h, max_side_cap())
}
pub const DEFAULT_MAX_SIDE: u32 = 960;
fn max_side_cap() -> u32 {
static CAP: std::sync::OnceLock<u32> = std::sync::OnceLock::new();
*CAP.get_or_init(|| {
docling_core::env::parse::<u32>("DOCLING_RS_OCR_DET_MAX_SIDE").unwrap_or(DEFAULT_MAX_SIDE)
})
}
pub fn det_input_size_capped(w: u32, h: u32, max_side: u32) -> Option<(u32, u32)> {
if w == 0 || h == 0 {
return None;
}
let (w, h) = (w as f32, h as f32);
let mut ratio = if w.min(h) < LIMIT_SIDE_LEN as f32 {
LIMIT_SIDE_LEN as f32 / w.min(h)
} else {
1.0
};
if max_side > 0 && w.max(h) * ratio > max_side as f32 {
ratio = max_side as f32 / w.max(h);
}
let round32 = |v: f32| ((v as i64 as f32 / 32.0).round() * 32.0) as i64;
let (rw, rh) = (round32(w * ratio), round32(h * ratio));
(rw > 0 && rh > 0).then_some((rw as u32, rh as u32))
}
pub fn prep_det_input(img: &RgbImage) -> Option<(Vec<f32>, u32, u32)> {
let (w, h) = det_input_size(img.width(), img.height())?;
let resized = if (w, h) == img.dimensions() {
img.clone()
} else {
resize_bilinear(img, w, h)
};
let n = (w * h) as usize;
let mut data = vec![0f32; 3 * n];
for (i, px) in resized.pixels().enumerate() {
data[i] = px[2] as f32 / 127.5 - 1.0;
data[n + i] = px[1] as f32 / 127.5 - 1.0;
data[2 * n + i] = px[0] as f32 / 127.5 - 1.0;
}
Some((data, w, h))
}
fn resize_bilinear(img: &RgbImage, w: u32, h: u32) -> RgbImage {
#[cfg(feature = "ml")]
{
use fast_image_resize as fir;
static SLOW: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
let slow = *SLOW.get_or_init(|| docling_core::env::flag("DOCLING_RS_SLOW_RESIZE"));
if !slow {
let fast = || {
let src = fir::images::ImageRef::new(
img.width(),
img.height(),
img.as_raw(),
fir::PixelType::U8x3,
)
.ok()?;
let mut dst = fir::images::Image::new(w, h, fir::PixelType::U8x3);
fir::Resizer::new()
.resize(
&src,
&mut dst,
&fir::ResizeOptions::new()
.resize_alg(fir::ResizeAlg::Convolution(fir::FilterType::Bilinear)),
)
.ok()?;
RgbImage::from_raw(w, h, dst.into_vec())
};
if let Some(out) = fast() {
return out;
}
}
}
image::imageops::resize(img, w, h, image::imageops::FilterType::Triangle)
}
pub fn db_boxes(prob: &[f32], w: usize, h: usize, dest_w: u32, dest_h: u32) -> Vec<DetBox> {
if prob.len() < w * h || w == 0 || h == 0 {
return Vec::new();
}
let seg = |x: usize, y: usize| prob[y * w + x] > THRESH;
let mut mask = vec![false; w * h];
for y in 0..h {
for x in 0..w {
mask[y * w + x] = seg(x, y)
|| (x > 0 && seg(x - 1, y))
|| (y > 0 && seg(x, y - 1))
|| (x > 0 && y > 0 && seg(x - 1, y - 1));
}
}
let mut quads: Vec<([(f32, f32); 4], f32)> = Vec::new();
let mut seen = vec![false; w * h];
let mut stack = Vec::new();
let mut component = Vec::new();
for start in 0..w * h {
if !mask[start] || seen[start] {
continue;
}
component.clear();
seen[start] = true;
stack.push(start);
while let Some(i) = stack.pop() {
component.push(i);
let (x, y) = (i % w, i / w);
for dy in -1i64..=1 {
for dx in -1i64..=1 {
let (nx, ny) = (x as i64 + dx, y as i64 + dy);
if nx < 0 || ny < 0 || nx >= w as i64 || ny >= h as i64 {
continue;
}
let j = ny as usize * w + nx as usize;
if mask[j] && !seen[j] {
seen[j] = true;
stack.push(j);
}
}
}
}
if quads.len() >= 1000 {
break;
}
let boundary: Vec<(f32, f32)> = component
.iter()
.copied()
.filter(|&i| {
let (x, y) = (i % w, i / w);
x == 0
|| y == 0
|| x + 1 == w
|| y + 1 == h
|| !mask[i - 1]
|| !mask[i + 1]
|| !mask[i - w]
|| !mask[i + w]
})
.map(|i| ((i % w) as f32, (i / w) as f32))
.collect();
let Some((corners, sside)) = min_area_rect(&boundary) else {
continue;
};
if sside < MIN_SIZE {
continue;
}
let score = box_score_fast(prob, w, h, &corners);
if score < BOX_THRESH {
continue;
}
let Some((expanded, sside)) = unclip(&corners) else {
continue;
};
if sside < MIN_SIZE + 2.0 {
continue;
}
let mapped: [(f32, f32); 4] = std::array::from_fn(|k| {
let (x, y) = expanded[k];
(
(x / w as f32 * dest_w as f32)
.round()
.clamp(0.0, dest_w as f32),
(y / h as f32 * dest_h as f32)
.round()
.clamp(0.0, dest_h as f32),
)
});
quads.push((mapped, score));
}
let mut boxes: Vec<DetBox> = quads
.into_iter()
.filter_map(|(q, score)| {
let side =
|a: (f32, f32), b: (f32, f32)| ((a.0 - b.0).powi(2) + (a.1 - b.1).powi(2)).sqrt();
let (rw, rh) = (side(q[0], q[1]).floor(), side(q[0], q[3]).floor());
if rw <= 3.0 || rh <= 3.0 {
return None;
}
let xs = q.iter().map(|p| p.0);
let ys = q.iter().map(|p| p.1);
Some(DetBox {
l: xs.clone().fold(f32::MAX, f32::min),
t: ys.clone().fold(f32::MAX, f32::min),
r: xs.fold(f32::MIN, f32::max),
b: ys.fold(f32::MIN, f32::max),
score,
})
})
.collect();
sort_boxes(&mut boxes);
boxes
}
pub fn sort_boxes(boxes: &mut [DetBox]) {
boxes.sort_by(|a, b| a.t.total_cmp(&b.t));
let mut row = 0usize;
let mut rows = Vec::with_capacity(boxes.len());
for i in 0..boxes.len() {
if i > 0 && boxes[i].t - boxes[i - 1].t >= BOX_SORT_Y_THRESHOLD {
row += 1;
}
rows.push(row);
}
let mut order: Vec<usize> = (0..boxes.len()).collect();
order.sort_by(|&a, &b| {
rows[a]
.cmp(&rows[b])
.then(boxes[a].l.total_cmp(&boxes[b].l))
});
let sorted: Vec<DetBox> = order.iter().map(|&i| boxes[i]).collect();
boxes.copy_from_slice(&sorted);
}
type RectCandidate = (f32, [(f32, f32); 4], f32);
pub fn min_area_rect(points: &[(f32, f32)]) -> Option<([(f32, f32); 4], f32)> {
let hull = convex_hull(points);
if hull.is_empty() {
return None;
}
if hull.len() <= 2 {
let (a, b) = (hull[0], *hull.last().unwrap());
return Some((order_corners([a, b, b, a]), 0.0));
}
let mut best: Option<RectCandidate> = None;
for i in 0..hull.len() {
let (p, q) = (hull[i], hull[(i + 1) % hull.len()]);
let (ex, ey) = (q.0 - p.0, q.1 - p.1);
let len = (ex * ex + ey * ey).sqrt();
if len < 1e-6 {
continue;
}
let (ux, uy) = (ex / len, ey / len);
let (vx, vy) = (-uy, ux);
let (mut umin, mut umax, mut vmin, mut vmax) = (f32::MAX, f32::MIN, f32::MAX, f32::MIN);
for &(x, y) in &hull {
let u = x * ux + y * uy;
let v = x * vx + y * vy;
umin = umin.min(u);
umax = umax.max(u);
vmin = vmin.min(v);
vmax = vmax.max(v);
}
let area = (umax - umin) * (vmax - vmin);
if best.as_ref().is_none_or(|(a, _, _)| area < *a) {
let corner = |u: f32, v: f32| (u * ux + v * vx, u * uy + v * vy);
let corners = [
corner(umin, vmin),
corner(umax, vmin),
corner(umax, vmax),
corner(umin, vmax),
];
best = Some((area, corners, (umax - umin).min(vmax - vmin)));
}
}
best.map(|(_, corners, sside)| (order_corners(corners), sside))
}
fn order_corners(mut c: [(f32, f32); 4]) -> [(f32, f32); 4] {
c.sort_by(|a, b| a.0.total_cmp(&b.0));
let (i1, i4) = if c[1].1 > c[0].1 { (0, 1) } else { (1, 0) };
let (i2, i3) = if c[3].1 > c[2].1 { (2, 3) } else { (3, 2) };
[c[i1], c[i2], c[i3], c[i4]]
}
fn convex_hull(points: &[(f32, f32)]) -> Vec<(f32, f32)> {
let mut pts: Vec<(f32, f32)> = points.to_vec();
pts.sort_by(|a, b| a.0.total_cmp(&b.0).then(a.1.total_cmp(&b.1)));
pts.dedup();
if pts.len() < 3 {
return pts;
}
let cross = |o: (f32, f32), a: (f32, f32), b: (f32, f32)| {
(a.0 - o.0) * (b.1 - o.1) - (a.1 - o.1) * (b.0 - o.0)
};
let mut lower: Vec<(f32, f32)> = Vec::new();
for &p in &pts {
while lower.len() >= 2 && cross(lower[lower.len() - 2], lower[lower.len() - 1], p) <= 0.0 {
lower.pop();
}
lower.push(p);
}
let mut upper: Vec<(f32, f32)> = Vec::new();
for &p in pts.iter().rev() {
while upper.len() >= 2 && cross(upper[upper.len() - 2], upper[upper.len() - 1], p) <= 0.0 {
upper.pop();
}
upper.push(p);
}
lower.pop();
upper.pop();
lower.extend(upper);
lower
}
fn box_score_fast(prob: &[f32], w: usize, h: usize, quad: &[(f32, f32); 4]) -> f32 {
let xmin = quad
.iter()
.map(|p| p.0)
.fold(f32::MAX, f32::min)
.floor()
.clamp(0.0, (w - 1) as f32) as usize;
let xmax = quad
.iter()
.map(|p| p.0)
.fold(f32::MIN, f32::max)
.ceil()
.clamp(0.0, (w - 1) as f32) as usize;
let ymin = quad
.iter()
.map(|p| p.1)
.fold(f32::MAX, f32::min)
.floor()
.clamp(0.0, (h - 1) as f32) as usize;
let ymax = quad
.iter()
.map(|p| p.1)
.fold(f32::MIN, f32::max)
.ceil()
.clamp(0.0, (h - 1) as f32) as usize;
let (mut sum, mut n) = (0f64, 0usize);
for y in ymin..=ymax {
for x in xmin..=xmax {
if inside_convex(quad, (x as f32, y as f32)) {
sum += prob[y * w + x] as f64;
n += 1;
}
}
}
if n == 0 {
0.0
} else {
(sum / n as f64) as f32
}
}
fn inside_convex(quad: &[(f32, f32); 4], p: (f32, f32)) -> bool {
let mut pos = false;
let mut neg = false;
for i in 0..4 {
let (a, b) = (quad[i], quad[(i + 1) % 4]);
let cross = (b.0 - a.0) * (p.1 - a.1) - (b.1 - a.1) * (p.0 - a.0);
pos |= cross > 1e-6;
neg |= cross < -1e-6;
}
!(pos && neg)
}
fn unclip(quad: &[(f32, f32); 4]) -> Option<([(f32, f32); 4], f32)> {
let side = |a: (f32, f32), b: (f32, f32)| ((a.0 - b.0).powi(2) + (a.1 - b.1).powi(2)).sqrt();
let (wlen, hlen) = (side(quad[0], quad[1]), side(quad[1], quad[2]));
let perimeter = 2.0 * (wlen + hlen);
if perimeter < 1e-6 {
return None;
}
let d = wlen * hlen * UNCLIP_RATIO / perimeter;
let (cx, cy) = (
quad.iter().map(|p| p.0).sum::<f32>() / 4.0,
quad.iter().map(|p| p.1).sum::<f32>() / 4.0,
);
let (ux, uy) = if wlen > 1e-6 {
(
(quad[1].0 - quad[0].0) / wlen,
(quad[1].1 - quad[0].1) / wlen,
)
} else {
(1.0, 0.0)
};
let (vx, vy) = if hlen > 1e-6 {
(
(quad[2].0 - quad[1].0) / hlen,
(quad[2].1 - quad[1].1) / hlen,
)
} else {
(-uy, ux)
};
let (hw, hh) = (wlen / 2.0 + d, hlen / 2.0 + d);
let corner = |su: f32, sv: f32| {
(
cx + su * hw * ux + sv * hh * vx,
cy + su * hw * uy + sv * hh * vy,
)
};
let corners = order_corners([
corner(-1.0, -1.0),
corner(1.0, -1.0),
corner(1.0, 1.0),
corner(-1.0, 1.0),
]);
Some((corners, (wlen + 2.0 * d).min(hlen + 2.0 * d)))
}
pub fn uncovered_lines(
detected: &[DetBox],
scale: f32,
regions: &[crate::layout::Region],
cells: &[crate::pdfium_backend::TextCell],
) -> Vec<crate::layout::Region> {
let mut accepted: Vec<crate::layout::Region> = Vec::new();
for d in detected.iter().map(|d| crate::layout::Region {
label: "text",
score: d.score,
l: d.l / scale,
t: d.t / scale,
r: d.r / scale,
b: d.b / scale,
}) {
let da = ((d.r - d.l) * (d.b - d.t)).max(1.0);
let inter = |l: f32, t: f32, r: f32, b: f32| {
(d.r.min(r) - d.l.max(l)).max(0.0) * (d.b.min(b) - d.t.max(t)).max(0.0)
};
let in_region = regions.iter().any(|r| {
(crate::ocr_prep::is_text_label(r.label) || crate::assemble::is_table_like(r.label))
&& inter(r.l, r.t, r.r, r.b) / da > 0.5
});
let by_cells: f32 = cells.iter().map(|c| inter(c.l, c.t, c.r, c.b)).sum::<f32>() / da;
let by_accepted = accepted
.iter()
.any(|u| inter(u.l, u.t, u.r, u.b) / da > 0.3);
if !in_region && by_cells <= 0.3 && !by_accepted {
accepted.push(d);
}
}
accepted
}
#[cfg(feature = "ml")]
pub use session::DetModel;
#[cfg(feature = "ml")]
mod session {
use super::{db_boxes, prep_det_input, DetBox};
use image::RgbImage;
use ort::session::Session;
use ort::value::Tensor;
pub struct DetModel {
session: Session,
}
pub(crate) fn resolve_det_path() -> String {
docling_core::env::nonempty("DOCLING_OCR_DET_ONNX")
.unwrap_or_else(|| crate::resolve_asset(".models/ocr_det.onnx"))
}
impl DetModel {
pub fn load(intra: usize) -> Result<Self, String> {
let path = resolve_det_path();
if !std::path::Path::new(&path).exists() {
return Err(format!("text detection model not found at {path}"));
}
let builder = Session::builder()
.map_err(|e| format!("ocr-det: builder: {e}"))?
.with_intra_threads(intra.max(1))
.map_err(|e| format!("ocr-det: intra_threads: {e}"))?;
let builder = docling_onnx::apply(builder).map_err(|e| format!("ocr-det: {e}"))?;
let session = docling_onnx::commit(builder, &path, "det")
.map_err(|e| format!("ocr-det: load {path}: {e}"))?;
Ok(Self { session })
}
pub fn detect(&mut self, img: &RgbImage) -> Result<Vec<DetBox>, String> {
let Some((data, w, h)) = crate::timing::timed("ocr.det.prep", || prep_det_input(img))
else {
return Ok(Vec::new());
};
let input = Tensor::from_array(([1usize, 3, h as usize, w as usize], data))
.map_err(|e| format!("ocr-det: input: {e}"))?;
let name = self.session.inputs()[0].name().to_string();
let outputs = crate::timing::timed("ocr.det.net", || {
self.session
.run(ort::inputs![name.as_str() => input])
.map_err(|e| format!("ocr-det: run: {e}"))
})?;
let (shape, prob) = outputs[0]
.try_extract_tensor::<f32>()
.map_err(|e| format!("ocr-det: output: {e}"))?;
let dims: Vec<usize> = shape.iter().map(|&d| d as usize).collect();
let (ph, pw) = match dims.as_slice() {
[_, _, ph, pw] => (*ph, *pw),
_ => return Err(format!("ocr-det: unexpected output shape {dims:?}")),
};
Ok(crate::timing::timed("ocr.det.post", || {
db_boxes(prob, pw, ph, img.width(), img.height())
}))
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn input_size_scales_the_short_side_to_736_in_multiples_of_32() {
assert_eq!(det_input_size_capped(445, 884, 0), Some((736, 1472)));
assert_eq!(det_input_size_capped(1335, 2652, 0), Some((1344, 2656)));
assert_eq!(det_input_size(0, 10), None);
assert_eq!(det_input_size(1224, 1584), Some((736, 960)));
assert_eq!(det_input_size_capped(1224, 1584, 960), Some((736, 960)));
assert_eq!(det_input_size_capped(445, 884, 1500), Some((736, 1472)));
assert_eq!(det_input_size_capped(445, 884, 960), Some((480, 960)));
assert_eq!(det_input_size_capped(1224, 1584, 0), Some((1216, 1600)));
}
#[test]
fn det_input_is_bgr_normalized() {
let mut img = RgbImage::new(736, 736);
img.put_pixel(0, 0, image::Rgb([255, 0, 128]));
let (data, w, h) = prep_det_input(&img).unwrap();
assert_eq!((w, h), (736, 736));
let n = (w * h) as usize;
assert!((data[0] - (128.0 / 127.5 - 1.0)).abs() < 1e-6);
assert_eq!(data[n], -1.0);
assert_eq!(data[2 * n], 1.0);
}
#[test]
fn min_area_rect_of_an_upright_and_a_tilted_blob() {
let pts: Vec<(f32, f32)> = (0..20)
.flat_map(|x| (0..5).map(move |y| (x as f32, y as f32)))
.collect();
let (c, sside) = min_area_rect(&pts).unwrap();
assert!((sside - 4.0).abs() < 1e-3);
assert!(
(c[0].0 - 0.0).abs() < 1e-3 && (c[0].1 - 0.0).abs() < 1e-3,
"{c:?}"
);
assert!(
(c[2].0 - 19.0).abs() < 1e-3 && (c[2].1 - 4.0).abs() < 1e-3,
"{c:?}"
);
let s = std::f32::consts::FRAC_1_SQRT_2;
let rot: Vec<(f32, f32)> = pts
.iter()
.map(|&(x, y)| (x * s - y * s + 50.0, x * s + y * s + 50.0))
.collect();
let (_, sside) = min_area_rect(&rot).unwrap();
assert!((sside - 4.0).abs() < 1e-2, "{sside}");
}
#[test]
fn db_boxes_from_a_synthetic_probability_map() {
let (w, h) = (128usize, 64usize);
let mut prob = vec![0f32; w * h];
let blob = |prob: &mut Vec<f32>, l: usize, t: usize, r: usize, b: usize, p: f32| {
for y in t..b {
for x in l..r {
prob[y * w + x] = p;
}
}
};
blob(&mut prob, 70, 10, 110, 20, 0.9); blob(&mut prob, 10, 12, 50, 22, 0.9); blob(&mut prob, 10, 40, 60, 48, 0.35); blob(&mut prob, 100, 50, 102, 52, 0.9); let boxes = db_boxes(&prob, w, h, 256, 128);
assert_eq!(boxes.len(), 2, "{boxes:?}");
let a = &boxes[0];
assert!(a.l < boxes[1].l);
assert!(a.score > 0.75 && a.score < 0.9, "{}", a.score);
assert!(
(a.l - 7.0).abs() <= 2.0 && (a.r - 113.0).abs() <= 2.0,
"{a:?}"
);
assert!(
(a.t - 11.0).abs() <= 2.0 && (a.b - 57.0).abs() <= 2.0,
"{a:?}"
);
}
#[test]
fn uncovered_lines_skip_what_the_region_pass_read() {
use crate::layout::Region;
use crate::pdfium_backend::TextCell;
let bx = |l: f32, t: f32, r: f32, b: f32| DetBox {
l,
t,
r,
b,
score: 0.9,
};
let regions = vec![Region {
label: "text",
score: 0.9,
l: 0.0,
t: 0.0,
r: 100.0,
b: 20.0,
}];
let cell = |l: f32, r: f32| TextCell {
text: "x".into(),
l,
t: 50.0,
r,
b: 60.0,
};
let cells = vec![cell(0.0, 50.0), cell(50.0, 100.0)];
let detected = vec![
bx(0.0, 0.0, 200.0, 40.0), bx(0.0, 100.0, 200.0, 120.0), bx(0.0, 300.0, 200.0, 320.0), bx(20.0, 302.0, 100.0, 318.0), ];
let out = uncovered_lines(&detected, 2.0, ®ions, &cells);
assert_eq!(out.len(), 1, "{out:?}");
assert_eq!(
(out[0].l, out[0].t, out[0].r, out[0].b),
(0.0, 150.0, 100.0, 160.0)
);
assert_eq!(out[0].label, "text");
}
#[test]
fn boxes_sort_by_row_then_column() {
let bx = |l: f32, t: f32| DetBox {
l,
t,
r: l + 10.0,
b: t + 10.0,
score: 1.0,
};
let mut boxes = vec![
bx(50.0, 100.0),
bx(10.0, 105.0),
bx(30.0, 20.0),
bx(5.0, 200.0),
];
sort_boxes(&mut boxes);
let order: Vec<(f32, f32)> = boxes.iter().map(|b| (b.l, b.t)).collect();
assert_eq!(
order,
vec![(30.0, 20.0), (10.0, 105.0), (50.0, 100.0), (5.0, 200.0)]
);
}
}