use anyhow::{Context, Result};
use std::path::Path;
use crate::EmbedEngine;
use crate::encoder_weights::EncBatch;
use crate::gliner::GlinerDevice;
use crate::weights::LazySt;
#[derive(Default)]
struct Linear {
w: Vec<f32>,
b: Vec<f32>,
n: usize,
k: usize,
packed: std::sync::OnceLock<crate::cpu_gemm::PackedWeight>,
}
impl Linear {
fn load(st: &LazySt, prefix: &str) -> Result<Self> {
let w = st.tensor_f32(&format!("{prefix}.weight"))?;
let b = st.tensor_f32(&format!("{prefix}.bias"))?;
let n = b.len();
let k = w.len() / n;
Ok(Self {
w,
b,
n,
k,
packed: std::sync::OnceLock::new(),
})
}
fn forward(&self, x: &[f32]) -> Vec<f32> {
let (n, k) = (self.n, self.k);
let m = x.len() / k;
let mut out = vec![0f32; m * n];
let packed = self
.packed
.get_or_init(|| crate::cpu_gemm::PackedWeight::new(&self.w, n, k));
crate::cpu_gemm::gemm_packed(&mut out, x, packed, m, Some(&self.b));
out
}
}
struct Mlp {
up: Linear,
down: Linear,
}
impl Mlp {
fn load(st: &LazySt, prefix: &str) -> Result<Self> {
Ok(Self {
up: Linear::load(st, &format!("{prefix}.0"))?,
down: Linear::load(st, &format!("{prefix}.3"))?,
})
}
fn forward(&self, x: &[f32]) -> Vec<f32> {
let mut h = self.up.forward(x);
for v in h.iter_mut() {
*v = v.max(0.0);
}
self.down.forward(&h)
}
}
struct SpanUp {
left: Linear,
right: Linear,
}
impl SpanUp {
fn load(st: &LazySt, prefix: &str) -> Result<Self> {
let up = Linear::load(st, prefix)?;
let (n, k) = (up.n, up.k);
let half = k / 2;
let mut left = Vec::with_capacity(n * half);
let mut right = Vec::with_capacity(n * half);
for row in up.w.chunks_exact(k) {
left.extend_from_slice(&row[..half]);
right.extend_from_slice(&row[half..]);
}
Ok(Self {
left: Linear {
w: left,
b: vec![0.0; n],
n,
k: half,
..Default::default()
},
right: Linear {
w: right,
b: up.b,
n,
k: half,
..Default::default()
},
})
}
}
struct LstmDir {
in_proj: Linear,
w_hh: Vec<f32>,
b_hh: Vec<f32>,
hidden: usize,
input: usize,
}
impl LstmDir {
fn load(st: &LazySt, prefix: &str, suffix: &str) -> Result<Self> {
let w_ih = st.tensor_f32(&format!("{prefix}.weight_ih_l0{suffix}"))?;
let w_hh = st.tensor_f32(&format!("{prefix}.weight_hh_l0{suffix}"))?;
let b_ih = st.tensor_f32(&format!("{prefix}.bias_ih_l0{suffix}"))?;
let b_hh = st.tensor_f32(&format!("{prefix}.bias_hh_l0{suffix}"))?;
let hidden = b_ih.len() / 4;
let input = w_ih.len() / (4 * hidden);
Ok(Self {
in_proj: Linear {
n: 4 * hidden,
k: input,
w: w_ih,
b: b_ih,
..Default::default()
},
w_hh,
b_hh,
hidden,
input,
})
}
fn run(&self, x: &[f32], order: impl Iterator<Item = usize>) -> Vec<f32> {
let (h_n, i_n) = (self.hidden, self.input);
let t = x.len() / i_n;
let xg = self.in_proj.forward(x);
let mut out = vec![0f32; t * h_n];
let mut h = vec![0f32; h_n];
let mut c = vec![0f32; h_n];
let mut gates = vec![0f32; 4 * h_n];
for step in order {
let xr = &xg[step * 4 * h_n..(step + 1) * 4 * h_n];
for (g, gate) in gates.iter_mut().enumerate() {
*gate = xr[g] + self.b_hh[g];
}
crate::simd::gemv_acc(&mut gates, &self.w_hh, &h);
for j in 0..h_n {
let i_g = sigmoid(gates[j]);
let f_g = sigmoid(gates[h_n + j]);
let g_g = gates[2 * h_n + j].tanh();
let o_g = sigmoid(gates[3 * h_n + j]);
c[j] = f_g * c[j] + i_g * g_g;
h[j] = o_g * c[j].tanh();
}
out[step * h_n..(step + 1) * h_n].copy_from_slice(&h);
}
out
}
}
struct BiLstm {
fwd: LstmDir,
rev: LstmDir,
hidden: usize,
}
impl BiLstm {
fn load(st: &LazySt, prefix: &str) -> Result<Self> {
let fwd = LstmDir::load(st, prefix, "")?;
let rev = LstmDir::load(st, prefix, "_reverse")?;
let hidden = fwd.hidden;
Ok(Self { fwd, rev, hidden })
}
fn forward(&self, x: &[f32]) -> Vec<f32> {
let t = x.len() / self.fwd.input;
let (f, r) = (self.fwd.run(x, 0..t), self.rev.run(x, (0..t).rev()));
let h = self.hidden;
let mut out = vec![0f32; t * 2 * h];
for step in 0..t {
out[step * 2 * h..step * 2 * h + h].copy_from_slice(&f[step * h..(step + 1) * h]);
out[step * 2 * h + h..(step + 1) * 2 * h].copy_from_slice(&r[step * h..(step + 1) * h]);
}
out
}
}
fn sigmoid(v: f32) -> f32 {
if v >= 0.0 {
1.0 / (1.0 + (-v).exp())
} else {
let e = v.exp();
e / (1.0 + e)
}
}
#[derive(Debug, Clone, PartialEq)]
pub struct JointEntity {
pub text: String,
pub label: String,
pub start: usize,
pub end: usize,
pub score: f32,
}
#[derive(Debug, Clone, PartialEq)]
pub struct JointRelation {
pub relation: String,
pub head_idx: usize,
pub tail_idx: usize,
pub score: f32,
}
pub struct GlinerRelex {
backbone: EmbedEngine,
projection: Linear, rnn: BiLstm,
project_start: Mlp,
project_end: Mlp,
span_up: SpanUp,
span_down: Linear,
prompt_rep: Mlp, pair_rep: Mlp, max_width: usize,
hidden: usize,
ent_token_id: u32,
rel_token_id: u32,
#[cfg(feature = "cli")]
tokenizer: tokenizers::Tokenizer,
#[cfg(feature = "cli")]
splitter: regex::Regex,
}
impl GlinerRelex {
pub fn load(dir: &Path) -> Result<Self> {
Self::load_on(dir, GlinerDevice::Cpu)
}
pub fn load_on(dir: &Path, device: GlinerDevice) -> Result<Self> {
let backbone = match device {
GlinerDevice::Auto => EmbedEngine::auto(dir, 8192)?,
GlinerDevice::Cpu => EmbedEngine::cpu(dir)?,
};
let st = LazySt::open(dir)?;
let head: serde_json::Value = std::fs::read(dir.join("glinerrelex_head.json"))
.ok()
.and_then(|b| serde_json::from_slice(&b).ok())
.unwrap_or(serde_json::Value::Null);
let hidden = head
.get("hidden_size")
.and_then(|x| x.as_u64())
.unwrap_or(768) as usize;
let max_width = head.get("max_width").and_then(|x| x.as_u64()).unwrap_or(12) as usize;
let ent_token_id = head
.get("class_token_index")
.and_then(|x| x.as_u64())
.unwrap_or(128001) as u32;
let rel_token_id = head
.get("rel_token_index")
.and_then(|x| x.as_u64())
.unwrap_or(128003) as u32;
let sp = "span_rep_layer.span_rep_layer";
Ok(Self {
projection: Linear::load(&st, "token_rep_layer.projection")?,
rnn: BiLstm::load(&st, "rnn.lstm")?,
project_start: Mlp::load(&st, &format!("{sp}.project_start"))?,
project_end: Mlp::load(&st, &format!("{sp}.project_end"))?,
span_up: SpanUp::load(&st, &format!("{sp}.out_project.0"))?,
span_down: Linear::load(&st, &format!("{sp}.out_project.3"))?,
prompt_rep: Mlp::load(&st, "prompt_rep_layer")?,
pair_rep: Mlp::load(&st, "pair_rep_layer")?,
max_width,
hidden,
ent_token_id,
rel_token_id,
backbone,
#[cfg(feature = "cli")]
tokenizer: tokenizers::Tokenizer::from_file(dir.join("tokenizer.json"))
.map_err(|e| anyhow::anyhow!("gliner-relex tokenizer: {e}"))?,
#[cfg(feature = "cli")]
splitter: regex::Regex::new(r"\w+(?:[-_]\w+)*|\S")?,
})
}
pub fn device(&self) -> String {
self.backbone.device()
}
#[cfg(feature = "cli")]
pub fn inference(
&mut self,
text: &str,
entity_labels: &[impl AsRef<str>],
relation_labels: &[impl AsRef<str>],
threshold: f32,
relation_threshold: f32,
) -> Result<(Vec<JointEntity>, Vec<JointRelation>)> {
let h = self.hidden;
let mut words: Vec<(String, usize, usize)> = Vec::new();
for m in self.splitter.find_iter(text) {
let cs = text[..m.start()].chars().count();
let ce = cs + m.as_str().chars().count();
words.push((m.as_str().to_string(), cs, ce));
}
let mut seq: Vec<String> = Vec::new();
for e in entity_labels {
seq.push("<<ENT>>".into());
seq.push(e.as_ref().to_string());
}
seq.push("<<SEP>>".into());
for r in relation_labels {
seq.push("<<REL>>".into());
seq.push(r.as_ref().to_string());
}
seq.push("<<SEP>>".into());
let prompt_len = seq.len();
seq.extend(words.iter().map(|(w, _, _)| w.clone()));
let enc = self
.tokenizer
.encode(tokenizers::InputSequence::from(seq), true)
.map_err(|e| anyhow::anyhow!("gliner-relex tokenize: {e}"))?;
let ids = enc.get_ids();
let word_ids = enc.get_word_ids();
let n_words = words.len();
let mut first_tok = vec![usize::MAX; n_words];
for (pos, wid) in word_ids.iter().enumerate() {
if let Some(wi) = wid {
let wi = *wi as usize;
if wi >= prompt_len {
let tw = wi - prompt_len;
if tw < n_words && first_tok[tw] == usize::MAX {
first_tok[tw] = pos;
}
}
}
}
let states = self
.backbone
.forward_hidden(&EncBatch::from_seqs([ids.to_vec()]))
.context("gliner-relex backbone")?;
let tokens = self.projection.forward(&states);
let mut ent_prompts: Vec<f32> = Vec::new();
let mut rel_prompts: Vec<f32> = Vec::new();
for (i, &id) in ids.iter().enumerate() {
if id == self.ent_token_id {
ent_prompts.extend_from_slice(&tokens[i * h..(i + 1) * h]);
} else if id == self.rel_token_id {
rel_prompts.extend_from_slice(&tokens[i * h..(i + 1) * h]);
}
}
let c_ent = ent_prompts.len() / h;
let c_rel = rel_prompts.len() / h;
let mut words_emb = vec![0f32; n_words * h];
for (wi, &ft) in first_tok.iter().enumerate() {
if ft != usize::MAX {
words_emb[wi * h..(wi + 1) * h].copy_from_slice(&tokens[ft * h..(ft + 1) * h]);
}
}
let words_emb = self.rnn.forward(&words_emb); let ent_reps = self.prompt_rep.forward(&ent_prompts);
let start = self.project_start.forward(&words_emb);
let end = self.project_end.forward(&words_emb);
let relu = |v: &[f32]| -> Vec<f32> { v.iter().map(|x| x.max(0.0)).collect() };
let a = self.span_up.left.forward(&relu(&start));
let b = self.span_up.right.forward(&relu(&end));
let up_w = self.span_up.left.n;
let valid: Vec<(usize, usize)> = (0..n_words)
.flat_map(|l| (0..self.max_width).map(move |k| (l, k)))
.filter(|&(l, k)| l + k < n_words)
.collect();
let mut hbuf = vec![0f32; valid.len() * up_w];
for (row, &(l, k)) in valid.iter().enumerate() {
for j in 0..up_w {
hbuf[row * up_w + j] = (a[l * up_w + j] + b[(l + k) * up_w + j]).max(0.0);
}
}
let span_reps = self.span_down.forward(&hbuf);
let ner_scorer = Linear {
w: ent_reps,
b: vec![0.0; c_ent],
n: c_ent,
k: h,
..Default::default()
};
let mut ner = ner_scorer.forward(&span_reps); for v in ner.iter_mut() {
*v = sigmoid(*v);
}
let mut cands: Vec<(usize, usize, usize, f32)> = Vec::new(); for (row, &(l, k)) in valid.iter().enumerate() {
for c in 0..c_ent {
let s = ner[row * c_ent + c];
if s > threshold {
cands.push((l, l + k, c, s));
}
}
}
let mut order: Vec<usize> = (0..cands.len()).collect();
order.sort_by(|&a, &b| {
cands[b]
.3
.partial_cmp(&cands[a].3)
.unwrap_or(std::cmp::Ordering::Equal)
});
let mut kept: Vec<(usize, usize, usize, f32)> = Vec::new();
for &oi in &order {
let cnd = cands[oi];
if !kept
.iter()
.any(|k| overlaps_nested((cnd.0, cnd.1), (k.0, k.1)))
{
kept.push(cnd);
}
}
kept.sort_by_key(|c| c.0);
let orig: Vec<char> = text.chars().collect();
let entities: Vec<JointEntity> = kept
.iter()
.map(|&(s, e, c, sc)| {
let cs = words[s].1;
let ce = words[e].2;
JointEntity {
text: orig
.get(cs..ce)
.map(|c| c.iter().collect())
.unwrap_or_default(),
label: entity_labels[c].as_ref().to_string(),
start: cs,
end: ce,
score: sc,
}
})
.collect();
let decoded_map: std::collections::HashMap<(usize, usize), usize> = kept
.iter()
.enumerate()
.map(|(i, &(s, e, _, _))| ((s, e), i))
.collect();
let mut selected: Vec<(usize, usize, usize)> = Vec::new(); for (row, &(l, k)) in valid.iter().enumerate() {
let maxs = (0..c_ent)
.map(|c| ner[row * c_ent + c])
.fold(0.0f32, f32::max);
if maxs > threshold {
selected.push((l, l + k, row));
}
}
let mut relations: Vec<JointRelation> = Vec::new();
if c_rel > 0 && selected.len() >= 2 {
let rel_scorer = Linear {
w: rel_prompts,
b: vec![0.0; c_rel],
n: c_rel,
k: h,
..Default::default()
};
for i in 0..selected.len() {
for j in 0..selected.len() {
if i == j {
continue;
}
let (hs, he, hr) = selected[i];
let (ts, te, tr) = selected[j];
let mut cat = Vec::with_capacity(2 * h);
cat.extend_from_slice(&span_reps[hr * h..(hr + 1) * h]);
cat.extend_from_slice(&span_reps[tr * h..(tr + 1) * h]);
let pr = self.pair_rep.forward(&cat); let scores = rel_scorer.forward(&pr); let (Some(&hi), Some(&ti)) =
(decoded_map.get(&(hs, he)), decoded_map.get(&(ts, te)))
else {
continue; };
for c in 0..c_rel {
let s = sigmoid(scores[c]);
if s > relation_threshold {
relations.push(JointRelation {
relation: relation_labels[c].as_ref().to_string(),
head_idx: hi,
tail_idx: ti,
score: s,
});
}
}
}
}
}
Ok((entities, relations))
}
}
fn overlaps_nested(a: (usize, usize), b: (usize, usize)) -> bool {
if a == b {
return true;
}
if a.0 > b.1 || b.0 > a.1 {
return false; }
if (a.0 <= b.0 && b.1 <= a.1) || (b.0 <= a.0 && a.1 <= b.1) {
return false; }
true }