1use super::validation::ensure_non_empty_images;
6use crate::ConfigValidator;
7use crate::core::OCRError;
8use crate::core::traits::TaskDefinition;
9use crate::core::traits::task::{ImageTaskInput, Task, TaskType};
10use crate::processors::BoundingBox;
11use crate::utils::{ScoreValidator, validate_max_value};
12use serde::{Deserialize, Serialize};
13use std::collections::HashMap;
14
15#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Default)]
17pub enum MergeBboxMode {
18 #[default]
20 Large,
21 Union,
23 Small,
25}
26
27#[derive(Debug, Clone, Serialize, Deserialize)]
30pub enum UnclipRatio {
31 Uniform(f32),
33 Separate(f32, f32),
35 PerClass(HashMap<usize, (f32, f32)>),
37}
38
39impl Default for UnclipRatio {
40 fn default() -> Self {
41 UnclipRatio::Separate(1.0, 1.0)
42 }
43}
44
45#[derive(Debug, Clone, Serialize, Deserialize, ConfigValidator)]
47pub struct LayoutDetectionConfig {
48 #[validate(range(min = 0.0, max = 1.0))]
52 pub score_threshold: f32,
53 #[validate(min = 1)]
55 pub max_elements: usize,
56 #[serde(default)]
64 pub class_thresholds: Option<HashMap<String, f32>>,
65 #[serde(default)]
67 pub class_merge_modes: Option<HashMap<String, MergeBboxMode>>,
68 #[serde(default = "default_layout_nms")]
70 pub layout_nms: bool,
71 #[serde(default = "default_nms_threshold")]
73 pub nms_threshold: f32,
74 #[serde(default)]
77 pub layout_unclip_ratio: Option<UnclipRatio>,
78}
79
80fn default_layout_nms() -> bool {
81 true
82}
83
84fn default_nms_threshold() -> f32 {
85 0.5
86}
87
88impl Default for LayoutDetectionConfig {
89 fn default() -> Self {
90 Self {
91 score_threshold: 0.5,
92 max_elements: 100,
93 class_thresholds: None,
94 class_merge_modes: None,
95 layout_nms: true,
96 nms_threshold: 0.5,
97 layout_unclip_ratio: None,
98 }
99 }
100}
101
102impl LayoutDetectionConfig {
103 pub fn with_pp_structurev3_thresholds() -> Self {
112 let mut class_thresholds = HashMap::new();
113 class_thresholds.insert("paragraph_title".to_string(), 0.3);
114 class_thresholds.insert("formula".to_string(), 0.3);
115 class_thresholds.insert("text".to_string(), 0.4);
116 class_thresholds.insert("seal".to_string(), 0.45);
117
118 Self {
119 score_threshold: 0.3,
122 max_elements: 100,
123 class_thresholds: Some(class_thresholds),
124 class_merge_modes: None,
125 layout_nms: true,
126 nms_threshold: 0.5,
127 layout_unclip_ratio: Some(UnclipRatio::Separate(1.0, 1.0)),
128 }
129 }
130
131 pub fn with_pp_doclayoutv2_defaults() -> Self {
141 let mut class_thresholds = HashMap::new();
142 class_thresholds.insert("abstract".to_string(), 0.5);
143 class_thresholds.insert("algorithm".to_string(), 0.5);
144 class_thresholds.insert("aside_text".to_string(), 0.5);
145 class_thresholds.insert("chart".to_string(), 0.5);
146 class_thresholds.insert("content".to_string(), 0.5);
147 class_thresholds.insert("display_formula".to_string(), 0.4);
148 class_thresholds.insert("doc_title".to_string(), 0.4);
149 class_thresholds.insert("figure_title".to_string(), 0.5);
150 class_thresholds.insert("footer".to_string(), 0.5);
151 class_thresholds.insert("footer_image".to_string(), 0.5);
152 class_thresholds.insert("footnote".to_string(), 0.5);
153 class_thresholds.insert("formula_number".to_string(), 0.5);
154 class_thresholds.insert("header".to_string(), 0.5);
155 class_thresholds.insert("header_image".to_string(), 0.5);
156 class_thresholds.insert("image".to_string(), 0.5);
157 class_thresholds.insert("inline_formula".to_string(), 0.4);
158 class_thresholds.insert("number".to_string(), 0.5);
159 class_thresholds.insert("paragraph_title".to_string(), 0.4);
160 class_thresholds.insert("reference".to_string(), 0.5);
161 class_thresholds.insert("reference_content".to_string(), 0.5);
162 class_thresholds.insert("seal".to_string(), 0.45);
163 class_thresholds.insert("table".to_string(), 0.5);
164 class_thresholds.insert("text".to_string(), 0.4);
165 class_thresholds.insert("vertical_text".to_string(), 0.4);
166 class_thresholds.insert("vision_footnote".to_string(), 0.5);
167
168 let mut merge_modes = HashMap::new();
169 merge_modes.insert("abstract".to_string(), MergeBboxMode::Union);
170 merge_modes.insert("algorithm".to_string(), MergeBboxMode::Union);
171 merge_modes.insert("aside_text".to_string(), MergeBboxMode::Union);
172 merge_modes.insert("chart".to_string(), MergeBboxMode::Large);
173 merge_modes.insert("content".to_string(), MergeBboxMode::Union);
174 merge_modes.insert("display_formula".to_string(), MergeBboxMode::Large);
175 merge_modes.insert("doc_title".to_string(), MergeBboxMode::Large);
176 merge_modes.insert("figure_title".to_string(), MergeBboxMode::Union);
177 merge_modes.insert("footer".to_string(), MergeBboxMode::Union);
178 merge_modes.insert("footer_image".to_string(), MergeBboxMode::Union);
179 merge_modes.insert("footnote".to_string(), MergeBboxMode::Union);
180 merge_modes.insert("formula_number".to_string(), MergeBboxMode::Union);
181 merge_modes.insert("header".to_string(), MergeBboxMode::Union);
182 merge_modes.insert("header_image".to_string(), MergeBboxMode::Union);
183 merge_modes.insert("image".to_string(), MergeBboxMode::Union);
184 merge_modes.insert("inline_formula".to_string(), MergeBboxMode::Large);
185 merge_modes.insert("number".to_string(), MergeBboxMode::Union);
186 merge_modes.insert("paragraph_title".to_string(), MergeBboxMode::Large);
187 merge_modes.insert("reference".to_string(), MergeBboxMode::Union);
188 merge_modes.insert("reference_content".to_string(), MergeBboxMode::Union);
189 merge_modes.insert("seal".to_string(), MergeBboxMode::Union);
190 merge_modes.insert("table".to_string(), MergeBboxMode::Union);
191 merge_modes.insert("text".to_string(), MergeBboxMode::Union);
192 merge_modes.insert("vertical_text".to_string(), MergeBboxMode::Union);
193 merge_modes.insert("vision_footnote".to_string(), MergeBboxMode::Union);
194
195 Self {
196 score_threshold: 0.4,
197 max_elements: 100,
198 class_thresholds: Some(class_thresholds),
199 class_merge_modes: Some(merge_modes),
200 layout_nms: true,
201 nms_threshold: 0.5,
202 layout_unclip_ratio: Some(UnclipRatio::Separate(1.0, 1.0)),
203 }
204 }
205
206 pub fn with_pp_doclayoutv3_defaults() -> Self {
211 let mut merge_modes = HashMap::new();
212 merge_modes.insert("abstract".to_string(), MergeBboxMode::Union);
213 merge_modes.insert("algorithm".to_string(), MergeBboxMode::Union);
214 merge_modes.insert("aside_text".to_string(), MergeBboxMode::Union);
215 merge_modes.insert("chart".to_string(), MergeBboxMode::Large);
216 merge_modes.insert("content".to_string(), MergeBboxMode::Union);
217 merge_modes.insert("display_formula".to_string(), MergeBboxMode::Large);
218 merge_modes.insert("doc_title".to_string(), MergeBboxMode::Large);
219 merge_modes.insert("figure_title".to_string(), MergeBboxMode::Union);
220 merge_modes.insert("footer".to_string(), MergeBboxMode::Union);
221 merge_modes.insert("footer_image".to_string(), MergeBboxMode::Union);
222 merge_modes.insert("footnote".to_string(), MergeBboxMode::Union);
223 merge_modes.insert("formula_number".to_string(), MergeBboxMode::Union);
224 merge_modes.insert("header".to_string(), MergeBboxMode::Union);
225 merge_modes.insert("header_image".to_string(), MergeBboxMode::Union);
226 merge_modes.insert("image".to_string(), MergeBboxMode::Union);
227 merge_modes.insert("inline_formula".to_string(), MergeBboxMode::Large);
228 merge_modes.insert("number".to_string(), MergeBboxMode::Union);
229 merge_modes.insert("paragraph_title".to_string(), MergeBboxMode::Large);
230 merge_modes.insert("reference".to_string(), MergeBboxMode::Union);
231 merge_modes.insert("reference_content".to_string(), MergeBboxMode::Union);
232 merge_modes.insert("seal".to_string(), MergeBboxMode::Union);
233 merge_modes.insert("table".to_string(), MergeBboxMode::Union);
234 merge_modes.insert("text".to_string(), MergeBboxMode::Union);
235 merge_modes.insert("vertical_text".to_string(), MergeBboxMode::Union);
236 merge_modes.insert("vision_footnote".to_string(), MergeBboxMode::Union);
237
238 Self {
239 score_threshold: 0.3,
240 max_elements: 100,
241 class_thresholds: None,
242 class_merge_modes: Some(merge_modes),
243 layout_nms: true,
244 nms_threshold: 0.5,
245 layout_unclip_ratio: Some(UnclipRatio::Separate(1.0, 1.0)),
246 }
247 }
248
249 pub fn with_pp_structurev3_defaults() -> Self {
255 let mut cfg = Self::with_pp_structurev3_thresholds();
256
257 let mut merge_modes = HashMap::new();
258 merge_modes.insert("paragraph_title".to_string(), MergeBboxMode::Large);
259 merge_modes.insert("image".to_string(), MergeBboxMode::Large);
260 merge_modes.insert("text".to_string(), MergeBboxMode::Union);
261 merge_modes.insert("number".to_string(), MergeBboxMode::Union);
262 merge_modes.insert("abstract".to_string(), MergeBboxMode::Union);
263 merge_modes.insert("content".to_string(), MergeBboxMode::Union);
264 merge_modes.insert("figure_table_chart_title".to_string(), MergeBboxMode::Union);
265 merge_modes.insert("formula".to_string(), MergeBboxMode::Large);
266 merge_modes.insert("table".to_string(), MergeBboxMode::Union);
267 merge_modes.insert("reference".to_string(), MergeBboxMode::Union);
268 merge_modes.insert("doc_title".to_string(), MergeBboxMode::Union);
269 merge_modes.insert("footnote".to_string(), MergeBboxMode::Union);
270 merge_modes.insert("header".to_string(), MergeBboxMode::Union);
271 merge_modes.insert("algorithm".to_string(), MergeBboxMode::Union);
272 merge_modes.insert("footer".to_string(), MergeBboxMode::Union);
273 merge_modes.insert("seal".to_string(), MergeBboxMode::Union);
274 merge_modes.insert("chart".to_string(), MergeBboxMode::Large);
275 merge_modes.insert("formula_number".to_string(), MergeBboxMode::Union);
276 merge_modes.insert("aside_text".to_string(), MergeBboxMode::Union);
277 merge_modes.insert("reference_content".to_string(), MergeBboxMode::Union);
278
279 cfg.class_merge_modes = Some(merge_modes);
280 cfg.layout_unclip_ratio = Some(UnclipRatio::Separate(1.0, 1.0));
281 cfg
282 }
283
284 pub fn get_class_threshold(&self, class_name: &str) -> f32 {
288 self.class_thresholds
289 .as_ref()
290 .and_then(|thresholds| thresholds.get(class_name).copied())
291 .unwrap_or(self.score_threshold)
292 }
293
294 pub fn get_class_merge_mode(&self, class_name: &str) -> MergeBboxMode {
298 self.class_merge_modes
299 .as_ref()
300 .and_then(|modes| modes.get(class_name).copied())
301 .unwrap_or(MergeBboxMode::Large)
302 }
303}
304
305#[derive(Debug, Clone)]
310pub struct LayoutDetectionElement {
311 pub bbox: BoundingBox,
313 pub element_type: String,
315 pub score: f32,
317}
318
319#[derive(Debug, Clone)]
321pub struct LayoutDetectionOutput {
322 pub elements: Vec<Vec<LayoutDetectionElement>>,
324 pub is_reading_order_sorted: bool,
329}
330
331impl LayoutDetectionOutput {
332 pub fn empty() -> Self {
334 Self {
335 elements: Vec::new(),
336 is_reading_order_sorted: false,
337 }
338 }
339
340 pub fn with_capacity(capacity: usize) -> Self {
342 Self {
343 elements: Vec::with_capacity(capacity),
344 is_reading_order_sorted: false,
345 }
346 }
347
348 pub fn with_reading_order_sorted(mut self, sorted: bool) -> Self {
350 self.is_reading_order_sorted = sorted;
351 self
352 }
353}
354
355impl TaskDefinition for LayoutDetectionOutput {
356 const TASK_NAME: &'static str = "layout_detection";
357 const TASK_DOC: &'static str = "Layout detection/analysis";
358
359 fn empty() -> Self {
360 LayoutDetectionOutput::empty()
361 }
362}
363
364#[derive(Debug, Default)]
366pub struct LayoutDetectionTask {
367 config: LayoutDetectionConfig,
368}
369
370impl LayoutDetectionTask {
371 pub fn new(config: LayoutDetectionConfig) -> Self {
373 Self { config }
374 }
375}
376
377impl Task for LayoutDetectionTask {
378 type Config = LayoutDetectionConfig;
379 type Input = ImageTaskInput;
380 type Output = LayoutDetectionOutput;
381
382 fn task_type(&self) -> TaskType {
383 TaskType::LayoutDetection
384 }
385
386 fn validate_input(&self, input: &Self::Input) -> Result<(), OCRError> {
387 ensure_non_empty_images(&input.images, "No images provided for layout detection")?;
388
389 Ok(())
390 }
391
392 fn validate_output(&self, output: &Self::Output) -> Result<(), OCRError> {
393 let validator = ScoreValidator::new_unit_range("score");
394
395 for (idx, elements) in output.elements.iter().enumerate() {
396 validate_max_value(
398 elements.len(),
399 self.config.max_elements,
400 "element count",
401 &format!("Image {}", idx),
402 )?;
403
404 let scores: Vec<f32> = elements.iter().map(|e| e.score).collect();
406 validator.validate_scores_with(&scores, |elem_idx| {
407 format!("Image {}, element {}", idx, elem_idx)
408 })?;
409 }
410
411 Ok(())
412 }
413
414 fn empty_output(&self) -> Self::Output {
415 LayoutDetectionOutput::empty()
416 }
417}
418
419#[cfg(test)]
420mod tests {
421 use super::*;
422 use crate::processors::Point;
423 use image::RgbImage;
424
425 #[test]
426 fn test_layout_detection_task_creation() {
427 let task = LayoutDetectionTask::default();
428 assert_eq!(task.task_type(), TaskType::LayoutDetection);
429 }
430
431 #[test]
432 fn test_input_validation() {
433 let task = LayoutDetectionTask::default();
434
435 let empty_input = ImageTaskInput::new(vec![]);
437 assert!(task.validate_input(&empty_input).is_err());
438
439 let valid_input = ImageTaskInput::new(vec![RgbImage::new(100, 100)]);
441 assert!(task.validate_input(&valid_input).is_ok());
442 }
443
444 #[test]
445 fn test_output_validation() {
446 let task = LayoutDetectionTask::default();
447
448 let box1 = BoundingBox::new(vec![
450 Point::new(0.0, 0.0),
451 Point::new(10.0, 0.0),
452 Point::new(10.0, 10.0),
453 Point::new(0.0, 10.0),
454 ]);
455 let element = LayoutDetectionElement {
456 bbox: box1,
457 element_type: "text".to_string(),
458 score: 0.95,
459 };
460 let output = LayoutDetectionOutput {
461 elements: vec![vec![element]],
462 is_reading_order_sorted: false,
463 };
464 assert!(task.validate_output(&output).is_ok());
465 }
466}