use std::sync::Mutex;
use docling_core::redact::{PiiDetector, PiiKind, Span};
use ort::session::Session;
use ort::value::Tensor;
use tokenizers::Tokenizer;
const MAX_TOKENS: usize = 512;
const CHUNK_BYTES: usize = 1200;
pub struct NerDetector {
session: Mutex<Session>,
tokenizer: Tokenizer,
labels: Vec<String>,
input_names: Vec<String>,
}
pub fn model_dir() -> String {
docling_core::env::nonempty("DOCLING_RS_NER_DIR")
.unwrap_or_else(|| crate::resolve_asset(".models/ner"))
}
pub fn models_available() -> bool {
let dir = model_dir();
let p = |n: &str| std::path::Path::new(&dir).join(n).exists();
(p("model.onnx") || p("model_int8.onnx")) && p("tokenizer.json") && p("config.json")
}
impl NerDetector {
pub fn load() -> Result<Self, String> {
let dir = model_dir();
let file = |n: &str| format!("{dir}/{n}");
let int8 = file("model_int8.onnx");
let model = if !crate::prefer_fp32() && std::path::Path::new(&int8).exists() {
int8
} else {
file("model.onnx")
};
for f in [&model, &file("tokenizer.json"), &file("config.json")] {
if !std::path::Path::new(f).exists() {
return Err(format!("NER model file not found: {f}"));
}
}
let config: serde_json::Value = serde_json::from_slice(
&std::fs::read(file("config.json")).map_err(|e| format!("ner: config.json: {e}"))?,
)
.map_err(|e| format!("ner: config.json: {e}"))?;
let id2label = config
.get("id2label")
.and_then(|v| v.as_object())
.ok_or("ner: config.json has no id2label")?;
let mut labels = vec![String::new(); id2label.len()];
for (k, v) in id2label {
let i: usize = k.parse().map_err(|_| format!("ner: id2label key {k:?}"))?;
if i >= labels.len() {
return Err(format!("ner: id2label index {i} out of range"));
}
labels[i] = v.as_str().unwrap_or("O").to_string();
}
let mut tokenizer = Tokenizer::from_file(file("tokenizer.json"))
.map_err(|e| format!("ner: tokenizer: {e}"))?;
tokenizer
.with_truncation(Some(tokenizers::TruncationParams {
max_length: MAX_TOKENS,
..Default::default()
}))
.map_err(|e| format!("ner: tokenizer truncation: {e}"))?;
let builder = docling_onnx::session_builder()?
.with_intra_threads(crate::intra_threads())
.map_err(|e| format!("ner: {e}"))?;
let builder = docling_onnx::apply(builder)?;
let session = docling_onnx::commit_uncached(builder, &model)
.map_err(|e| format!("ner: load {model}: {e}"))?;
let input_names = session
.inputs()
.iter()
.map(|i| i.name().to_string())
.collect();
Ok(Self {
session: Mutex::new(session),
tokenizer,
labels,
input_names,
})
}
fn detect_chunk(&self, chunk: &str) -> Result<Vec<Span>, String> {
let enc = self
.tokenizer
.encode(chunk, true)
.map_err(|e| format!("ner: tokenize: {e}"))?;
let ids: Vec<i64> = enc.get_ids().iter().map(|&v| v as i64).collect();
let n = ids.len();
if n == 0 {
return Ok(Vec::new());
}
let mask: Vec<i64> = enc.get_attention_mask().iter().map(|&v| v as i64).collect();
let types: Vec<i64> = enc.get_type_ids().iter().map(|&v| v as i64).collect();
let offsets = enc.get_offsets();
let specials = enc.get_special_tokens_mask();
let logits: Vec<f32> = {
let mut session = self.session.lock().unwrap_or_else(|p| p.into_inner());
let mut inputs: Vec<(String, ort::value::DynValue)> = Vec::new();
for name in &self.input_names {
let data = match name.as_str() {
"input_ids" => ids.clone(),
"attention_mask" => mask.clone(),
"token_type_ids" => types.clone(),
other => return Err(format!("ner: unexpected model input {other:?}")),
};
let t = Tensor::from_array(([1usize, n], data))
.map_err(|e| format!("ner: input {name}: {e}"))?;
inputs.push((name.clone(), t.into()));
}
let outputs = session.run(inputs).map_err(|e| format!("ner: run: {e}"))?;
let (_, data) = outputs[0]
.try_extract_tensor::<f32>()
.map_err(|e| format!("ner: output: {e}"))?;
data.to_vec()
};
let classes = self.labels.len();
if logits.len() != n * classes {
return Err(format!(
"ner: {} logits for {n} tokens × {classes} labels",
logits.len()
));
}
let mut spans: Vec<Span> = Vec::new();
let mut current: Option<(PiiKind, usize, usize, f32, usize)> = None; let flush = |cur: &mut Option<(PiiKind, usize, usize, f32, usize)>, out: &mut Vec<Span>| {
if let Some((kind, s, e, sum, cnt)) = cur.take() {
if e > s {
out.push(Span {
start: s,
end: e,
kind,
score: sum / cnt as f32,
name: None,
});
}
}
};
for t in 0..n {
if specials[t] != 0 {
continue;
}
let row = &logits[t * classes..(t + 1) * classes];
let (best, prob) = softmax_argmax(row);
let label = &self.labels[best];
let (prefix, kind) = match label.split_once('-') {
Some((p, "PER")) => (p, Some(PiiKind::Person)),
Some((p, "ORG")) => (p, Some(PiiKind::Organization)),
Some((p, "LOC")) => (p, Some(PiiKind::Location)),
_ => ("O", None),
};
let (ts, te) = offsets[t];
match (kind, current.as_mut()) {
(Some(k), Some(cur)) if prefix == "I" && cur.0 == k => {
cur.2 = te;
cur.3 += prob;
cur.4 += 1;
}
(Some(k), _) => {
flush(&mut current, &mut spans);
current = Some((k, ts, te, prob, 1));
}
(None, _) => flush(&mut current, &mut spans),
}
}
flush(&mut current, &mut spans);
for s in spans.iter_mut() {
s.start = snap_back(chunk, s.start);
s.end = snap_forward(chunk, s.end);
}
spans.retain(|s| {
let text = &chunk[s.start..s.end];
let letters = text.chars().filter(|c| c.is_alphabetic()).count();
!(text.chars().all(|c| !c.is_lowercase()) && letters <= 4 && !text.contains(' '))
});
Ok(spans)
}
}
fn is_word_char(c: char) -> bool {
c.is_alphanumeric() || c == '\''
}
fn snap_back(text: &str, mut i: usize) -> usize {
while i > 0 {
let Some(c) = text[..i].chars().next_back() else {
break;
};
if !is_word_char(c) {
break;
}
i -= c.len_utf8();
}
i
}
fn snap_forward(text: &str, mut i: usize) -> usize {
while i < text.len() {
let Some(c) = text[i..].chars().next() else {
break;
};
if !is_word_char(c) {
break;
}
i += c.len_utf8();
}
i
}
fn softmax_argmax(row: &[f32]) -> (usize, f32) {
let max = row.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
let sum: f32 = row.iter().map(|v| (v - max).exp()).sum();
let (best, v) =
row.iter().enumerate().fold(
(0, f32::NEG_INFINITY),
|acc, (i, &v)| if v > acc.1 { (i, v) } else { acc },
);
(best, (v - max).exp() / sum)
}
fn chunks(text: &str) -> Vec<(usize, &str)> {
let mut out = Vec::new();
let mut start = 0;
while start < text.len() {
let mut end = (start + CHUNK_BYTES).min(text.len());
if end < text.len() {
match text[start..end].rfind(char::is_whitespace) {
Some(i) if i > 0 => end = start + i,
_ => {
while !text.is_char_boundary(end) {
end += 1;
}
}
}
}
out.push((start, &text[start..end]));
start = end;
}
out
}
impl PiiDetector for NerDetector {
fn detect(&self, text: &str) -> Vec<Span> {
if !text.chars().any(|c| c.is_alphabetic()) {
return Vec::new();
}
let mut out = Vec::new();
for (base, chunk) in chunks(text) {
match self.detect_chunk(chunk) {
Ok(spans) => out.extend(spans.into_iter().map(|mut s| {
s.start += base;
s.end += base;
s
})),
Err(e) => {
eprintln!("warning: NER detection failed: {e}");
break;
}
}
}
out
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn chunks_split_on_whitespace_and_cover_the_text() {
let word = "wordpiece ";
let text: String = word.repeat(400);
let parts = chunks(&text);
assert!(parts.len() > 1);
let joined: String = parts.iter().map(|(_, s)| *s).collect();
assert_eq!(joined, text);
for (base, s) in &parts {
assert_eq!(&text[*base..*base + s.len()], *s);
assert!(s.len() <= CHUNK_BYTES);
}
assert_eq!(chunks(""), Vec::<(usize, &str)>::new());
}
#[test]
fn softmax_argmax_picks_the_top_class() {
let (i, p) = softmax_argmax(&[0.0, 2.0, 1.0]);
assert_eq!(i, 1);
assert!(p > 0.6 && p < 0.7, "{p}");
}
fn model_ready() -> bool {
if !models_available() {
let root = std::path::Path::new(env!("CARGO_MANIFEST_DIR")).join("../../.models/ner");
if root.join("model.onnx").exists() {
std::env::set_var("DOCLING_RS_NER_DIR", root);
}
}
models_available()
}
#[test]
fn snapping_covers_whole_words() {
let t = "see Angela Merkel, SSN";
assert_eq!(snap_back(t, 13), 11);
assert_eq!(snap_back(t, 11), 11);
assert_eq!(snap_forward(t, 15), 17);
assert_eq!(snap_forward(t, 17), 17);
assert_eq!(snap_back(t, 0), 0);
assert_eq!(snap_forward(t, t.len()), t.len());
}
#[test]
fn detects_seeded_entities() {
if !model_ready() {
eprintln!("skipping: NER model not found");
return;
}
let det = NerDetector::load().unwrap();
let text = "Contact john.doe@example.com or +1 (555) 123-4567 (Angela Merkel).";
let found: Vec<&str> = det
.detect(text)
.iter()
.map(|s| &text[s.start..s.end])
.collect();
assert_eq!(found, vec!["Angela Merkel"], "{found:?}");
let text = "Card 4111 1111 1111 1111, SSN 123-45-6789, IBAN DE89 3704 0044 0532 0130 00.";
assert!(
det.detect(text).is_empty(),
"abbreviations are not entities"
);
let text = "Café note: Angela Merkel met Siemens AG in Berlin on Monday.";
let spans = det.detect(text);
let found: Vec<(PiiKind, &str)> = spans
.iter()
.map(|s| (s.kind, &text[s.start..s.end]))
.collect();
assert!(
found.contains(&(PiiKind::Person, "Angela Merkel")),
"{found:?}"
);
assert!(found.contains(&(PiiKind::Location, "Berlin")), "{found:?}");
assert!(
found
.iter()
.any(|(k, t)| *k == PiiKind::Organization && t.starts_with("Siemens")),
"{found:?}"
);
assert!(spans.iter().all(|s| s.score > 0.5));
}
}