use anyhow::{Context, Result};
use std::path::Path;
use crate::weights::LazySt;
#[derive(Default)]
struct Linear {
w: Vec<f32>,
b: Option<Vec<f32>>,
n: usize,
k: usize,
packed: std::sync::OnceLock<crate::cpu_gemm::PackedWeight>,
}
impl Linear {
fn load(st: &LazySt, prefix: &str, bias: bool) -> Result<Self> {
let w = st.tensor_f32(&format!("{prefix}.weight"))?;
let b = if bias {
Some(st.tensor_f32(&format!("{prefix}.bias"))?)
} else {
None
};
let n = if let Some(b) = &b { b.len() } else { 0 };
let (n, k) = if n > 0 { (n, w.len() / n) } else { (0, 0) };
Ok(Self {
w,
b,
n,
k,
packed: std::sync::OnceLock::new(),
})
}
fn from_parts(w: Vec<f32>, b: Option<Vec<f32>>, n: usize, k: usize) -> Self {
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, self.b.as_deref());
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"), true)?,
down: Linear::load(st, &format!("{prefix}.3"), true)?,
})
}
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)
}
}
fn layer_norm(x: &mut [f32], h: usize, w: &[f32], b: &[f32], eps: f32) {
for row in x.chunks_exact_mut(h) {
let mean = row.iter().sum::<f32>() / h as f32;
let var = row.iter().map(|v| (v - mean) * (v - mean)).sum::<f32>() / h as f32;
let inv = 1.0 / (var + eps).sqrt();
for (j, v) in row.iter_mut().enumerate() {
*v = (*v - mean) * inv * w[j] + b[j];
}
}
}
fn gelu(x: &mut [f32]) {
const INV_SQRT2: f32 = std::f32::consts::FRAC_1_SQRT_2;
for v in x.iter_mut() {
*v = 0.5 * *v * (1.0 + libm::erff(*v * INV_SQRT2));
}
}
fn softmax_row(row: &mut [f32]) {
let m = row.iter().copied().fold(f32::NEG_INFINITY, f32::max);
let mut s = 0.0;
for v in row.iter_mut() {
*v = (*v - m).exp();
s += *v;
}
let inv = 1.0 / s;
for v in row.iter_mut() {
*v *= inv;
}
}
struct V1Layer {
qw: Linear, kw: Linear, vw: Linear, pos_proj: Linear, pos_q_proj: Linear, att_dense: Linear,
att_ln_w: Vec<f32>,
att_ln_b: Vec<f32>,
inter: Linear,
out_dense: Linear,
out_ln_w: Vec<f32>,
out_ln_b: Vec<f32>,
}
struct DebertaV1 {
word_emb: Vec<f32>, emb_ln_w: Vec<f32>,
emb_ln_b: Vec<f32>,
rel_emb: Vec<f32>, layers: Vec<V1Layer>,
h: usize,
heads: usize,
hd: usize,
max_pos: usize, eps: f32,
}
impl DebertaV1 {
fn load(
st: &LazySt,
prefix: &str,
h: usize,
heads: usize,
n_layers: usize,
max_pos: usize,
eps: f32,
) -> Result<Self> {
let hd = h / heads;
let word_emb = st.tensor_f32(&format!("{prefix}.embeddings.word_embeddings.weight"))?;
let emb_ln_w = st.tensor_f32(&format!("{prefix}.embeddings.LayerNorm.weight"))?;
let emb_ln_b = st.tensor_f32(&format!("{prefix}.embeddings.LayerNorm.bias"))?;
let rel_emb = st.tensor_f32(&format!("{prefix}.encoder.rel_embeddings.weight"))?;
let mut layers = Vec::with_capacity(n_layers);
for i in 0..n_layers {
let lp = format!("{prefix}.encoder.layer.{i}");
let inw = st.tensor_f32(&format!("{lp}.attention.self.in_proj.weight"))?;
let (mut qw, mut kw, mut vw) =
(vec![0f32; h * h], vec![0f32; h * h], vec![0f32; h * h]);
for g in 0..heads {
for d in 0..hd {
let out_row = g * hd + d; let src_q = (g * 3 * hd + d) * h;
let src_k = (g * 3 * hd + hd + d) * h;
let src_v = (g * 3 * hd + 2 * hd + d) * h;
qw[out_row * h..out_row * h + h].copy_from_slice(&inw[src_q..src_q + h]);
kw[out_row * h..out_row * h + h].copy_from_slice(&inw[src_k..src_k + h]);
vw[out_row * h..out_row * h + h].copy_from_slice(&inw[src_v..src_v + h]);
}
}
let q_bias = st.tensor_f32(&format!("{lp}.attention.self.q_bias"))?;
let v_bias = st.tensor_f32(&format!("{lp}.attention.self.v_bias"))?;
let pos_proj = {
let w = st.tensor_f32(&format!("{lp}.attention.self.pos_proj.weight"))?;
Linear::from_parts(w, None, h, h)
};
let pos_q_proj = {
let w = st.tensor_f32(&format!("{lp}.attention.self.pos_q_proj.weight"))?;
let b = st.tensor_f32(&format!("{lp}.attention.self.pos_q_proj.bias"))?;
Linear::from_parts(w, Some(b), h, h)
};
layers.push(V1Layer {
qw: Linear::from_parts(qw, Some(q_bias), h, h),
kw: Linear::from_parts(kw, None, h, h),
vw: Linear::from_parts(vw, Some(v_bias), h, h),
pos_proj,
pos_q_proj,
att_dense: Linear::load(st, &format!("{lp}.attention.output.dense"), true)?,
att_ln_w: st.tensor_f32(&format!("{lp}.attention.output.LayerNorm.weight"))?,
att_ln_b: st.tensor_f32(&format!("{lp}.attention.output.LayerNorm.bias"))?,
inter: Linear::load(st, &format!("{lp}.intermediate.dense"), true)?,
out_dense: Linear::load(st, &format!("{lp}.output.dense"), true)?,
out_ln_w: st.tensor_f32(&format!("{lp}.output.LayerNorm.weight"))?,
out_ln_b: st.tensor_f32(&format!("{lp}.output.LayerNorm.bias"))?,
});
}
Ok(Self {
word_emb,
emb_ln_w,
emb_ln_b,
rel_emb,
layers,
h,
heads,
hd,
max_pos,
eps,
})
}
fn forward(&self, ids: &[u32]) -> Vec<f32> {
let (h, t) = (self.h, ids.len());
let mut x = vec![0f32; t * h];
for (i, &id) in ids.iter().enumerate() {
x[i * h..(i + 1) * h]
.copy_from_slice(&self.word_emb[id as usize * h..(id as usize + 1) * h]);
}
layer_norm(&mut x, h, &self.emb_ln_w, &self.emb_ln_b, self.eps);
let span = t.min(self.max_pos);
let rel_slice: Vec<f32> =
self.rel_emb[(self.max_pos - span) * h..(self.max_pos + span) * h].to_vec();
let two_span = 2 * span;
for layer in &self.layers {
let ctx = self.attention(layer, &x, t, &rel_slice, span, two_span);
let mut ao = layer.att_dense.forward(&ctx);
for (o, r) in ao.iter_mut().zip(x.iter()) {
*o += *r;
}
layer_norm(&mut ao, h, &layer.att_ln_w, &layer.att_ln_b, self.eps);
let mut inter = layer.inter.forward(&ao);
gelu(&mut inter);
let mut out = layer.out_dense.forward(&inter);
for (o, r) in out.iter_mut().zip(ao.iter()) {
*o += *r;
}
layer_norm(&mut out, h, &layer.out_ln_w, &layer.out_ln_b, self.eps);
x = out;
}
x
}
fn attention(
&self,
l: &V1Layer,
x: &[f32],
t: usize,
rel_slice: &[f32],
span: usize,
two_span: usize,
) -> Vec<f32> {
let (h, heads, hd) = (self.h, self.heads, self.hd);
let mut q = l.qw.forward(x); let k = l.kw.forward(x); let v = l.vw.forward(x); let scale = ((hd as f32) * 3.0).sqrt();
for vq in q.iter_mut() {
*vq /= scale;
}
let pos_key = l.pos_proj.forward(rel_slice);
let mut pos_query = l.pos_q_proj.forward(rel_slice);
for vq in pos_query.iter_mut() {
*vq /= scale;
}
let mut ctx = vec![0f32; t * h];
for hh in 0..heads {
let ho = hh * hd;
let mut scores = vec![0f32; t * t];
for i in 0..t {
let qi = &q[i * h + ho..i * h + ho + hd];
for j in 0..t {
let kj = &k[j * h + ho..j * h + ho + hd];
let mut s = 0.0f32;
for d in 0..hd {
s += qi[d] * kj[d];
}
let m = ((i as isize - j as isize) + span as isize)
.clamp(0, two_span as isize - 1) as usize;
let pk = &pos_key[m * h + ho..m * h + ho + hd];
let pq = &pos_query[m * h + ho..m * h + ho + hd];
let mut c2p = 0.0f32;
let mut p2c = 0.0f32;
for d in 0..hd {
c2p += qi[d] * pk[d]; p2c += kj[d] * pq[d]; }
scores[i * t + j] = s + c2p + p2c;
}
}
for i in 0..t {
softmax_row(&mut scores[i * t..(i + 1) * t]);
}
for i in 0..t {
let ci = &mut ctx[i * h + ho..i * h + ho + hd];
for j in 0..t {
let p = scores[i * t + j];
let vj = &v[j * h + ho..j * h + ho + hd];
for d in 0..hd {
ci[d] += p * vj[d];
}
}
}
}
ctx
}
}
#[derive(Debug, Clone, PartialEq)]
pub struct LinkEntity {
pub text: String,
pub label: String,
pub start: usize,
pub end: usize,
pub score: f32,
}
pub struct GlinkerIntermediates {
pub words_embedding: Vec<f32>, pub label_emb: Vec<f32>, pub scores: Vec<f32>, pub span_idx: Vec<(usize, usize)>,
pub span_rep: Vec<f32>, pub span_logits: Vec<f32>, pub entities: Vec<LinkEntity>,
pub w: usize,
pub c: usize,
}
pub struct Glinker {
text_enc: DebertaV1,
label_enc: DebertaV1,
proj_token: Linear, proj_label: Linear,
out_mlp0: Linear, out_mlp3: Linear, project_start: Mlp,
project_end: Mlp,
span_up: Linear, span_down: Linear, h: usize,
#[cfg(feature = "cli")]
text_tok: tokenizers::Tokenizer,
#[cfg(feature = "cli")]
label_tok: tokenizers::Tokenizer,
#[cfg(feature = "cli")]
splitter: regex::Regex,
}
impl Glinker {
pub fn load(dir: &Path) -> Result<Self> {
let cfg: serde_json::Value = serde_json::from_slice(
&std::fs::read(dir.join("config.json")).context("config.json")?,
)?;
let g = |k: &str| cfg.get(k).and_then(|x| x.as_u64()).map(|v| v as usize);
let h = g("hidden_size").context("hidden_size")?;
let heads = g("num_attention_heads").context("num_attention_heads")?;
let n_layers = g("num_hidden_layers").context("num_hidden_layers")?;
let max_pos = g("max_position_embeddings").unwrap_or(512);
let eps = cfg
.get("layer_norm_eps")
.and_then(|x| x.as_f64())
.unwrap_or(1e-7) as f32;
let st = LazySt::open(dir)?;
let text_enc = DebertaV1::load(
&st,
"token_rep_layer.bert_layer.model",
h,
heads,
n_layers,
max_pos,
eps,
)?;
let label_enc = DebertaV1::load(
&st,
"token_rep_layer.labels_encoder.model",
h,
heads,
n_layers,
max_pos,
eps,
)?;
let sp = "span_rep_layer.span_rep_layer";
Ok(Self {
text_enc,
label_enc,
proj_token: Linear::load(&st, "scorer.proj_token", true)?,
proj_label: Linear::load(&st, "scorer.proj_label", true)?,
out_mlp0: Linear::load(&st, "scorer.out_mlp.0", true)?,
out_mlp3: Linear::load(&st, "scorer.out_mlp.3", true)?,
project_start: Mlp::load(&st, &format!("{sp}.project_start"))?,
project_end: Mlp::load(&st, &format!("{sp}.project_end"))?,
span_up: Linear::load(&st, &format!("{sp}.out_project.0"), true)?,
span_down: Linear::load(&st, &format!("{sp}.out_project.3"), true)?,
h,
#[cfg(feature = "cli")]
text_tok: tokenizers::Tokenizer::from_file(dir.join("tokenizer.json"))
.map_err(|e| anyhow::anyhow!("glinker text tokenizer: {e}"))?,
#[cfg(feature = "cli")]
label_tok: tokenizers::Tokenizer::from_file(
dir.join("labels_tokenizer/tokenizer.json"),
)
.map_err(|e| anyhow::anyhow!("glinker label tokenizer: {e}"))?,
#[cfg(feature = "cli")]
splitter: regex::Regex::new(r"\w+(?:[-_]\w+)*|\S")?,
})
}
#[cfg(feature = "cli")]
fn encode_label(&self, label: &str) -> Vec<f32> {
let enc = self.label_tok.encode(label, true).expect("label tokenize");
let ids = enc.get_ids();
let hs = self.label_enc.forward(ids); let (h, t) = (self.h, ids.len());
let mut m = vec![0f32; h];
for row in hs.chunks_exact(h) {
for (a, b) in m.iter_mut().zip(row) {
*a += *b;
}
}
for a in m.iter_mut() {
*a /= t as f32;
}
m
}
#[cfg(feature = "cli")]
pub fn predict_entities(
&self,
text: &str,
labels: &[impl AsRef<str>],
threshold: f32,
) -> Vec<LinkEntity> {
self.predict_debug(text, labels, threshold).0
}
#[cfg(feature = "cli")]
pub fn predict_debug(
&self,
text: &str,
labels: &[impl AsRef<str>],
threshold: f32,
) -> (Vec<LinkEntity>, GlinkerIntermediates) {
let h = self.h;
let mut seen = std::collections::HashSet::new();
let labels: Vec<String> = labels
.iter()
.map(|s| s.as_ref().to_string())
.filter(|s| seen.insert(s.clone()))
.collect();
let c = labels.len();
let orig: Vec<char> = text.chars().collect();
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 word_strs: Vec<&str> = words.iter().map(|(w, _, _)| w.as_str()).collect();
let enc = self
.text_tok
.encode(word_strs, true)
.expect("text tokenize");
let ids = enc.get_ids();
let word_ids = enc.get_word_ids();
let w = words.len();
let mut first_tok = vec![usize::MAX; w];
for (pos, wid) in word_ids.iter().enumerate() {
if let Some(wi) = wid {
let wi = *wi as usize;
if wi < w && first_tok[wi] == usize::MAX {
first_tok[wi] = pos;
}
}
}
let hs = self.text_enc.forward(ids); let mut words_embedding = vec![0f32; w * h];
for wi in 0..w {
let p = first_tok[wi];
words_embedding[wi * h..(wi + 1) * h].copy_from_slice(&hs[p * h..(p + 1) * h]);
}
let mut label_emb = vec![0f32; c * h];
for (ci, lab) in labels.iter().enumerate() {
label_emb[ci * h..(ci + 1) * h].copy_from_slice(&self.encode_label(lab));
}
let tp = self.proj_token.forward(&words_embedding); let lp = self.proj_label.forward(&label_emb); let two_h = 2 * h;
let mut cat = vec![0f32; w * c * 3 * h];
for wi in 0..w {
for ci in 0..c {
let base = (wi * c + ci) * 3 * h;
cat[base..base + h].copy_from_slice(&tp[wi * two_h..wi * two_h + h]);
cat[base + h..base + 2 * h].copy_from_slice(&lp[ci * two_h..ci * two_h + h]);
for d in 0..h {
cat[base + 2 * h + d] = tp[wi * two_h + h + d] * lp[ci * two_h + h + d];
}
}
}
let mut sc = self.out_mlp0.forward(&cat); for v in sc.iter_mut() {
*v = v.max(0.0);
}
let scores = self.out_mlp3.forward(&sc);
let sig = |x: f32| 1.0 / (1.0 + (-x).exp());
let mut start_m = vec![false; w * c];
let mut end_m = vec![false; w * c];
let mut inside_m = vec![false; w * c];
for wi in 0..w {
for ci in 0..c {
let s = &scores[(wi * c + ci) * 3..(wi * c + ci) * 3 + 3];
start_m[wi * c + ci] = sig(s[0]) > threshold;
end_m[wi * c + ci] = sig(s[1]) > threshold;
inside_m[wi * c + ci] = sig(s[2]) > threshold;
}
}
let starts: Vec<(usize, usize)> = (0..w)
.flat_map(|p| (0..c).map(move |cl| (p, cl)))
.filter(|&(p, cl)| start_m[p * c + cl])
.collect();
let ends: Vec<(usize, usize)> = (0..w)
.flat_map(|p| (0..c).map(move |cl| (p, cl)))
.filter(|&(p, cl)| end_m[p * c + cl])
.collect();
let mut span_idx: Vec<(usize, usize)> = Vec::new();
for &(sp, scl) in &starts {
for &(ep, ecl) in &ends {
if scl == ecl && sp <= ep {
let covered = (sp..=ep).all(|p| inside_m[p * c + scl]);
if covered {
span_idx.push((sp, ep));
}
}
}
}
let start_rep = self.project_start.forward(&words_embedding); let end_rep = self.project_end.forward(&words_embedding); let n = span_idx.len();
let mut mcat = vec![0f32; n * two_h];
for (ni, &(s, e)) in span_idx.iter().enumerate() {
for d in 0..h {
mcat[ni * two_h + d] = start_rep[s * h + d].max(0.0);
mcat[ni * two_h + h + d] = end_rep[e * h + d].max(0.0);
}
}
let mut mup = self.span_up.forward(&mcat); for v in mup.iter_mut() {
*v = v.max(0.0);
}
let span_rep = self.span_down.forward(&mup); let mut span_logits = vec![0f32; n * c];
for ni in 0..n {
for ci in 0..c {
let mut acc = 0.0f32;
for d in 0..h {
acc += span_rep[ni * h + d] * label_emb[ci * h + d];
}
span_logits[ni * c + ci] = acc;
}
}
let mut spans: Vec<(usize, usize, usize, f32)> = Vec::new(); for ni in 0..n {
let (s, e) = span_idx[ni];
for ci in 0..c {
let p = sig(span_logits[ni * c + ci]);
if p > threshold {
spans.push((s, e, ci, p));
}
}
}
let mut order: Vec<usize> = (0..spans.len()).collect();
order.sort_by(|&a, &b| {
spans[b]
.3
.partial_cmp(&spans[a].3)
.unwrap_or(std::cmp::Ordering::Equal)
});
let mut kept: Vec<(usize, usize, usize, f32)> = Vec::new();
for &oi in &order {
let (s, e, cl, sco) = spans[oi];
let overlap = kept.iter().any(|&(ks, ke, _, _)| {
if ks == s && ke == e {
true
} else {
!(s > ke || ks > e)
}
});
if !overlap {
kept.push((s, e, cl, sco));
}
}
kept.sort_by_key(|&(s, _, _, _)| s);
let entities: Vec<LinkEntity> = kept
.iter()
.map(|&(s, e, cl, sco)| {
let cs = words[s].1;
let ce = words[e].2;
LinkEntity {
text: orig
.get(cs..ce)
.map(|c| c.iter().collect())
.unwrap_or_default(),
label: labels[cl].clone(),
start: cs,
end: ce,
score: sco,
}
})
.collect();
let inter = GlinkerIntermediates {
words_embedding,
label_emb,
scores,
span_idx,
span_rep,
span_logits,
entities: entities.clone(),
w,
c,
};
(entities, inter)
}
}