1use image::RgbImage;
21use ort::session::Session;
22use ort::value::Tensor;
23use tokenizers::Tokenizer;
24
25use docling_core::PictureClass;
26
27pub const CLASSIFIER_SCALE: f32 = 2.0;
30pub const CODE_FORMULA_SCALE: f32 = 1.67;
31pub const CODE_FORMULA_EXPANSION: f32 = 0.18;
34
35const PICTURE_CLASSES: [&str; 26] = [
42 "logo",
43 "photograph",
44 "icon",
45 "engineering_drawing",
46 "line_chart",
47 "bar_chart",
48 "other",
49 "table",
50 "flow_chart",
51 "screenshot_from_computer",
52 "signature",
53 "screenshot_from_manual",
54 "geographical_map",
55 "pie_chart",
56 "page_thumbnail",
57 "stamp",
58 "music",
59 "calendar",
60 "qr_code",
61 "bar_code",
62 "full_page_image",
63 "scatter_plot",
64 "chemistry_structure",
65 "topographical_map",
66 "crossword_puzzle",
67 "box_plot",
68];
69
70const CLASSIFIER_SIDE: u32 = 224;
71const CLASSIFIER_MEAN: [f32; 3] = [0.485, 0.456, 0.406];
73const CLASSIFIER_STD: [f32; 3] = [0.478_539_44, 0.473_286_4, 0.474_341_63];
74
75pub struct PictureClassifier {
76 session: Session,
77}
78
79impl PictureClassifier {
80 pub fn load_with(intra: usize) -> Option<Self> {
84 let path = crate::model_path(
85 "DOCLING_PICTURE_CLASSIFIER_ONNX",
86 ".models/picture_classifier.onnx",
87 ".models/picture_classifier_int8.onnx",
88 );
89 if !std::path::Path::new(&path).exists() {
90 eprintln!(
91 "docling-pdf: picture classifier model not found ({path}); \
92 picture classification skipped. Run scripts/install/download_dependencies.sh."
93 );
94 return None;
95 }
96 let builder = docling_onnx::session_builder()
97 .map_err(|e| eprintln!("docling-pdf: picture classifier: {e}"))
98 .ok()?
99 .with_intra_threads(intra)
100 .ok()?;
101 let builder = docling_onnx::apply(builder)
102 .and_then(docling_onnx::cap_before_1_29)
103 .map_err(|e| eprintln!("docling-pdf: picture classifier: {e}"))
104 .ok()?;
105 let session = docling_onnx::commit_uncached(builder, &path)
106 .map_err(|e| eprintln!("docling-pdf: picture classifier load {path}: {e}"))
107 .ok()?;
108 Some(Self { session })
109 }
110
111 pub fn classify(&mut self, crop: &RgbImage) -> Result<Vec<PictureClass>, String> {
114 let resized = image::imageops::resize(
115 crop,
116 CLASSIFIER_SIDE,
117 CLASSIFIER_SIDE,
118 image::imageops::FilterType::Triangle,
119 );
120 let n = (CLASSIFIER_SIDE * CLASSIFIER_SIDE) as usize;
121 let mut data = vec![0f32; 3 * n];
122 for (i, px) in resized.pixels().enumerate() {
123 for c in 0..3 {
124 data[c * n + i] = (px[c] as f32 / 255.0 - CLASSIFIER_MEAN[c]) / CLASSIFIER_STD[c];
125 }
126 }
127 let input = Tensor::from_array((
128 [
129 1usize,
130 3,
131 CLASSIFIER_SIDE as usize,
132 CLASSIFIER_SIDE as usize,
133 ],
134 data,
135 ))
136 .map_err(|e| format!("picture classifier: input: {e}"))?;
137 let outputs = self
138 .session
139 .run(ort::inputs!["input" => input])
140 .map_err(|e| format!("picture classifier: inference: {e}"))?;
141 let (_, logits) = outputs[0]
142 .try_extract_tensor::<f32>()
143 .map_err(|e| format!("picture classifier: output: {e}"))?;
144 let max = logits.iter().copied().fold(f32::NEG_INFINITY, f32::max);
146 let exp: Vec<f32> = logits.iter().map(|&v| (v - max).exp()).collect();
147 let sum: f32 = exp.iter().sum();
148 let mut preds: Vec<PictureClass> = exp
149 .iter()
150 .enumerate()
151 .map(|(i, &e)| PictureClass {
152 class_name: PICTURE_CLASSES.get(i).copied().unwrap_or("other").into(),
153 confidence: e / sum,
154 })
155 .collect();
156 preds.sort_by(|a, b| b.confidence.total_cmp(&a.confidence));
157 Ok(preds)
158 }
159}
160
161#[derive(Debug, Clone, Copy, PartialEq, Eq)]
167pub enum CodeFormulaKind {
168 Code,
169 Formula,
170}
171
172const TILE: u32 = 512; const LONGEST_EDGE: u32 = 2048; const MAX_IMAGE_SIZE: u32 = 4096; const IMAGE_SEQ_LEN: usize = 64; const IMAGE_TOKEN_ID: i64 = 100270; const EOS_ID: i64 = 100338; const MODEL_MAX_LEN: usize = 8192;
181const HIDDEN: usize = 576;
182const N_LAYERS: usize = 30;
183const N_KV: usize = 3;
184const HEAD_DIM: usize = 64;
185
186pub struct CodeFormula {
187 vision: Session,
188 embed: Session,
189 decoder: Session,
190 tokenizer: Tokenizer,
191}
192
193impl CodeFormula {
194 pub fn load_with(intra: usize) -> Option<Self> {
198 let dir = docling_core::env::nonempty("DOCLING_CODE_FORMULA_DIR")
199 .unwrap_or_else(|| crate::resolve_asset(".models/code_formula"));
200 let file = |name: &str| format!("{dir}/{name}");
201 let graph = |base: &str| {
203 let int8 = file(&format!("{base}_int8.onnx"));
204 if !crate::prefer_fp32() && std::path::Path::new(&int8).exists() {
205 int8
206 } else {
207 file(&format!("{base}.onnx"))
208 }
209 };
210 for f in [&graph("vision"), &graph("embed"), &graph("decoder_kv")] {
211 if !std::path::Path::new(f.as_str()).exists() {
212 eprintln!(
213 "docling-pdf: CodeFormula model not found ({f}); code/formula \
214 enrichment skipped. Run scripts/install/download_dependencies.sh."
215 );
216 return None;
217 }
218 }
219 let load = |p: String| {
220 let builder = docling_onnx::session_builder()
221 .map_err(|e| eprintln!("docling-pdf: CodeFormula: {e}"))
222 .ok()?
223 .with_intra_threads(intra)
224 .ok()?;
225 let builder = docling_onnx::apply(builder)
226 .map_err(|e| eprintln!("docling-pdf: CodeFormula: {e}"))
227 .ok()?;
228 docling_onnx::commit_uncached(builder, &p)
229 .map_err(|e| eprintln!("docling-pdf: CodeFormula load {p}: {e}"))
230 .ok()
231 };
232 let tokenizer = Tokenizer::from_file(file("tokenizer.json"))
233 .map_err(|e| eprintln!("docling-pdf: CodeFormula tokenizer: {e}"))
234 .ok()?;
235 Some(Self {
236 vision: load(graph("vision"))?,
237 embed: load(graph("embed"))?,
238 decoder: load(graph("decoder_kv"))?,
239 tokenizer,
240 })
241 }
242
243 pub fn predict(&mut self, crop: &RgbImage, kind: CodeFormulaKind) -> Result<String, String> {
247 if let Some(dir) = docling_core::env::nonempty("DOCLING_RS_ENRICH_DEBUG") {
250 use std::sync::atomic::{AtomicUsize, Ordering};
251 static N: AtomicUsize = AtomicUsize::new(0);
252 let n = N.fetch_add(1, Ordering::Relaxed);
253 let _ = crop.save(format!("{dir}/rs_crop_{n}.png"));
254 }
255 let (tiles, rows, cols) = preprocess_idefics3(crop);
256 let n_tiles = tiles.len() / (3 * (TILE * TILE) as usize);
257
258 let feats: Vec<f32> = {
261 let input = Tensor::from_array(([n_tiles, 3, TILE as usize, TILE as usize], tiles))
262 .map_err(|e| format!("code-formula: vision input: {e}"))?;
263 let outputs = self
264 .vision
265 .run(ort::inputs!["pixel_values" => input])
266 .map_err(|e| format!("code-formula: vision: {e}"))?;
267 let (_, feats) = outputs["image_features"]
268 .try_extract_tensor::<f32>()
269 .map_err(|e| format!("code-formula: vision output: {e}"))?;
270 feats.to_vec()
271 };
272
273 let query = match kind {
276 CodeFormulaKind::Code => "<code>",
277 CodeFormulaKind::Formula => "<formula>",
278 };
279 let prompt = format!(
280 "<|start_of_role|>user:{}{query}<end_of_utterance>\nassistant:",
281 image_prompt(rows, cols)
282 );
283 let enc = self
284 .tokenizer
285 .encode(prompt, false)
286 .map_err(|e| format!("code-formula: tokenize: {e}"))?;
287 let ids: Vec<i64> = enc.get_ids().iter().map(|&v| v as i64).collect();
288 let seq = ids.len();
289
290 let mut embeds = self.embed_ids(&ids)?;
293 let image_positions: Vec<usize> = ids
294 .iter()
295 .enumerate()
296 .filter(|(_, &t)| t == IMAGE_TOKEN_ID)
297 .map(|(i, _)| i)
298 .collect();
299 if image_positions.len() != n_tiles * IMAGE_SEQ_LEN {
300 return Err(format!(
301 "code-formula: {} image tokens for {} tiles",
302 image_positions.len(),
303 n_tiles
304 ));
305 }
306 for (v, &pos) in image_positions.iter().enumerate() {
307 embeds[pos * HIDDEN..(pos + 1) * HIDDEN]
308 .copy_from_slice(&feats[v * HIDDEN..(v + 1) * HIDDEN]);
309 }
310
311 let mut cache: Option<(ort::value::DynValue, ort::value::DynValue)> = None;
319 let empty = {
320 let mk = || {
321 Tensor::<f32>::new(
322 self.decoder.allocator(),
323 [N_LAYERS, 1, N_KV, 0usize, HEAD_DIM],
324 )
325 .map_err(|e| format!("code-formula: empty kv cache: {e}"))
326 };
327 (mk()?, mk()?)
328 };
329 let mut past_len = 0usize;
330 let mut positions: Vec<i64> = (0..seq as i64).collect();
331 let mut x = embeds;
332 let mut x_seq = seq;
333 let mut out_ids: Vec<u32> = Vec::new();
334 let max_new = MODEL_MAX_LEN.saturating_sub(seq);
335 for _ in 0..max_new {
336 let embeds_t = Tensor::from_array(([1usize, x_seq, HIDDEN], x))
337 .map_err(|e| format!("code-formula: embeds: {e}"))?;
338 let pos_t = Tensor::from_array(([1usize, positions.len()], positions.clone()))
339 .map_err(|e| format!("code-formula: positions: {e}"))?;
340 let next = {
341 let mut out = match cache.as_ref() {
342 Some((k, v)) => self.decoder.run(ort::inputs![
343 "inputs_embeds" => embeds_t, "position_ids" => pos_t,
344 "past_k" => k, "past_v" => v]),
345 None => self.decoder.run(ort::inputs![
346 "inputs_embeds" => embeds_t, "position_ids" => pos_t,
347 "past_k" => &empty.0, "past_v" => &empty.1]),
348 }
349 .map_err(|e| format!("code-formula: decoder: {e}"))?;
350 let (_, logits) = out["logits"]
351 .try_extract_tensor::<f32>()
352 .map_err(|e| format!("code-formula: logits: {e}"))?;
353 let next = logits
354 .iter()
355 .enumerate()
356 .max_by(|a, b| a.1.total_cmp(b.1))
357 .map(|(i, _)| i as i64)
358 .unwrap_or(EOS_ID);
359 cache = Some((
360 out.remove("new_k")
361 .ok_or_else(|| "code-formula: new_k missing".to_string())?,
362 out.remove("new_v")
363 .ok_or_else(|| "code-formula: new_v missing".to_string())?,
364 ));
365 next
366 };
367 past_len += x_seq;
368 if next == EOS_ID {
369 break;
370 }
371 out_ids.push(next as u32);
372 x = self.embed_ids(&[next])?;
373 x_seq = 1;
374 positions = vec![past_len as i64];
375 }
376
377 let text = self
378 .tokenizer
379 .decode(&out_ids, false)
380 .map_err(|e| format!("code-formula: decode: {e}"))?;
381 Ok(post_process(&text))
382 }
383
384 fn embed_ids(&mut self, ids: &[i64]) -> Result<Vec<f32>, String> {
385 let input = Tensor::from_array(([1usize, ids.len()], ids.to_vec()))
386 .map_err(|e| format!("code-formula: ids: {e}"))?;
387 let out = self
388 .embed
389 .run(ort::inputs!["input_ids" => input])
390 .map_err(|e| format!("code-formula: embed: {e}"))?;
391 let (_, embeds) = out["inputs_embeds"]
392 .try_extract_tensor::<f32>()
393 .map_err(|e| format!("code-formula: embed output: {e}"))?;
394 Ok(embeds.to_vec())
395 }
396}
397
398fn image_prompt(rows: u32, cols: u32) -> String {
401 let img = "<image>".repeat(IMAGE_SEQ_LEN);
402 let mut s = String::new();
403 for r in 1..=rows {
404 for c in 1..=cols {
405 s.push_str(&format!("<fake_token_around_image><row_{r}_col_{c}>{img}"));
406 }
407 s.push('\n');
408 }
409 s.push_str(&format!(
410 "\n<fake_token_around_image><global-img>{img}<fake_token_around_image>"
411 ));
412 s
413}
414
415fn preprocess_idefics3(crop: &RgbImage) -> (Vec<f32>, u32, u32) {
420 use image::imageops::FilterType;
421 let (w0, h0) = crop.dimensions();
424 let (mut w, mut h) = rescale_to_max_len(w0, h0, LONGEST_EDGE);
425 (h, w) = scale_below_upper_bound(h, w, MAX_IMAGE_SIZE);
426 let img = image::imageops::resize(crop, w, h, FilterType::Lanczos3);
427
428 let (tw, th) = if w >= h {
430 let tw = w.div_ceil(TILE) * TILE;
431 let th0 = (tw as f64 / (w as f64 / h as f64)) as u32;
432 (tw, th0.div_ceil(TILE) * TILE)
433 } else {
434 let th = h.div_ceil(TILE) * TILE;
435 let tw0 = (th as f64 * (w as f64 / h as f64)) as u32;
436 (tw0.div_ceil(TILE) * TILE, th)
437 };
438 let img = image::imageops::resize(&img, tw, th, FilterType::Lanczos3);
439
440 let (rows, cols) = (th / TILE, tw / TILE);
442 let mut tensor = Vec::with_capacity(((rows * cols + 1) * 3 * TILE * TILE) as usize);
443 for r in 0..rows {
444 for c in 0..cols {
445 let tile = image::imageops::crop_imm(&img, c * TILE, r * TILE, TILE, TILE).to_image();
446 push_normalized(&mut tensor, &tile);
447 }
448 }
449 let global = image::imageops::resize(&img, TILE, TILE, FilterType::Lanczos3);
450 push_normalized(&mut tensor, &global);
451 (tensor, rows, cols)
452}
453
454fn rescale_to_max_len(w0: u32, h0: u32, max_len: u32) -> (u32, u32) {
457 let aspect = w0 as f64 / h0 as f64;
458 let (w, h) = if w0 >= h0 {
459 let w = max_len;
460 let mut h = (w as f64 / aspect) as u32;
461 if !h.is_multiple_of(2) {
462 h += 1;
463 }
464 (w, h)
465 } else {
466 let h = max_len;
467 let mut w = (h as f64 * aspect) as u32;
468 if !w.is_multiple_of(2) {
469 w += 1;
470 }
471 (w, h)
472 };
473 (w.max(1), h.max(1))
474}
475
476fn scale_below_upper_bound(h0: u32, w0: u32, max_len: u32) -> (u32, u32) {
479 let aspect = w0 as f64 / h0 as f64;
480 let (h, w) = if w0 >= h0 && w0 > max_len {
481 let w = max_len;
482 (((w as f64 / aspect) as u32).max(1), w)
483 } else if h0 > w0 && h0 > max_len {
484 let h = max_len;
485 (h, ((h as f64 * aspect) as u32).max(1))
486 } else {
487 (h0, w0)
488 };
489 (h.max(1), w.max(1))
490}
491
492fn push_normalized(tensor: &mut Vec<f32>, tile: &RgbImage) {
494 let n = (TILE * TILE) as usize;
495 let base = tensor.len();
496 tensor.resize(base + 3 * n, 0.0);
497 for (i, px) in tile.pixels().enumerate() {
498 for c in 0..3 {
499 tensor[base + c * n + i] = px[c] as f32 / 255.0 * 2.0 - 1.0;
500 }
501 }
502}
503
504fn post_process(text: &str) -> String {
508 let mut t = match text.find("<end_of_utterance>") {
509 Some(i) => &text[..i],
510 None => text,
511 }
512 .to_string();
513 for tok in ["</code>", "</formula>", "<loc_0><loc_0><loc_500><loc_500>"] {
514 t = t.replace(tok, "");
515 }
516 t.trim_start().to_string()
517}
518
519pub fn extract_code_language(s: &str) -> (String, Option<String>) {
522 let rest = match s.strip_prefix("<_") {
523 Some(r) => r,
524 None => return (s.to_string(), None),
525 };
526 match rest.find("_>") {
529 Some(end) if !rest[..end].is_empty() && !rest[..end].contains(['_', '>']) => {
530 let lang = rest[..end].to_string();
531 let remainder = rest[end + 2..].trim_start().to_string();
532 (remainder, Some(lang))
533 }
534 _ => (s.to_string(), None),
535 }
536}
537
538#[cfg(test)]
539mod tests {
540 use super::*;
541
542 #[test]
543 fn language_prefix_extraction() {
544 assert_eq!(
545 extract_code_language("<_JavaScript_> function f() {}"),
546 (
547 "function f() {}".to_string(),
548 Some("JavaScript".to_string())
549 )
550 );
551 assert_eq!(
552 extract_code_language("plain text"),
553 ("plain text".to_string(), None)
554 );
555 assert_eq!(
556 extract_code_language("<_x_y_> t"),
557 ("<_x_y_> t".to_string(), None)
558 );
559 }
560
561 #[test]
562 fn idefics3_grid_matches_processor() {
563 let img = RgbImage::new(800, 300);
566 let (tensor, rows, cols) = preprocess_idefics3(&img);
567 assert_eq!((rows, cols), (2, 4));
568 assert_eq!(tensor.len(), 9 * 3 * 512 * 512);
569 }
570
571 #[test]
572 fn image_prompt_layout() {
573 let p = image_prompt(1, 2);
574 assert!(p.starts_with("<fake_token_around_image><row_1_col_1><image>"));
575 assert!(p.contains("<row_1_col_2>"));
576 let tail = "<fake_token_around_image><global-img>".to_owned()
577 + &"<image>".repeat(64)
578 + "<fake_token_around_image>";
579 assert!(p.ends_with(&tail));
580 assert_eq!(p.matches("<image>").count(), 3 * 64);
581 }
582}